from collections.abc import AsyncGenerator from contextlib import asynccontextmanager from functools import cache from typing import Annotated from fastapi import Depends from sqlalchemy.ext.asyncio import AsyncSessionTransaction from sqlalchemy.ext.asyncio import async_sessionmaker from sqlmodel.ext.asyncio.session import AsyncSession from ..config import Settings from ..config import get_settings from .engine import dispose_engine from .engine import get_database_url from .engine import get_engine type SessionFactory = async_sessionmaker[AsyncSession] @cache def get_session_factory(database_url: str) -> SessionFactory: return async_sessionmaker( bind=get_engine(database_url), class_=AsyncSession, expire_on_commit=False, ) def resolve_session_factory( database_url: str | None = None, *, settings: Settings | None = None, ) -> SessionFactory: if database_url is not None: return get_session_factory(database_url) return get_session_factory(get_database_url(settings or get_settings())) type SessionFactoryDep = Annotated[SessionFactory, Depends(resolve_session_factory)] async def dispose_session_factory(database_url: str) -> None: get_session_factory.cache_clear() await dispose_engine(database_url) @asynccontextmanager async def session_scope( *, settings: Settings | None = None, database_url: str | None = None, session_factory: SessionFactory | None = None, session: AsyncSession | None = None, ) -> AsyncGenerator[AsyncSession]: if session is not None: yield session return active_session_factory = session_factory or resolve_session_factory( database_url, settings=settings, ) async with active_session_factory() as owned_session: yield owned_session type SessionScopeDep = Annotated[AsyncSession, Depends(session_scope)] @asynccontextmanager async def transaction_scope( *, settings: Settings | None = None, database_url: str | None = None, session_factory: SessionFactory | None = None, session: AsyncSession | AsyncSessionTransaction | None = None, ) -> AsyncGenerator[AsyncSession | AsyncSessionTransaction]: match session: case AsyncSession() as async_session: if not async_session.in_transaction(): raise RuntimeError("A supplied session must have an active transaction") yield async_session return case AsyncSessionTransaction() as async_transaction: yield async_transaction return active_session_factory = session_factory or resolve_session_factory( database_url, settings=settings, ) async with active_session_factory.begin() as owned_session: yield owned_session type TransactionScopeDep = Annotated[ AsyncSession | AsyncSessionTransaction, Depends(transaction_scope), ]