Files
transcription/src/transcription/db/session.py
T

100 lines
2.9 KiB
Python

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),
]