"""High-level logic for fetching entities and preparing for sync.""" import base64 import json from typing import List from typing import Type from uuid import uuid4 from aws_kinesis_agg.aggregator import RecordAggregator from src.common.models.metadata import ARSyncEvent from src.common.models.metadata import ArtRelationsTable from src.common.models.metadata import InconsistencyReport from src.common.models.metadata import InconsistencyType def filter_report_type( report_list: List[InconsistencyReport], report_type: InconsistencyType ) -> InconsistencyReport: """Get InconsistencyReport list filtered to only a single report_type. Args: report_list: List of InconsistencyReports report_type: InconsistencyType to retrieve Returns: InconsistencyReport """ filtered = [report for report in report_list if report.inconsistency_type == report_type.value] if not filtered: raise ValueError( f'Provided inconsistence type {report_type.value} not found' ) return filtered[0] def get_ids_per_table( report: InconsistencyReport, table: Type[ArtRelationsTable] = None ) -> List[int]: """Get ids from InconsistencyReport. Args: report: InconsistencyReport table: (Optional) filter by specific table in InconsistencyReport Returns: List of ids from InconsistencyReport """ if table: return [ record.record_id for record in report.missing_data if record.table_name == table.name ] return [record.record_id for record in report.missing_data] def encode_records(records: List[ARSyncEvent]) -> str: """Encode records for sync payload. Args: records: List of AR records for sync. Returns: List of encoded records for sync payload. """ aggregator = RecordAggregator() for record in records: aggregator.add_user_record( str(uuid4()), json.dumps(record.dict(), default=str) ) agg_content = aggregator.current_record.get_contents()[-1] encoded_records = base64.b64encode(agg_content).decode() return encoded_records