"""Connector for DynamoDB service.""" import copy import decimal from boto3.dynamodb import conditions from botocore import exceptions from accounting.flows.reserve_payouts import connectors from accounting.flows.reserve_payouts import setting from accounting.flows.reserve_payouts.constants import ( dynamodb as dynamodb_constants) from accounting.util import dynamodb def health_check(): """Perform a simple query to do the health check. Returns: namedtuple: with bool and message attributes, (True, '') if connection is ok, (False, 'Error message') otherwise. """ try: dynamodb.get_table_count(setting.DYNAMODB_TABLE) except (exceptions.BotoCoreError, exceptions.ClientError) as e: result = connectors.HealthCheckResult(False, str(e)) else: result = connectors.HealthCheckResult(True, '') return result def get_table_items_count(): """Get the items coun for reserve payout table. Returns: int: item count """ return dynamodb.get_table_count(setting.DYNAMODB_TABLE) def get_table_label_ids(): """Get label ids from table using scan. Returns: list: scan results, list of lable_id's """ table = dynamodb.get_dynamodb_table(setting.DYNAMODB_TABLE) resp = dynamodb.full_scan( table, Select='SPECIFIC_ATTRIBUTES', ProjectionExpression='label_id') return resp['Items'] def format_decimal_fields(vendor_data): """Cast types according to configuration. Args: vendor_data (dict): dict with data Returns: dict: a dict with correct types """ new_data = copy.deepcopy(vendor_data) for decimal_field in dynamodb_constants.DECIMAL_FIELDS: if decimal_field not in new_data: continue if type(new_data[decimal_field]) == decimal.Decimal: decimal_value = new_data[decimal_field] else: decimal_value = decimal.Decimal.from_float( new_data[decimal_field]) decimal_value = decimal_value.quantize( setting.DECIMAL_QUANTIZE, rounding=setting.DECIMAL_ROUNDING) new_data[decimal_field] = decimal_value return new_data def construct_vendor_item(key, vendor_data): """Construct a dict item corresponding to table schema. Args: key (int): vendor/label id vendor_data (dict): vendor data with transactions sum and contract terms Returns: dict: constructed item """ constructed_item = { 'label_id': key, 'processing_status': dynamodb_constants.DEFAULT_STATUS, 'ttl': dynamodb.get_ttl_value_for_item(setting.DYNAMODB_TTL) } new_data = format_decimal_fields(vendor_data) constructed_item.update(new_data) return constructed_item def batch_write_vendor_data(data): """Batch write items into DynamoDB table. Args: data (dict): dict with vendor id as a key """ table = dynamodb.get_dynamodb_table(setting.DYNAMODB_TABLE) with table.batch_writer() as batch_writer: for label_id, vendor_data in data.items(): batch_writer.put_item( Item=construct_vendor_item(label_id, vendor_data)) def update_vendor_status(vendor_id, processing_status): """Update a vendor item processing status in the table. Args: vendor_id (int): vendor/label id processing_status (str): new status to be set Returns: dict: DynamoDB response dict """ table = dynamodb.get_dynamodb_table(setting.DYNAMODB_TABLE) key = { 'label_id': vendor_id, } condition = conditions.Key('label_id').eq(vendor_id) expression = 'SET processing_status = :processing_status' expr_attributes = {':processing_status': processing_status} resp = table.update_item( Key=key, ConditionExpression=condition, ReturnValues='NONE', UpdateExpression=expression, ExpressionAttributeValues=expr_attributes, ) return resp def validate_successful_items_count(): """Validate the count of success items against other processing status. Returns: bool: True if total item count matches the SUCCESS item count, False otherwise """ table = dynamodb.get_dynamodb_table(setting.DYNAMODB_TABLE) condition = conditions.Attr( 'processing_status').eq(dynamodb_constants.SUCCESS) resp = dynamodb.full_scan( table, Select='COUNT', FilterExpression=condition) item_count = resp['Count'] scanned_count = resp['ScannedCount'] return item_count == scanned_count