"""DB Connector. Manages interactions with Snowflake.""" import sys from contextlib import contextmanager from functools import wraps # from integration_scripts.connectors import sentry from connectors.logging import logger import config from snowflake.sqlalchemy import URL from sqlalchemy import create_engine from sqlalchemy import pool from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import scoped_session from sqlalchemy.orm import sessionmaker # from .pidguard import add_engine_pidguard BaseModel = declarative_base() current_env = config.ENVIRONMENT # Database config if current_env == config.TEST_ENVIRONMENT: POOL_CLASS = pool.StaticPool DB_URL = 'sqlite://' else: # DB Pool starts with specified pool size, but will temporarily increase # if there is a surge of requests. POOL_CLASS = pool.QueuePool POOL_SIZE = 5 POOL_RECYCLE_MS = 3600 # Avoids connections going stale POOL_MAX_OVERFLOW = -1 POOL_PRE_PING = True POOL_RECYCLE = 3600 # Avoids connections going stale, recycle after 1 hour # Define small helper function for user error reporting def config_complain(field=None): """Print an error message about config files / environment and quit.""" logger.error( '\n\nThere is an error in the environment: ', sys.exc_info()[1] or '') logger.error( 'Please make sure your environment variables are configured properly.') logger.error(""" Required variables: SNOWFLAKE_ACCOUNT SNOWFLAKE_USER SNOWFLAKE_ROLE SNOWFLAKE_PASSWORD SNOWFLAKE_WAREHOUSE SNOWFLAKE_SCHEMA SNOWFLAKE_DATABASE and (...if dev) SNOWFLAKE_KEY_PASSPHRASE Optional: (None) """) if field: logger.error('{} has a bad value.'.format(field)) sys.exit() # Perform sanity checks against environment snow_db_config = {} if current_env != config.TEST_ENVIRONMENT: if not len(config.SNOWFLAKE_CONNECT_ARGS): logger.info('Snowflake will connect using password.') snow_db_config = { 'account': config.SNOWFLAKE_ACCOUNT, 'user': config.SNOWFLAKE_USER, 'role': config.SNOWFLAKE_ROLE, 'password': config.SNOWFLAKE_PASSWORD, 'warehouse': config.SNOWFLAKE_WAREHOUSE, 'schema': config.SNOWFLAKE_SCHEMA, 'database': config.SNOWFLAKE_DATABASE, # 'connect_args': config.SNOWFLAKE_CONNECT_ARGS, } else: logger.info('Snowflake will connect using PSK.') snow_db_config = { 'account': config.SNOWFLAKE_ACCOUNT, 'user': config.SNOWFLAKE_USER, 'role': config.SNOWFLAKE_ROLE, # 'password': config.SNOWFLAKE_PASSWORD, 'warehouse': config.SNOWFLAKE_WAREHOUSE, 'schema': config.SNOWFLAKE_SCHEMA, 'database': config.SNOWFLAKE_DATABASE, 'connect_args': config.SNOWFLAKE_CONNECT_ARGS, } for key, value in snow_db_config.items(): if value == '' or value is None: config_complain(key) else: snow_db_config = { 'account': 'test_acc', 'role': 'test_role', 'host': 'test_host', 'warehouse': 'test_wh', 'port': 10, 'user': 'test_user', 'password': 'test_pass', 'database': 'test_db', 'schema': 'test_schema', 'connect_args': {}, } def _create_engine(snow_db_cfg): """Create engine based on configuration settings. Returns: obj: db_engine """ if not len(config.SNOWFLAKE_CONNECT_ARGS): db_url = URL( account=snow_db_cfg['account'], user=snow_db_cfg['user'], password=snow_db_cfg['password'], warehouse=snow_db_cfg['warehouse'], timezone='America/New_York', schema=snow_db_cfg['schema'], database=snow_db_cfg['database']) connect_args = {} else: db_url = URL( account=snow_db_cfg['account'], user=snow_db_cfg['user'], password='', warehouse=snow_db_cfg['warehouse'], timezone='America/New_York', schema=snow_db_cfg['schema'], database=snow_db_cfg['database']) connect_args = snow_db_cfg['connect_args'] engine = create_engine( db_url, poolclass=pool.NullPool, pool_pre_ping=True, connect_args=connect_args) # add_engine_pidguard(engine) return engine # Do not use these variables directly other than running unit tests _db_engine = _create_engine(snow_db_config) _db_session = scoped_session(sessionmaker(bind=_db_engine)) @contextmanager def db_session(): """Provide a transactional scope around a series of operations. Taken from http://docs.sqlalchemy.org/en/latest/orm/session_basics.html. This handles rollback and closing of session, so there is no need to do that throughout the code. Usage: with db_session() as session: session.execute(query) """ session = _db_session() try: yield session session.commit() except Exception as e: logger.error( '------------------------------------------------------------\n' '------------------------- DB ERROR -----------------------\n' '------------------------------------------------------------') logger.error(str(e)) session.rollback() raise finally: session.close() # @sentry.sentry_wrap def db_session_wrap(func): """DB Session Wrapper. Creates a new session if one isn't passed in. """ @wraps(func) def wrapper(*args, **kwargs): session = kwargs.pop('session', None) if session: result = func(*args, session=session, **kwargs) else: with db_session() as session: result = func(*args, session=session, **kwargs) return result return wrapper