"""DB Connector. Manages interactions with Snowflake.""" import sys from contextlib import contextmanager from functools import wraps from connectors import sentry from connectors.logger 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