"""DB Connector. Manages interactions with Relational Database (MySQL and SQLite). """ from contextlib import contextmanager from functools import wraps, partial from time import sleep from urllib import parse import sys from integration_scripts import logger from integration_scripts.connectors import sentry from integration_scripts.constants import mysql as mysql_consts from sqlalchemy import create_engine from sqlalchemy import exc from sqlalchemy import pool from sqlalchemy.orm import declarative_base from sqlalchemy.orm import scoped_session from sqlalchemy.orm import sessionmaker from integration_scripts import common_config from .pidguard import add_engine_pidguard # Shorten vars AR_MYSQL_USER = common_config.AR_MYSQL_USER AR_MYSQL_PASSWORD = common_config.AR_MYSQL_PASSWORD AR_MYSQL_HOST = common_config.AR_MYSQL_HOST AR_MYSQL_DB = common_config.AR_MYSQL_DB RDS_MYSQL_USER = common_config.RDS_MYSQL_USER RDS_MYSQL_PASSWORD = common_config.RDS_MYSQL_PASSWORD RDS_MYSQL_HOST = common_config.RDS_MYSQL_HOST RDS_MYSQL_DB = common_config.RDS_MYSQL_DB ENVIRONMENT = common_config.ENVIRONMENT TEST_ENVIRONMENT = common_config.TEST_ENVIRONMENT env_present = all(x not in ['', None] for x in [ AR_MYSQL_USER, AR_MYSQL_PASSWORD, AR_MYSQL_HOST, AR_MYSQL_DB ]) if not env_present and ENVIRONMENT != TEST_ENVIRONMENT: sys.exit( 'AR_MYSQL_USER, AR_MYSQL_PASSWORD, AR_MYSQL_HOST, AR_MYSQL_DB must not' ' be blank.\nTarget environment: {}'.format(ENVIRONMENT)) current_env = common_config.ENVIRONMENT # Database config if current_env == common_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.NullPool 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(prefix, field=None): """Print an error message about config files / environment and quit.""" prefixes = [prefix for _ in range(4)] 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: {}_MYSQL_USER {}_MYSQL_PASSWORD {}_MYSQL_HOST {}_MYSQL_DB Optional: (None) """.format(*prefixes)) if field: logger.error('{} has a bad value.'.format(field)) sys.exit() # Perform sanity checks against environment rds_mysl_db_config = {} if current_env != common_config.TEST_ENVIRONMENT: rds_mysl_db_config = { 'user': common_config.RDS_MYSQL_USER, 'password': parse.quote(common_config.RDS_MYSQL_PASSWORD), 'host': common_config.RDS_MYSQL_HOST, 'database': common_config.RDS_MYSQL_DB, } for key, value in rds_mysl_db_config.items(): if value == '' or value is None: config_complain('RDS', key) else: rds_mysl_db_config = { 'user': 'test_user', 'password': 'test_pass', 'host': 'test_host', 'database': 'test_db', 'port': 3306, } ar_mysl_db_config = {} if current_env != common_config.TEST_ENVIRONMENT: ar_mysl_db_config = { 'user': common_config.AR_MYSQL_USER, 'password': parse.quote(common_config.AR_MYSQL_PASSWORD), 'host': common_config.AR_MYSQL_HOST, 'database': common_config.AR_MYSQL_DB, } for key, value in ar_mysl_db_config.items(): if value == '' or value is None: config_complain('AR', key) else: ar_mysl_db_config = { 'user': 'test_user', 'password': 'test_pass', 'host': 'test_host', 'database': 'test_db', 'port': 3306, } def _create_engine(db_url, pool_class=None): """Create engine based on configuration settings.""" if not pool_class: pool_class = common_config.POOL_CLASS engine = create_engine( db_url, poolclass=pool_class, pool_pre_ping=True) add_engine_pidguard(engine) return engine # DEBUG - SPIT OUT THE HOSTNAME if common_config.DEBUG_LOG: logger.debug('Using Art Relations Host: {}', AR_MYSQL_HOST) logger.debug('Using RDS Host: {}', RDS_MYSQL_HOST) # Do not use these variables directly other than running unit tests ar_db_engine = _create_engine(mysql_consts.AR_DB_URL) rds_db_engine = _create_engine(mysql_consts.RDS_DB_URL, pool_class=POOL_CLASS) # please don't use sessions directly; # instead use db_session sessions = { mysql_consts.ART_RELATIONS_SESSION: scoped_session(sessionmaker(bind=ar_db_engine)), mysql_consts.INTEGRATIONS_RDS_SESSION: scoped_session(sessionmaker(bind=rds_db_engine)), } BaseModel = declarative_base() @contextmanager def db_session(session_name, read_only=False): """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: session_name (session): Name of database to connect to read_only (bool): Determines if commit is performed. Usage: with db_session() as session: session.execute(query) """ session = sessions[session_name]() try: yield session if not read_only: session.commit() except Exception: session.rollback() raise finally: session.close() @sentry.sentry_wrap def _db_session_wrap(function, target): """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_session(target) as session: return function(*args, session=session, **kwargs) return wrapper @sentry.sentry_wrap def _db_session_retry_wrap(function, target): """DB Session Wrapper with retry. Creates a new session if one isn't passed in, and does 1 retry if an internal error occurs while executing query. db_session_wrap is still preferred since this can rollback the session. """ @wraps(function) def wrapper(*args, **kwargs): session = kwargs.pop('session', None) try: if session: return function(*args, session=session, **kwargs) else: with db_session(target) as session: return function(*args, session=session, **kwargs) except exc.InternalError as e: logger.warning(e) sleep(mysql_consts.INTERNAL_ERROR_SLEEP_TIME) if session: session.rollback() return function(*args, session=session, **kwargs) else: with db_session(target) as session: return function(*args, session=session, **kwargs) return wrapper ar_db_session_wrap = partial( _db_session_wrap, target=mysql_consts.ART_RELATIONS_SESSION) rds_db_session_wrap = partial( _db_session_wrap, target=mysql_consts.INTEGRATIONS_RDS_SESSION) def invalidate_pool(db_engine): """Dispose of the current pool. Useful for multiprocessing.""" db_engine.dispose()