"""DynamoDB connector.""" import time from typing import Any import boto3 from boto3.dynamodb.types import TypeDeserializer, TypeSerializer from mypy_boto3_dynamodb.client import DynamoDBClient from assets import config DDB_BATCH_LIMIT = 100 class DDBClient: """Shared DynamoDB Resource.""" _client: DynamoDBClient | None = None _region: str _serializer: TypeSerializer _deserializer: TypeDeserializer def __init__(self, region: str) -> None: """Init.""" self._region = region self._serializer = TypeSerializer() self._deserializer = TypeDeserializer() def batch_get_items( self, table_name: str, table_key: str, item_ids: list[Any], projection_expression: str, consistent_read: bool = False, ) -> list[dict[str, Any]]: """Get multiple items by id.""" key_chunks = ( item_ids[i : i + DDB_BATCH_LIMIT] for i in range(0, len(item_ids), DDB_BATCH_LIMIT) ) all_items = [] for keys_chunk in key_chunks: items_chunk = self._get_items( table_name, table_key, keys_chunk, projection_expression, consistent_read, ) all_items.extend(items_chunk) return all_items def _get_items( self, table_name: str, table_key: str, item_ids: list[Any], projection_expression: str, consistent_read: bool = False, ) -> list[dict[str, Any]]: if self._client is None: self._client = boto3.session.Session().client( "dynamodb", region_name=self._region ) items_batch = [] processing_keys = list(item_ids) while processing_keys: result = self._client.batch_get_item( RequestItems={ table_name: { "Keys": [ self.python_to_dynamo({table_key: item_id}) for item_id in processing_keys ], "ConsistentRead": consistent_read, "ProjectionExpression": projection_expression, } } ) processing_keys = self.get_unprocessed_keys(result, table_name, table_key) # type: ignore[arg-type] result_items = [ self.dynamo_to_python(item) for item in result["Responses"][table_name] ] items_batch.extend(result_items) if processing_keys: time.sleep(0.3) return items_batch def python_to_dynamo(self, python_object: dict[str, Any]) -> dict[str, Any]: """Convert python dict to dynamo syntax.""" return {k: self._serializer.serialize(v) for k, v in python_object.items()} def dynamo_to_python(self, dynamo_object: dict[str, Any]) -> dict[str, Any]: """Convert dynamo syntax to python dict.""" return {k: self._deserializer.deserialize(v) for k, v in dynamo_object.items()} def get_unprocessed_keys( self, result: dict[str, Any], table_name: str, table_key: str ) -> list[Any]: """Get unprocessed keys.""" if ( "UnprocessedKeys" not in result or not result["UnprocessedKeys"] or table_name not in result["UnprocessedKeys"] or not result["UnprocessedKeys"][table_name] or "Keys" not in result["UnprocessedKeys"][table_name] or not result["UnprocessedKeys"][table_name]["Keys"] ): return [] return [ self.dynamo_to_python(item)[table_key] for item in result["UnprocessedKeys"][table_name]["Keys"] ] ddb_connector = DDBClient(config.AWS_REGION)