"""Integration test configuration.""" import logging import os from contextlib import asynccontextmanager from posixpath import abspath, dirname, join, pardir import sqlalchemy from environs import Env from secrets_manager.python_ext import PythonSecretsManager from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine env = Env() logger = logging.getLogger(__name__) # Load environment variables from a .env file if present dotenv_path = abspath(join(dirname(__file__), pardir, "../.env")) env.read_env(dotenv_path, recurse=False) secrets_manager_client = PythonSecretsManager( application_context=False, environment="qa", service_name="ows-product-staging-integration-test", force_remote=True, ) MYSQL_DB_PASSWORD = os.environ.get( "INTEGRATION_DB_PASSWORD", None ) or secrets_manager_client.get_cred("MYSQL_DB_PASSWORD") POOL_CLASS = sqlalchemy.pool.QueuePool POOL_SIZE = 15 POOL_RECYCLE_MS = 3600 # Avoids connections going stale POOL_MAX_OVERFLOW = 0 DB_URL = "mysql+aiomysql://{user}:{password}@{host}/{db_name}?charset={charset}".format( user=os.environ.get("INTEGRATION_DB_USER"), password=MYSQL_DB_PASSWORD, host=os.environ.get("INTEGRATION_DB_HOST"), db_name=os.environ.get("INTEGRATION_DB_NAME"), charset="utf8mb4", ) INTEGRATION_S3_BUCKET = os.environ.get("INTEGRATION_S3_BUCKET") assert INTEGRATION_S3_BUCKET def get_engine(): """Create engine based on config settings.""" if POOL_CLASS == sqlalchemy.pool.QueuePool: _db_engine = create_async_engine( DB_URL, pool_size=POOL_SIZE, max_overflow=POOL_MAX_OVERFLOW, pool_recycle=POOL_RECYCLE_MS, ) return _db_engine return create_async_engine(DB_URL, poolclass=POOL_CLASS) db_engine = get_engine() db_session_maker = async_sessionmaker(db_engine, expire_on_commit=False) @asynccontextmanager async 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_maker() try: yield session await session.commit() except: # noqa await session.rollback() raise finally: await session.close() QA_BASE_URL = ( os.environ.get("QA_BASE_URL") or "https://qa-ows-product-staging.theorchard.io" )