"""Distribution Fee ETL specific utility functions.""" from collections import defaultdict import json import boto3 from flows import datastore from flows.distribution_fee import config from flows.distribution_fee import queries def send_sns_notification(correlation_id, report_date, upcs): """Send SNS message after ETL completes. Args: correlation_id (str): pre-chained correlation id. report_date (str): YYYY-MM-DD date of etl execution. upcs (list): upcs in the ETL. Returns: dict: with MessageId string. """ payload = { 'action': config.SNS_ACTION_SUCCESS, 'correlation_id': correlation_id, 'report_date': report_date, 'source': config.SNS_SOURCE, 'upcs': upcs} sns = boto3.client('sns', region_name=config.SNS_REGION_NAME) return sns.publish( TopicArn=config.SNS_TOPIC_ARN, Message=json.dumps(payload), Subject=payload['action']) def get_etl_vendor_contracts(table_name): """Get vendor contracts of the ETL from temporary table. Args: table_name (str): temporary vendor contracts table. Returns: list, dict: vendor contract IDs, vendor ID to UPC mapping. """ sql = queries.SELECT_TEMP_CONTRACTS.format(table_name=table_name) results = datastore.query(sql) contracts = results.fetchall() upcs = set() vendor_contract_ids = set() vendor_upc_map = defaultdict(set) for contract_id, vendor_id, upc in contracts: assert upc not in upcs, 'Duplicate UPCs for vendor contracts' upcs.add(upc) vendor_contract_ids.add(contract_id) vendor_upc_map[vendor_id].add(upc) return list(vendor_contract_ids), vendor_upc_map