import os from contextlib import contextmanager from sqlalchemy import create_engine from sqlalchemy.orm import Session, sessionmaker from apollo_notifications.constants import ReplicaType from apollo_notifications.logger import logger class DBConfig: """Database configuration.""" DB_NAME = os.environ["MYSQL_DB_NAME"] DB_HOST = os.environ["MYSQL_DB_HOST"] DB_PORT = int(os.environ.get("MYSQL_DB_PORT", 3306)) DB_USER = os.environ["MYSQL_DB_USER"] DB_PASS = os.environ["MYSQL_DB_PASS"] DB_POOL_RECYCLE = int(os.environ.get("DB_POOL_RECYCLE", 7200)) DB_POOL_SIZE = int(os.environ.get("DB_POOL_SIZE", 1)) ALLOW_DB_REPLICATION = os.environ.get("ALLOW_DB_REPLICATION", False) if ALLOW_DB_REPLICATION: DB_SLAVE_NAME = os.environ.get("MYSQL_SLAVE_DB_NAME", DB_NAME) DB_SLAVE_HOST = os.environ.get("MYSQL_SLAVE_DB_HOST", "localhost") DB_SLAVE_PORT = int(os.environ.get("MYSQL_SLAVE_DB_PORT", 3306)) DB_SLAVE_USER = os.environ.get("MYSQL_SLAVE_DB_USER", DB_USER) DB_SLAVE_PASS = os.environ.get("MYSQL_SLAVE_DB_PASS", DB_PASS) _config = DBConfig() engine = create_engine( "mysql+pymysql://{user}:{password}@{host}:{port}/{database_name}" "?binary_prefix=true".format( user=_config.DB_USER, password=_config.DB_PASS, host=_config.DB_HOST, port=_config.DB_PORT, database_name=_config.DB_NAME), pool_recycle=_config.DB_POOL_RECYCLE, pool_size=_config.DB_POOL_SIZE ) engines = { ReplicaType.MASTER: engine, } if DBConfig.ALLOW_DB_REPLICATION: slave_engine = create_engine( "mysql+pymysql://{user}:{password}@{host}:{port}/{database_name}" "?binary_prefix=true".format( user=_config.DB_SLAVE_USER, password=_config.DB_SLAVE_PASS, host=_config.DB_SLAVE_HOST, port=_config.DB_SLAVE_PORT, database_name=_config.DB_SLAVE_NAME), pool_recycle=_config.DB_POOL_RECYCLE, pool_size=_config.DB_POOL_SIZE ) engines[ReplicaType.SLAVE] = slave_engine class RoutingSession(Session): def get_bind(self, mapper=None, clause=None): if clause is not None and clause.is_selectable and DBConfig.ALLOW_DB_REPLICATION: return engines[ReplicaType.SLAVE] elif clause is None: if mapper is not None: logger.info(f"ReplicaSession:: master replica choosen for {mapper.tables} select") logger.info("ReplicaSession:: master replica reason - clause is None") else: if not clause.is_selectable: logger.info("ReplicaSession:: master replica reason - clause is not selectable") elif not DBConfig.ALLOW_DB_REPLICATION: logger.info("ReplicaSession:: master replica reason - ALLOW_DB_REPLICATION is False") return engines[ReplicaType.MASTER] Session = sessionmaker(class_=RoutingSession) session = Session() @contextmanager def session_scope(): """Provide a transactional scope around a series of operations. """ try: yield session session.commit() except: session.rollback() raise finally: session.close()