import hashlib import json from datetime import datetime, timedelta from threading import Lock, Thread from types import ModuleType from typing import List, Tuple from requests.exceptions import ConnectionError, HTTPError, ProxyError from smelog.factory import BoundLogger from vendor_image_caching import config from vendor_image_caching.constants import CollectionName, ObjectName, RecordField from vendor_image_caching.json_helper import date_hook from vendor_image_caching.s3 import Client as S3Client class RecordContainer: def __init__(self, records: List[dict]): self._records = list(records) self._lock = Lock() def pop(self): with self._lock: if not self._records: return None return self._records.pop() def make_singular(object_name: str or None) -> str or None: """Make object name singular. Args: object_name: Object name, can be plural. Returns: Singular object name or None. """ return object_name[:-1] if object_name and object_name[-1] == "s" else object_name def get_images( data: dict or list, current_node: str = None, inner_path: tuple = None ) -> List[Tuple[str, str, tuple, dict]]: """Get all images dict from data. Args: data: Data. current_node: Current element. inner_path: Path within original data. Returns: List of tuples with object ID, name, inner path and image dict. """ result = [] if isinstance(data, dict): if isinstance(data.get("images"), list): singular_node = make_singular(current_node) result.extend( (data.get("id"), singular_node, inner_path, record) for record in data["images"] if "url" in record ) for key, value in data.items(): if isinstance(value, (dict, list)): current_path = list(inner_path) if inner_path else [] current_path.append(key) result.extend(get_images(value, current_node if key == "items" else key, tuple(current_path))) elif isinstance(data, list): for index, item in enumerate(data): current_path = list(inner_path) if inner_path else [] current_path.append(str(index)) result.extend(get_images(item, current_node, tuple(current_path))) return result def get_images_from_record( collection_name: str, record_id: str, data: dict or list ) -> List[Tuple[str, str, tuple or None, dict]]: """Get images from cache record data. There are endpoints to get images directly like /v1/playlists/{playlist_id}/images, in this case (the mentioned endpoint + cache collection) we need to parse record_id to get object name and ID. Args: collection_name: Collection name. record_id: Cache record ID. data: Cache data. Returns: List of tuples with object ID, name, inner path and image dict. """ url_parts = get_url_parts(record_id) if isinstance(data, list) and url_parts[-1] == "images" and len(url_parts) >= 3: obj_id = url_parts[-2] obj_name = make_singular(url_parts[-3]) return [(obj_id, obj_name, None, image) for image in data] return get_images(data, get_object_name(collection_name, record_id)) def get_key( collection_name: str, record_id: str, obj_id: str, obj_node: str, inner_path: tuple, image_obj: dict ) -> str: """Get image name for caching. Args: collection_name: Collection name. record_id: Mongo record ID. obj_id: Image containing object ID. obj_node: Image containing object node or truncated collection name. inner_path: Path to parent image node. image_obj: Dict with image data. Returns: Image name. """ img_size = f"{image_obj.get('width')}x{image_obj.get('height')}" if config.FORCE_UNIQUE_IMAGE_NAME or not obj_id or not obj_node: img_name_base = f"{record_id}{'_' + '_'.join(inner_path) if inner_path else ''}" img_name = hashlib.sha3_256(img_name_base.encode("utf-8")).hexdigest() return f"{collection_name}/{img_name}_{img_size}" return f"{obj_node}/{obj_id}_{img_size}" def get_url_parts(record_id: str) -> List[str]: """Parse record ID to get URL parts for cache collection. Args: record_id: Record ID. Returns: URL parts. """ url = record_id.split("_")[0] return url.rstrip("/").split("/") def get_object_name_from_record_id(record_id: str) -> str: """Get object name from record ID for cache collection. Args: record_id: Cache record ID. Returns: Object name. """ url_parts = get_url_parts(record_id) obj_name = url_parts[-1] if not obj_name.isalpha() and len(url_parts) > 1: obj_name = url_parts[-2] return obj_name def get_object_name(collection_name: str, record_id: str) -> str or None: """Get object name by collection name and record ID. Args: collection_name: Collection name. record_id: Cache record ID. Returns: Object name or None. """ if collection_name in CollectionName.OBJECTS: return CollectionName.OBJECTS[collection_name] if collection_name == CollectionName.CACHE: obj_name = get_object_name_from_record_id(record_id) else: obj_name = collection_name.split("_")[-1] obj_name = make_singular(obj_name) if obj_name in ObjectName.ALL: return obj_name return None def worker_thread(records_container: RecordContainer, logger: BoundLogger, cache_backend: ModuleType): """Worker thread code. Args: records_container: Records to process, thread safe. logger: Logger. cache_backend: cache backend module """ cache_client = cache_backend.Client() s3_client = S3Client() while True: record = records_container.pop() if not record: break logger.debug(f"Processing {record}") collection_name: str = record[RecordField.COLLECTION_NAME] record_id: str = record[RecordField.RECORD_ID] created_at: datetime = record[RecordField.CREATED_AT] if (datetime.utcnow() - timedelta(seconds=config.MAX_DATA_TIMEOUT)) >= created_at: logger.debug("Too late to cache images for this record") continue cache_record = cache_client.get_document(collection_name, record_id, created_at) if not cache_record: logger.debug("Nothing to process here") continue images = get_images_from_record(collection_name, record_id, cache_record) try: for obj_id, obj_node, inner_path, image_obj in images: image_name = get_key(collection_name, record_id, obj_id, obj_node, inner_path, image_obj) storage_url = s3_client.upload_image(image_obj["url"], image_name) image_obj["url"] = storage_url logger.debug(f"{image_name} processed") except (ConnectionError, HTTPError, ProxyError) as e: logger.error(e) continue cache_client.save_document(collection_name, record_id, created_at, cache_record) def handler(logger: BoundLogger, event: dict, cache_backend: ModuleType): """Lambda job. Args: logger: Logger instance. event: Lambda event obj. cache_backend: cache backend module Returns: JSON serializable response """ records = [json.loads(r["body"], object_hook=date_hook) for r in event["Records"]] logger.debug(f"Got {len(records)} to process, initializing {config.WORKER_COUNT} workers") records_container = RecordContainer(records) if config.WORKER_COUNT > 1: threads = [ Thread(target=worker_thread, args=(records_container, logger, cache_backend)) for _ in range(config.WORKER_COUNT) ] [thread.start() for thread in threads] [thread.join() for thread in threads] else: worker_thread(records_container, logger, cache_backend)