import json from functools import cache from os import environ import boto3 import snowflake.connector as sc from cryptography.hazmat.backends import default_backend from cryptography.hazmat.primitives.serialization import ( Encoding, NoEncryption, PrivateFormat, load_pem_private_key, ) _AWS_SECRET_NAME = "prod/hive-regression-tracking/HIVE_MODEL_REPORTING_SERVICE_USER_KEY" _SNOWFLAKE_ACCOUNT = "delphi.us-east-1" _SNOWFLAKE_USER_NAME = "PROD_HIVE_MODEL_REPORTING_SERVICE_USER" @cache def _load_private_key_der(): client = boto3.client("secretsmanager", region_name=environ["AWS_REGION"]) secret = client.get_secret_value(SecretId=_AWS_SECRET_NAME) data = json.loads(secret["SecretString"]) private_key = load_pem_private_key( data["private_key"].encode(), password=data["passphrase"].encode(), backend=default_backend(), ) return private_key.private_bytes( encoding=Encoding.DER, format=PrivateFormat.PKCS8, encryption_algorithm=NoEncryption(), ) @cache def snowflake_connection(): ctx = sc.connect( account=_SNOWFLAKE_ACCOUNT, user=_SNOWFLAKE_USER_NAME, authenticator="SNOWFLAKE_JWT", private_key=_load_private_key_der(), warehouse=environ.get("SNOWFLAKE_WAREHOUSE") or None, database=environ.get("SNOWFLAKE_DATABASE") or None, schema=environ.get("SNOWFLAKE_SCHEMA") or None, ) return ctx def execute(sql_query): """Execute a query against Snowflake.""" print("QUERY", sql_query) ctx = snowflake_connection() cursor = ctx.cursor() return cursor.execute(sql_query)