"""Reference data loader service. Loads reference tables from Snowflake into DuckDB with intelligent loading strategy selection (bulk vs. batched) and parallel execution via DAG-based task graph. """ from __future__ import annotations from abacus_common_logic.concurrent import Task, TaskGraph, TaskGraphRunner from abacus_common_logic.utils.profiling import profile from src.connectors.duckdb import DuckDBConnection from src.enums import DuckDBTable from src.gateways import SnowflakeGateway from src.infra.log import logger from src.infra.resources import ResourceManager from src.sql import DuckDBQuery, load_sql from src.utils.db_utils import ( count_table_rows, drop_table, ) class ReferenceLoader: """Orchestrates the loading of all reference data for the application.""" def __init__(self, duck_conn: DuckDBConnection, gateway: SnowflakeGateway) -> None: """Initialize the service.""" self._duck_conn = duck_conn self._gateway = gateway @profile(logger=logger) def load_all(self) -> None: """Load all reference tables in parallel using DAG-based task execution.""" logger.info('Loading reference data') graph = self._build_task_graph() runner = TaskGraphRunner( graph=graph, pools={ 'cpu': ResourceManager.get_cpu_pool(), 'io': ResourceManager.get_io_pool(), }, default_pool='cpu', ) runner.run() logger.info('Reference data loaded') def _build_task_graph(self) -> TaskGraph: """Build task dependency graph for parallel loading with derived table transformations.""" tasks: list[Task] = [ Task( id='create_upc_lookup', handler=self._create_upc_lookup, depends_on={'load_account_upcs'}, pool='cpu', ), Task( id='drop_account_upcs', handler=self._drop_account_upcs, depends_on={'create_upc_lookup'}, pool='cpu', ), Task( id='load_accounts', handler=self._gateway.import_accounts, pool='io', ), Task( id='load_account_upcs', handler=self._gateway.import_account_upcs, pool='io', ), Task( id='load_account_contracts', handler=self._gateway.import_account_contracts, pool='io', ), Task( id='load_account_payment_terms', handler=self._gateway.import_account_payment_terms, pool='io', ), Task( id='load_close_balance_statuses', handler=self._gateway.import_close_balance_statuses, pool='io', ), Task( id='load_currency_codes', handler=self._gateway.import_currency_codes, pool='io', ), Task( id='load_flat_contract_terms', handler=self._gateway.import_flat_contract_terms, pool='io', ), Task( id='load_reference_adjustment_types', handler=self._gateway.import_reference_adjustment_types, pool='io', ), Task( id='load_statement_periods', handler=self._gateway.import_statement_periods, pool='io', ), ] return TaskGraph(tasks) def _create_upc_lookup(self) -> None: """Create UPC lookup table for efficient UPC validation.""" with self._duck_conn.cursor() as cursor: table_name = DuckDBTable.UPC_LOOKUP logger.info(f'Creating {table_name}') cursor.execute(load_sql(DuckDBQuery.CreateUpcLookup)) row_count = count_table_rows(cursor, table_name) logger.info(f'- {table_name}: {row_count}') def _drop_account_upcs(self) -> None: """Drop account UPCs table.""" with self._duck_conn.cursor() as cursor: table_name = DuckDBTable.ACCOUNT_UPCS logger.info(f'Dropping {table_name}') drop_table(cursor, table_name) logger.info(f'Dropped {table_name}')