from contextlib import asynccontextmanager from sqlalchemy import MetaData from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine from sqlalchemy.orm import Session, sessionmaker from server import config metadata = MetaData() engine = create_async_engine(config.DATABASE_URI, pool_recycle=config.DB_POOL_RECYCLE, pool_size=config.DB_POOL_SIZE) session_maker = sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) def session_proxy(session: Session): orig_execute = session.execute async def _execute(query, *args, flush: bool = True, commit: bool = False, **kwargs): result = await orig_execute(query, *args, **kwargs) if query is not None and hasattr(query, "is_selectable") and not query.is_selectable: if flush: await session.flush() if commit: await session.commit() return result session.execute = _execute return session @asynccontextmanager async def db_session(*args, session: Session = None, **kwargs): """Allows you to pass existing session or create new one if not passed.""" _session = None try: if session is None: _session = session_proxy(session_maker(*args, **kwargs)) yield session or _session if _session: await _session.commit() except Exception: if _session: await _session.rollback() raise finally: if _session: await _session.close()