"""DB Connector. Manages interactions with Snowflake.""" from contextlib import contextmanager from functools import wraps import os from snowflake.sqlalchemy import URL from sqlalchemy import create_engine from sqlalchemy import event from sqlalchemy import exc from sqlalchemy import pool from sqlalchemy import select from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import sessionmaker from social_analytics import config BaseModel = declarative_base() role = os.environ.get('SNOWFLAKE_ROLE') database = os.environ.get('SNOWFLAKE_DATABASE') schema = os.environ.get('SNOWFLAKE_SCHEMA') warehouse = os.environ.get('SNOWFLAKE_WAREHOUSE') def _create_engine(db_url): """Create engine based on configuration settings.""" if config.POOL_CLASS == pool.QueuePool: db_engine = create_engine( db_url, pool_size=config.POOL_SIZE, max_overflow=config.POOL_MAX_OVERFLOW, pool_recycle=config.POOL_RECYCLE) @event.listens_for(db_engine, 'connect') def set_snowflake_params(connection, connection_record): """Set session params. When using raw SQL with fully qualified table names (db.schema.table), USE DATABASE and USE SCHEMA are not required. """ with connection.cursor() as cur: if role: cur.execute('USE ROLE {};'.format(role)) if database: cur.execute('USE DATABASE {};'.format(database)) if schema: cur.execute('USE SCHEMA {};'.format(schema)) if warehouse: cur.execute('USE WAREHOUSE {};'.format(warehouse)) if config.POOL_PRE_PING: @event.listens_for(db_engine, 'engine_connect') def ping_connection(connection, branch): _ping_connection(connection, branch) return db_engine else: return create_engine(db_url, poolclass=config.POOL_CLASS) def _ping_connection(connection, branch): """Ping database connection after engine_connect event. This function is copied verbatim from http://docs.sqlalchemy.org/en/latest/core/pooling.html """ if branch: return save_should_close_with_result = connection.should_close_with_result connection.should_close_with_result = False try: connection.scalar(select([1])) except exc.DBAPIError as err: if err.connection_invalidated: connection.scalar(select([1])) else: raise finally: connection.should_close_with_result = save_should_close_with_result if config.ENVIRONMENT == config.TEST_ENVIRONMENT: snowflake_db_engine = create_engine('sqlite://') else: SNOWFLAKE_DB_URL = URL( account=config.SNOWFLAKE_ACCOUNT, user=config.SNOWFLAKE_USER, password=config.SNOWFLAKE_PASSWORD, database=config.SNOWFLAKE_DATABASE, schema=config.SNOWFLAKE_SCHEMA, warehouse=config.SNOWFLAKE_WAREHOUSE ) snowflake_db_engine = _create_engine(SNOWFLAKE_DB_URL) # please don't use sessions directly; instead use db_session sessions = { config.SNOWFLAKE_DATABASE: sessionmaker(bind=snowflake_db_engine) } @contextmanager def db_session(db_name=config.SNOWFLAKE_DATABASE): """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. Args: db_name (session): Name of database to connect to Usage: with db_session() as session: session.execute(query) """ session = sessions[db_name]() try: yield session session.commit() except: session.rollback() raise finally: session.close() def db_session_wrap(func): """DB Session Wrappper. Creates a new session if one isn't passed in. """ @wraps(func) def wrapper(*args, **kwargs): session = kwargs.pop('session', None) if session: return func(*args, session=session, **kwargs) else: with db_session() as session: return func(*args, session=session, **kwargs) return wrapper