"""Snowflake contract data query — bulk fetch via SnowflakeConnection.""" from __future__ import annotations import json from typing import Any from lambdacommon.common_config import logger import config from src.connectors.snowflake.connection import ( SnowflakeConnection, handle_snowflake_errors, ) from src.types import ContractData _SCHEMA_MAP = { config.PROD_ENVIRONMENT: 'prod', config.UAT_ENVIRONMENT: 'uat', } _VIEW_SCHEMA = _SCHEMA_MAP.get(config.ENVIRONMENT, 'qa') VIEW = f'royalty_accounting_reporting.{_VIEW_SCHEMA}.vw_abacus_balances_looker_v2' ROYALTY_ACCOUNTING_DATABASE = ( f'orchard_app_reporting_v2.{_VIEW_SCHEMA}_royalty_accounting_royalty_accounting' ) _SQL_TEMPLATE = """ WITH current_sp AS ( SELECT statement_period_id FROM {royalty_accounting_db}.statement_period WHERE statement_period_status = 'current' ORDER BY statement_period_id DESC LIMIT 1 ), ranked AS ( SELECT CONTRACT_ID, ACCOUNT_ID, ACCOUNT_NAME, CONTRACT_NAME, ACCOUNT_PAYEE_CURRENCY, EST_CLOSING_BAL_W_APPLIED_FT, GROSS_REVENUE, NET_REVENUE, ROW_NUMBER() OVER ( PARTITION BY CONTRACT_ID ORDER BY CASE WHEN STATEMENT_PERIOD_ID = (SELECT statement_period_id FROM current_sp) THEN 0 ELSE 1 END, STATEMENT_PERIOD_ID DESC ) AS rn FROM {view} WHERE CONTRACT_ID IN ({{placeholders}}) ) SELECT CONTRACT_ID, ACCOUNT_ID, ACCOUNT_NAME, CONTRACT_NAME, ACCOUNT_PAYEE_CURRENCY, MAX(CASE WHEN rn = 1 THEN EST_CLOSING_BAL_W_APPLIED_FT END) AS CLOSING_BALANCE, MAX(CASE WHEN rn = 1 THEN GROSS_REVENUE END) AS GROSS_REVENUE, MAX(CASE WHEN rn = 1 THEN NET_REVENUE END) AS NET_REVENUE, MAX(CASE WHEN rn = 2 THEN EST_CLOSING_BAL_W_APPLIED_FT END) AS PRIOR_CLOSING_BALANCE FROM ranked WHERE rn <= 2 GROUP BY CONTRACT_ID, ACCOUNT_ID, ACCOUNT_NAME, CONTRACT_NAME, ACCOUNT_PAYEE_CURRENCY """.strip().format(view=VIEW, royalty_accounting_db=ROYALTY_ACCOUNTING_DATABASE) def input_to_column(input_value: str | None) -> str: """Map a DB input value to the corresponding result column. 'closing_balance' -> CLOSING_BALANCE 'gross_revenue' -> GROSS_REVENUE 'net_revenue' -> NET_REVENUE """ if input_value == 'gross_revenue': return 'GROSS_REVENUE' if input_value == 'net_revenue': return 'NET_REVENUE' return 'CLOSING_BALANCE' def get_balance_for_input(data: ContractData, input_value: str | None) -> float | None: """Get the appropriate balance from ContractData for the given input.""" column = input_to_column(input_value) if column == 'GROSS_REVENUE': return data.gross_revenue if column == 'NET_REVENUE': return data.net_revenue return data.closing_balance def _safe_float(val: Any) -> float | None: """Convert a Snowflake result value to float or None.""" if val is None: return None return float(val) def _safe_int(val: Any) -> int | None: """Convert a Snowflake result value to int or None.""" if val is None: return None return int(val) BATCH_SIZE = 1000 @handle_snowflake_errors def fetch_contract_data( conn: SnowflakeConnection, contract_ids: list[int], ) -> dict[int, ContractData]: """Bulk-fetch contract metadata and balances from Snowflake. Uses a ranked CTE to get the current and prior statement period rows for each contract ID. IDs are batched in chunks of ``BATCH_SIZE`` to avoid Snowflake IN-clause limits. Args: conn: Active SnowflakeConnection. contract_ids: Contract IDs to fetch. Returns: Dict keyed by contract_id with ContractData values. """ logger.info( 'Fetching contract data from Snowflake: environment=%s view=%s contract_count=%s', config.ENVIRONMENT, VIEW, len(contract_ids), ) if not contract_ids: return {} # Validate all IDs are positive integers — bad IDs indicate upstream # data problems (malformed spreadsheet) that should be surfaced. invalid = [cid for cid in contract_ids if cid <= 0] if invalid: raise ValueError(f'Invalid contract IDs (must be positive): {invalid}') result: dict[int, ContractData] = {} for start in range(0, len(contract_ids), BATCH_SIZE): batch = contract_ids[start : start + BATCH_SIZE] placeholders = ', '.join(['%s'] * len(batch)) sql = _SQL_TEMPLATE.replace('{placeholders}', placeholders) cursor = conn.cursor() cursor.execute(sql, batch) col_names = [desc[0] for desc in cursor.description] rows = [dict(zip(col_names, row)) for row in cursor.fetchall()] for row in rows: contract_id = int(row['CONTRACT_ID']) row_serializable = {k: str(v) for k, v in row.items()} logger.info( 'Snowflake raw row: contract_id=%s raw=%s', contract_id, json.dumps(row_serializable), ) result[contract_id] = ContractData( account_id=_safe_int(row.get('ACCOUNT_ID')), account_name=str(row.get('ACCOUNT_NAME') or ''), contract_name=str(row.get('CONTRACT_NAME') or ''), currency=str(row.get('ACCOUNT_PAYEE_CURRENCY') or ''), closing_balance=_safe_float(row.get('CLOSING_BALANCE')), gross_revenue=_safe_float(row.get('GROSS_REVENUE')), net_revenue=_safe_float(row.get('NET_REVENUE')), prior_closing_balance=_safe_float(row.get('PRIOR_CLOSING_BALANCE')), ) for cid in contract_ids: if cid not in result: logger.info('No Snowflake data found for contract: contract_id=%s', cid) return result