import base64 import logging import os from pathlib import Path from cryptography.hazmat.backends import default_backend from cryptography.hazmat.primitives import serialization from snowflake.connector import connect from snowflake.connector.pandas_tools import write_pandas from tadas.platform.caching import cached_df logger = logging.getLogger(__name__) THIS_DIR = Path(__file__).parent script_dir = os.path.dirname(os.path.abspath(__file__)) def decode_snowflake_key(snowflake_key): """Decode the Snowflake key.""" if not snowflake_key: return decoded_key = base64.b64decode(bytes(snowflake_key, encoding='utf-8')) return decoded_key def get_snowflake_private_key(): """Get the private key for Snowflake.""" private_key_path = os.getenv('SNOWFLAKE_PRIVATE_KEY_PATH') private_key: str = os.getenv('SNOWFLAKE_KEY') if private_key: key_data = private_key.encode() elif private_key_path: private_key_abs_path = os.path.join(script_dir, private_key_path) with open(private_key_abs_path, 'r') as f: key_data = f.read().encode() else: raise ValueError("Nor SNOWFLAKE_KEY nor SNOWFLAKE_PRIVATE_KEY_PATH environment variables are set.") passphrase = os.getenv('SNOWFLAKE_KEY_PASSPHRASE') passphrase_bytes = passphrase.encode() if passphrase else None private_key = serialization.load_pem_private_key( key_data, password=passphrase_bytes, backend=default_backend(), ) return private_key def snowflake_connection(): connection = connect( user=os.getenv('SNOWFLAKE_USER'), account=os.getenv('SNOWFLAKE_ACCOUNT'), warehouse=os.getenv('SNOWFLAKE_WAREHOUSE'), role=os.getenv('SNOWFLAKE_ROLE'), database=os.getenv('SNOWFLAKE_DATABASE'), schema=os.getenv('SNOWFLAKE_SCHEMA'), private_key=get_snowflake_private_key(), ) return connection def saveto_snowflake(df, myschema, table, mode, fix_column_names=True): """ Save a dataframe to Snowflake. Column names get uppercased so they are case-insensitive in queries. Pass fix_column_names=False to keep the original case. """ connection = snowflake_connection() database = os.getenv('SNOWFLAKE_DATABASE') logger.info(f'Writing {df.shape[0]} rows to Snowflake table {table} in {database}.{myschema}') original_columns = df.columns if fix_column_names: df.columns = df.columns.str.upper() result = write_pandas( conn=connection, df=df, table_name=table.upper(), auto_create_table=True, database=database.upper(), schema=myschema.upper(), overwrite=mode == 'replace', table_type='transient', chunk_size=500_000, ) logger.info(f'result: {result}') df.columns = original_columns def query_snowflake_fetchall(connection, query, params=None): """Query Snowflake and return a list of tuples.""" logger.info(f'Executing in Snowflake: {query}') logger.info(f'{params=}') with connection.cursor() as cursor: return cursor.execute(query, params).fetchall() @cached_df(hashing_kwargs={'query', 'params'}) def query_snowflake_to_df(connection, query, params=None): """Query Snowflake and return a DataFrame.""" logger.info(f'Executing in Snowflake: {query}') logger.info(f'{params=}') with connection.cursor() as cursor: cursor.execute(query, params) df = cursor.fetch_pandas_all() df = post_process_df(df) return df def execute(connection, query, params=None): """Execute a query against Snowflake (no result).""" logger.info(f'Executing in Snowflake: {query}') logger.info(f'Params: {params}') with connection.cursor() as cursor: cursor.execute(query, params) def post_process_df(df): df.columns = [col.lower() for col in df.columns] return df def is_table_exists(table_name, conn): query = 'SHOW TABLES LIKE %(table_name)s' results = query_snowflake_fetchall( connection=conn, query=query, params={'table_name': table_name}, ) return bool(results) def clone_table(source_table, target_table, conn): return execute(conn, f"create or replace transient table {target_table} clone {source_table}") def drop_table(table_name, conn): execute(conn, f"drop table if exists {table_name}")