generated from john/python-template
100 lines
2.9 KiB
Python
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),
|
|
]
|