"""Snowflake utils module.""" import os import config from cryptography.hazmat.backends import default_backend from cryptography.hazmat.primitives import serialization from snowflake import connector from snowflake.connector import connect as snowflake_connect def get_sf_private_key(): """Return Snowflake private_key. Returns: str: Snowflake private_key. """ private_key = None private_key_string = config.secrets_manager_client_lambda.get_cred('SNOWFLAKE_KEY') passphrase = config.secrets_manager_client_lambda.get_cred('SNOWFLAKE_KEY_PASSPHRASE') if private_key_string: private_key = serialization.load_pem_private_key( bytes(private_key_string, 'utf8'), password=bytes(passphrase, 'utf8'), backend=default_backend() ) if not private_key and config.ENVIRONMENT == 'dev': # Default for local dev snowflake_private_key_path = os.environ['HOME'] + \ '/.ssh/snowflake/rsa_key.p8' snowflake_private_key_path = os.environ.get( 'SNOWFLAKE_PRIVATE_KEY_PATH', snowflake_private_key_path) snowflake_key_passphrase = os.environ.get( 'SNOWFLAKE_KEY_PASSPHRASE', None) if snowflake_key_passphrase: with open(snowflake_private_key_path, 'rb') as key: p_key = serialization.load_pem_private_key( key.read(), password=snowflake_key_passphrase.encode(), backend=default_backend() ) private_key = p_key.private_bytes( encoding=serialization.Encoding.DER, format=serialization.PrivateFormat.PKCS8, encryption_algorithm=serialization.NoEncryption()) return private_key def execute_snowflake_query(query: str, params: dict) -> list[dict]: """Extract data from Snowflake. Args: query (str): SQL statement for extracting data from Snowflake. params (dict): For formatting the query. Returns: List: List of tuples or empty list. """ config.app_logger.info('Get snowflake connection') query = query.format( db=config.snowflake_db_config['db'], schema=config.snowflake_db_config['schema'] ) with snowflake_connect( **config.snowflake_db_config, private_key=get_sf_private_key(), ) as connection: with connection.cursor(connector.DictCursor) as cursor: return cursor.execute(query, params).fetchmany()