"""MySQL connector for art_relations.""" import functools from contextlib import contextmanager from typing import Iterator from sqlalchemy import create_engine from sqlalchemy.orm import DeclarativeBase, Session, scoped_session, sessionmaker from sqlalchemy.pool import StaticPool from contributor import config def _create_engine(): if config.AR_DB_POOL_CLASS is StaticPool: return create_engine( config.AR_DB_URL, poolclass=StaticPool, connect_args={"check_same_thread": False}, ) return create_engine( config.AR_DB_URL, pool_size=config.AR_DB_POOL_SIZE, max_overflow=config.AR_DB_POOL_MAX_OVERFLOW, pool_recycle=config.AR_DB_POOL_RECYCLE, ) ar_db_engine = _create_engine() _session_factory = scoped_session(sessionmaker(ar_db_engine)) class BaseModel(DeclarativeBase): pass @contextmanager def ar_db_session(read_only: bool = True) -> Iterator[Session]: """Provide a transaction scope for art_relations queries.""" session = _session_factory() try: yield session if not read_only: session.commit() except Exception: session.rollback() raise finally: _session_factory.remove() def db_session_wrap(func): """Function wrapper for shared DB Session. Creates a new session if one isn't passed in. This lets us share/pass-in a common session across multiple functions, making them all transactional. """ @functools.wraps(func) def wrapper(*args, **kwargs): session = kwargs.pop("session", None) if session: return func(*args, session=session, **kwargs) else: with ar_db_session() as session: return func(*args, session=session, **kwargs) return wrapper