"""pubsalesacc/connectors/snowflake.py — Snowflake connector (SQLAlchemy, private key auth). Connection and private key loading are deferred to the first actual query so that importing this module does not fail if SFPRIVATEKEYPATH is not yet configured. """ import logging import os from functools import lru_cache import pandas as pd from cryptography.hazmat.backends import default_backend from cryptography.hazmat.primitives import serialization from dotenv import load_dotenv from snowflake.sqlalchemy import URL from sqlalchemy import create_engine, text load_dotenv() logger = logging.getLogger(__name__) # --------------------------------------------------------------------------- # Lazy private key loading # --------------------------------------------------------------------------- _pkb: bytes | None = None def _load_private_key() -> bytes: """Load and cache the Snowflake RSA private key bytes (DER format). Reads the key file path from SFPRIVATEKEYPATH in .env. Called on first use, not at import time. """ global _pkb if _pkb is None: path = os.getenv("SFPRIVATEKEYPATH") if not path: raise EnvironmentError( "SFPRIVATEKEYPATH is not set in .env. " "Set it to the absolute path of your Snowflake RSA private key (.p8) file." ) if not os.path.exists(path): raise FileNotFoundError( f"Snowflake private key file not found at: {path}\n" "Update SFPRIVATEKEYPATH in your .env file." ) with open(path, "rb") as key_file: p_key = serialization.load_pem_private_key( key_file.read(), password=None, backend=default_backend(), ) _pkb = p_key.private_bytes( encoding=serialization.Encoding.DER, format=serialization.PrivateFormat.PKCS8, encryption_algorithm=serialization.NoEncryption(), ) logger.debug("Snowflake private key loaded.") return _pkb # --------------------------------------------------------------------------- # Engine factory # --------------------------------------------------------------------------- def get_engine( database: str = "", schema: str = "", role: str | None = None, warehouse: str | None = None, ): """Create and return a SQLAlchemy engine for Snowflake. Falls back to SFROLE / SFWAREHOUSE from .env if role/warehouse not passed. """ engine = create_engine( URL( account=os.getenv("SFACCOUNT"), user=os.getenv("SFUSER"), database=database, schema=schema, warehouse=warehouse or os.getenv("SFWAREHOUSE"), role=role or os.getenv("SFROLE"), ), connect_args={ "private_key": _load_private_key(), "autocommit": True, }, ) return engine # --------------------------------------------------------------------------- # Public query functions # --------------------------------------------------------------------------- def sf_df( sql: str, role: str | None = None, warehouse: str | None = None, ) -> pd.DataFrame: """Execute a SELECT query and return results as a DataFrame. Column names are uppercased to match Snowflake conventions. """ engine = get_engine(role=role, warehouse=warehouse) try: df = pd.read_sql(sql, engine) df.columns = [c.upper() for c in df.columns] return df except Exception as exc: logger.error("Snowflake query failed: %s", exc) raise finally: engine.dispose(close=True) def sf_execute( sql: str, role: str | None = None, warehouse: str | None = None, ) -> pd.DataFrame: """Execute a non-SELECT statement (UPDATE, INSERT, DELETE, etc.). Returns a one-row DataFrame with 'status' and 'rowcount' columns. """ engine = get_engine(role=role, warehouse=warehouse) try: with engine.connect() as conn: conn = conn.execution_options(isolation_level="AUTOCOMMIT") result = conn.execute(text(sql)) return pd.DataFrame([{"status": "success", "rowcount": result.rowcount}]) except Exception as exc: logger.error("Snowflake execute failed: %s", exc) return pd.DataFrame([{"status": "error", "message": str(exc)}]) finally: engine.dispose(close=True) def sf_write_df( df: pd.DataFrame, database: str, schema: str, table: str, role: str | None = None, warehouse: str | None = None, if_exists: str = "append", chunksize: int = 10_000, fix_col_names: bool = True, ) -> None: """Write a DataFrame to a Snowflake table. Args: if_exists: 'append' (default), 'replace', or 'fail'. fix_col_names: Uppercase and underscore-normalise column names. """ database, schema, table = database.casefold(), schema.casefold(), table.casefold() engine = get_engine(database, schema, role, warehouse) try: if fix_col_names: df = df.copy() df.columns = [c.upper().replace(" ", "_").strip() for c in df.columns] logger.info("Writing %d rows to %s.%s.%s...", len(df), database, schema, table) df.to_sql( table, engine, index=False, if_exists=if_exists, chunksize=chunksize, method="multi", ) logger.info("Write complete.") except Exception as exc: logger.error("sf_write_df failed: %s", exc) raise finally: engine.dispose(close=True)