"""Snowflake database connector.""" from functools import wraps from cryptography.hazmat.backends import default_backend from cryptography.hazmat.primitives import serialization from sqlalchemy.orm import declarative_base, sessionmaker from collaborator import config from collaborator.utils import db as db_utils # https://docs.snowflake.com/user-guide/key-pair-auth # https://docs.snowflake.com/developer-guide/python-connector/python-connector-example#using-key-pair-authentication-key-pair-rotation # https://docs.snowflake.com/developer-guide/python-connector/sqlalchemy#key-pair-authentication-support if config.ENVIRONMENT != config.TEST_ENVIRONMENT: passphrase_bytes = bytes(config.SNOWFLAKE_KEY_PASSPHRASE, "utf8") if config.SNOWFLAKE_PRIVATE_KEY: private_key_bytes = bytes(config.SNOWFLAKE_PRIVATE_KEY, "utf8") elif config.SNOWFLAKE_PRIVATE_KEY_PATH: with open(config.SNOWFLAKE_PRIVATE_KEY_PATH, "rb") as key: private_key_bytes = key.read() else: raise Exception( "SNOWFLAKE_PRIVATE_KEY_PATH or SNOWFLAKE_PRIVATE_KEY must be set." ) decrypted_key = serialization.load_pem_private_key( private_key_bytes, password=(passphrase_bytes or None), backend=default_backend(), ) decrypted_key_bytes = decrypted_key.private_bytes( encoding=serialization.Encoding.DER, format=serialization.PrivateFormat.PKCS8, encryption_algorithm=serialization.NoEncryption(), ) connect_args = {"private_key": decrypted_key_bytes} else: connect_args = {} _db_engine = db_utils.create_engine( config.SNOWFLAKE_DB_URL, connect_args=connect_args, ) _db_writer_engine = db_utils.create_engine( config.SNOWFLAKE_DB_WRITER_URL, connect_args=connect_args, ) _session_maker = sessionmaker(bind=_db_engine) _writer_session_maker = sessionmaker(bind=_db_writer_engine) BaseModel = declarative_base() def db_session(function): """DB Session Wrapper. Creates a new session if one isn't passed in. """ @wraps(function) def wrapper(*args, **kwargs): session = kwargs.pop("session", None) if session: return function(*args, session=session, **kwargs) else: with db_utils.create_session(_session_maker) as session: return function(*args, session=session, **kwargs) return wrapper def db_writer_session(function): """DB Writer Session Wrapper. Creates a new session if one isn't passed in. Targets a warehouse tailored for writes to Snowflake. """ @wraps(function) def wrapper(*args, **kwargs): session = kwargs.pop("session", None) if session: return function(*args, session=session, **kwargs) else: with db_utils.create_session(_writer_session_maker) as session: return function(*args, session=session, **kwargs) return wrapper