"""Config file for lambda function.""" from __future__ import annotations import base64 import logging import os from dataclasses import dataclass, field from pathlib import Path from typing import Final from abacus_common_logic.constants.math import Bytes from dotenv import load_dotenv from secrets_manager.lambda_ext import LambdaSecretsManager from src.connectors.mysql import MySQLConfig from src.connectors.snowflake import SnowflakeConfig from src.enums import Environment # Load environment variables from a .env file if present load_dotenv(override=True) @dataclass(frozen=True) class DuckDBSettings: """Environment configuration for DuckDB.""" # Path where extensions are stored (baked into lambda) HOME_DIR: Final[Path] = Path( os.environ.get('DUCKDB_HOME_DIR', str(Path.cwd() / 'duckdb')) ) MAX_THREADS: int | None = ( int(os.environ['DUCKDB_MAX_THREADS']) if os.environ.get('DUCKDB_MAX_THREADS') else None ) # The % of total system RAM DuckDB is allowed to use MEM_PCT: Final[float] = float(os.environ.get('DUCKDB_MEM_PCT', '0.5')) # Path for temp files when RAM limits are exceeded TEMP_DIR: Final[Path] = Path(os.environ.get('DUCKDB_TEMP_DIR', '/tmp/.duckdb')) # The % of available /tmp disk space DuckDB can use TEMP_PCT: Final[float] = float(os.environ.get('DUCKDB_TEMP_PCT', '0.6')) @dataclass(frozen=True) class EncodingConfig: """File encoding detection and processing settings. Defines default encodings, chunk sizes for reading files, and sample sizes for encoding detection and checksum calculation. """ DEFAULT: str = os.environ.get('ENCODING_DEFAULT', 'utf-8') CSV_SAMPLE_ROWS: int = int(os.environ.get('ENCODING_CSV_SAMPLE_ROWS', '1_000')) MD5_CHUNK_BYTES: int = int( os.environ.get('ENCODING_MD5_CHUNK_BYTES', str(64 * Bytes.KB)) ) SAMPLE_BYTES: int = int( os.environ.get('ENCODING_SAMPLE_BYTES', str(512 * Bytes.KB)) ) SAMPLE_CHUNK_BYTES: int = int( os.environ.get('ENCODING_SAMPLE_CHUNK_BYTES', str(64 * Bytes.KB)) ) @dataclass(frozen=True) class ExecutionConfig: """Resource limits for processing.""" MAX_QUERY_FILTER_RATIO: float = float( os.environ.get('EXEC_MAX_QUERY_FILTER_RATIO', '0.8') ) # The point where filtered batching becomes more expensive than a full scan. # Based on Snowflake compilation overhead and DuckDB ingest performance. # The value is a percentage of local IDs vs total remote size. Should ideally # be between 15% - 25%. # # For example, a value of 0.2 (i.e. 20%) means: # - < 20%: Load what you need (ID-based batching) # Reduces memory, Snowflake data transfer costs and local ingest time # - >= 20%: Full table load # Reduces Snowflake I/O and network latencies. MIN_BULK_RATIO: float = float(os.environ.get('EXEC_MIN_BULK_RATIO', '0.2')) # The maximum row count for which a table is considered "small." # For tables under this limit, the overhead of calculating required IDs # is often slower than just downloading everything. # # For this particular lambda, rows are narrow; mostly just the ID, with # maybe < 5 additional fields. A value of 250K would consume # negligible memory (<50 MB). MAX_BULK_ROWS: int = int(os.environ.get('EXEC_MAX_BULK_ROWS', '250_000')) # Scaling rule for IO-bound work IO_THREADS_PER_VCPU: int = int(os.environ.get('EXEC_IO_THREADS_PER_VCPU', '4')) DEFAULT_CPU_COUNT: int = int(os.environ.get('EXEC_DEFAULT_CPU_COUNT', '1')) MAX_CPU_COUNT: int | None = ( int(os.environ['EXEC_MAX_CPU_COUNT']) if os.environ.get('EXEC_MAX_CPU_COUNT') else None ) @dataclass(frozen=True) class StorageSettings: """Settings for S3 staging and local storage.""" S3_STAGING_FORMAT: str = os.environ.get( 'STORAGE_S3_FORMAT', 'staging/{batch_id}/prepared.csv' ) TEMP_DIR: Path = Path(os.environ.get('STORAGE_TEMP_DIR', '/tmp')) @dataclass(frozen=True) class ValidationPolicy: """Business rules that might scale.""" MAX_COMMENT_LEN: int = int(os.environ.get('VAL_MAX_COMMENT_LENGTH', '180')) MAX_FILE_ROWS: int = int(os.environ.get('VAL_MAX_FILE_ROWS', '1_000_000')) MAX_FILE_BYTES: int = int(os.environ.get('VAL_MAX_FILE_BYTES', str(Bytes.GB))) MAX_VALUE_LEN: int = int(os.environ.get('VAL_MAX_VALUE_LENGTH', '250')) @dataclass class AppConfig: """App configuration.""" app_name: Final[str] = 'lambda-abacus-adjustment-file-prepare' aws_region: Final[str] = os.environ.get('AWS_REGION', 'us-east-1') env: Environment = Environment(os.environ.get('ENVIRONMENT', Environment.DEV)) log_level: int = field(default_factory=lambda: 20) sentry_dsn: str | None = os.environ.get('SENTRY_DSN') duckdb: DuckDBSettings = field(default_factory=DuckDBSettings) encoding: EncodingConfig = field(default_factory=EncodingConfig) exec: ExecutionConfig = field(default_factory=ExecutionConfig) policy: ValidationPolicy = field(default_factory=ValidationPolicy) storage: StorageSettings = field(default_factory=StorageSettings) mysql: MySQLConfig = field(init=False) snowflake: SnowflakeConfig = field(init=False) def __post_init__(self) -> None: """Initialize connectors and override with secrets for non-dev environments.""" self.log_level = self._get_log_level() self.mysql = self._init_mysql() self.snowflake = self._init_snowflake() if self.env.is_managed: self._apply_secrets() def _get_log_level(self) -> int: """Safely convert LOGGING_LEVEL env var to a numeric logging constant.""" raw_level = os.environ.get('LOGGING_LEVEL', 'INFO').upper() if raw_level.isdigit(): return int(raw_level) return getattr(logging, raw_level, logging.INFO) def _init_mysql(self) -> MySQLConfig: return MySQLConfig( host=os.environ.get('MYSQL_DB_HOST', ''), user=os.environ.get('MYSQL_DB_USER', ''), password=os.environ.get('MYSQL_DB_PASS', ''), database=os.environ.get('MYSQL_DB_NAME', 'royalty_accounting'), port=int(os.environ.get('MYSQL_DB_PORT', '3306')), local_infile=self.env.is_local, ) def _init_snowflake(self) -> SnowflakeConfig: private_key = os.environ.get('SNOWFLAKE_PRIVATE_KEY') # Decode base64-encoded private key if provided (e.g., local development) if private_key and private_key.startswith('LS0tLS'): private_key = base64.b64decode(private_key).decode('utf-8') private_key = private_key.replace('\\n', '\n') return SnowflakeConfig( host=os.environ.get( 'SNOWFLAKE_HOST', 'delphi.us-east-1.snowflakecomputing.com' ), account=os.environ.get('SNOWFLAKE_ACCOUNT', 'delphi.us-east-1'), user=os.environ.get('SNOWFLAKE_USER', ''), private_key=private_key, private_key_passphrase=os.environ.get('SNOWFLAKE_KEY_PASS'), warehouse=os.environ.get( 'SNOWFLAKE_WAREHOUSE', 'qa_accounting_dbt_warehouse' ), database=os.environ.get('SNOWFLAKE_DATABASE', 'orchard_app_reporting_v2'), schema_name=os.environ.get( 'SNOWFLAKE_SCHEMA', 'qa_royalty_accounting_royalty_accounting' ), role=os.environ.get('SNOWFLAKE_ROLE', 'dev_engineering'), ) def _apply_secrets(self) -> None: """Fetch credentials from AWS Secrets Manager for higher environments.""" client = LambdaSecretsManager( environment=self.env.value, service_name=self.app_name ) self.sentry_dsn = client.get_cred('SENTRY_DSN') # Update MySQL password self.mysql = MySQLConfig( **self.mysql.model_dump(exclude={'password'}), password=client.get_cred('MYSQL_DB_PASS'), ) # Update Snowflake private key/passphrase self.snowflake = SnowflakeConfig( **self.snowflake.model_dump( exclude={'private_key', 'private_key_passphrase'} ), private_key=client.get_cred('SNOWFLAKE_PRIVATE_KEY'), private_key_passphrase=client.get_cred('SNOWFLAKE_KEY_PASS'), ) def validate_config(cfg: AppConfig) -> None: """Validate configuration values on startup to fail fast. Checks resource percentages, thread counts, and required string formats to ensure the environment is fit for processing. Raises: ValueError: If any configuration value is logically invalid. """ errors = [] # DuckDB if cfg.env.is_managed and not cfg.duckdb.HOME_DIR.exists(): errors.append(f'duckdb.HOME_DIR does not exist: {cfg.duckdb.HOME_DIR}') if not (0 < cfg.duckdb.MEM_PCT <= 0.8): errors.append(f'duckdb.MEM_PCT must be (0, 0.8], got {cfg.duckdb.MEM_PCT}') if not (0 < cfg.duckdb.TEMP_PCT <= 1): errors.append(f'duckdb.TEMP_PCT must be (0, 1], got {cfg.duckdb.TEMP_PCT}') if cfg.duckdb.MAX_THREADS and not (1 <= cfg.duckdb.MAX_THREADS <= 16): errors.append( f'duckdb.MAX_THREADS must be [1, 16], got {cfg.duckdb.MAX_THREADS}' ) # Execution if not (1 <= cfg.exec.DEFAULT_CPU_COUNT <= 4): errors.append( f'exec.DEFAULT_CPU_COUNT must be [1, 4], got {cfg.exec.DEFAULT_CPU_COUNT}' ) if not (1 <= cfg.exec.IO_THREADS_PER_VCPU <= 8): errors.append( f'exec.IO_THREADS_PER_VCPU must be [1, 8], got {cfg.exec.IO_THREADS_PER_VCPU}' ) if cfg.exec.MAX_CPU_COUNT and not (1 <= cfg.exec.MAX_CPU_COUNT <= 16): errors.append( f'exec.MAX_CPU_COUNT must be [1, 16], got {cfg.exec.MAX_CPU_COUNT}' ) if not (0 < cfg.exec.MAX_QUERY_FILTER_RATIO <= 0.95): errors.append( f'exec.MAX_QUERY_FILTER_RATIO must be (0, 0.95], got {cfg.exec.MAX_QUERY_FILTER_RATIO}' ) if cfg.exec.MAX_BULK_ROWS < 1: errors.append('MAX_BULK_ROWS must be positive.') if not (0 <= cfg.exec.MIN_BULK_RATIO <= 1): errors.append(f'MIN_BULK_RATIO must be [0, 1], got {cfg.exec.MIN_BULK_RATIO}') effective_cpu_limit = cfg.exec.MAX_CPU_COUNT or cfg.exec.DEFAULT_CPU_COUNT potential_io_threads = effective_cpu_limit * cfg.exec.IO_THREADS_PER_VCPU if potential_io_threads > 50: errors.append( f'Total threads ({potential_io_threads}) exceeds safety limit of 50. ' f'Reduce exec.MAX_CPU_COUNT or exec.IO_THREADS_PER_VCPU' ) # Encoding if not cfg.encoding.DEFAULT: errors.append('encoding.DEFAULT cannot be empty') min_md5_chunk_bytes = 4 * Bytes.KB max_md5_chunk_bytes = 128 * Bytes.MB if not (min_md5_chunk_bytes <= cfg.encoding.MD5_CHUNK_BYTES <= max_md5_chunk_bytes): errors.append( f'encoding.MD5_CHUNK_BYTES must be [{min_md5_chunk_bytes}, ' f'{max_md5_chunk_bytes}], got {cfg.encoding.MD5_CHUNK_BYTES}' ) if cfg.encoding.CSV_SAMPLE_ROWS < 100: errors.append( f'encoding.CSV_SAMPLE_ROWS must be >= 100, ' f'got {cfg.encoding.CSV_SAMPLE_ROWS}' ) min_sample_chunk_bytes = 32 * Bytes.KB max_sample_chunk_bytes = 10 * Bytes.MB if not ( min_sample_chunk_bytes <= cfg.encoding.SAMPLE_CHUNK_BYTES <= max_sample_chunk_bytes ): errors.append( f'encoding.SAMPLE_CHUNK_BYTES must be [{min_sample_chunk_bytes}, ' f'{max_sample_chunk_bytes}], got {cfg.encoding.SAMPLE_CHUNK_BYTES}' ) min_sample_bytes = 32 * Bytes.KB max_sample_bytes = 10 * Bytes.MB if not (min_sample_bytes <= cfg.encoding.SAMPLE_BYTES <= max_sample_bytes): errors.append( f'encoding.SAMPLE_BYTES must be [{min_sample_bytes}, ' f'{max_sample_bytes}], got {cfg.encoding.SAMPLE_BYTES}' ) if cfg.encoding.SAMPLE_BYTES < cfg.encoding.SAMPLE_CHUNK_BYTES: errors.append( f'encoding.SAMPLE_BYTES ({cfg.encoding.SAMPLE_BYTES}) should be >= ' f'encoding.SAMPLE_CHUNK_BYTES ({cfg.encoding.SAMPLE_CHUNK_BYTES})' ) # Storage if cfg.storage.S3_STAGING_FORMAT.startswith('/'): errors.append('storage.S3_STAGING_FORMAT should not start with a slash') if '{batch_id}' not in cfg.storage.S3_STAGING_FORMAT: errors.append("storage.S3_STAGING_FORMAT must contain '{batch_id}' placeholder") # Validation Policy max_file_rows = 2_000_000 if not (1 <= cfg.policy.MAX_FILE_ROWS <= max_file_rows): errors.append( f'policy.MAX_FILE_ROWS must be [1, {max_file_rows}], got {cfg.policy.MAX_FILE_ROWS}' ) max_file_bytes = 2.5 * Bytes.GB if not (0 <= cfg.policy.MAX_FILE_BYTES <= max_file_bytes): errors.append( f'policy.MAX_FILE_BYTES must be [0, {max_file_bytes}], got {cfg.policy.MAX_FILE_BYTES}' ) if not (1 <= cfg.policy.MAX_COMMENT_LEN <= 1000): errors.append( f'policy.MAX_COMMENT_LEN must be [1, 1000], got {cfg.policy.MAX_COMMENT_LEN}' ) if not (1 <= cfg.policy.MAX_VALUE_LEN <= 1000): errors.append( f'policy.MAX_VALUE_LEN must be [1, 1000], got {cfg.policy.MAX_VALUE_LEN}' ) if cfg.policy.MAX_VALUE_LEN < cfg.policy.MAX_COMMENT_LEN: errors.append( f'policy.MAX_COMMENT_LEN ({cfg.policy.MAX_COMMENT_LEN}) must be <= ' f'policy.MAX_VALUE_LEN ({cfg.policy.MAX_VALUE_LEN})' ) # Lambda settings if cfg.env.is_managed: if not str(cfg.duckdb.TEMP_DIR).startswith('/tmp'): errors.append( f'duckdb.TEMP_DIR must be in /tmp for Lambda, got {cfg.duckdb.TEMP_DIR}' ) if not str(cfg.storage.TEMP_DIR).startswith('/tmp'): errors.append( f'storage.TEMP_DIR must be in /tmp for Lambda, got {cfg.storage.TEMP_DIR}' ) if not cfg.mysql.password: errors.append('mysql.password is empty after applying secrets') if not cfg.snowflake.private_key: errors.append('snowflake.private_key is empty after applying secrets') if errors: raise ValueError( 'Configuration validation failed:\n' + '\n'.join(f' - {e}' for e in errors) ) config = AppConfig() validate_config(config)