From d75083a6662d01d6a108878e5156901fa3777bcc Mon Sep 17 00:00:00 2001 From: John Lancaster <32917998+jsl12@users.noreply.github.com> Date: Sat, 1 Aug 2026 09:36:34 -0500 Subject: [PATCH] shutdown fixes --- src/transcription/config.py | 1 - src/transcription/db/session.py | 32 ++++++++++++++++++++++-------- src/transcription/services/base.py | 5 ++++- 3 files changed, 28 insertions(+), 10 deletions(-) diff --git a/src/transcription/config.py b/src/transcription/config.py index bf0b489..44e597f 100644 --- a/src/transcription/config.py +++ b/src/transcription/config.py @@ -74,7 +74,6 @@ class Settings(BaseSettings): # --- persistence --- database: DatabaseSettings = Field(default_factory=SqliteSettings) - database_url: str = "sqlite:///./transcription.db" bootstrap_schema_on_startup: bool = False sqlite_check_same_thread: bool = False diff --git a/src/transcription/db/session.py b/src/transcription/db/session.py index 9135f16..140a0c5 100644 --- a/src/transcription/db/session.py +++ b/src/transcription/db/session.py @@ -8,6 +8,7 @@ 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 @@ -25,8 +26,13 @@ def get_session_factory(database_url: str) -> SessionFactory: ) -def resolve_session_factory(database_url: str | None = None) -> SessionFactory: - return get_session_factory(database_url or get_database_url(get_settings())) +def resolve_session_factory( + database_url: str | None = None, + *, + settings: Settings | None = None, +) -> SessionFactory: + active_settings = settings or get_settings() + return get_session_factory(database_url or get_database_url(active_settings)) type SessionFactoryDep = Annotated[SessionFactory, Depends(resolve_session_factory)] @@ -40,15 +46,20 @@ async def dispose_session_factory(database_url: str) -> None: @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 - session_factory = resolve_session_factory(database_url) - async with session_factory() as owned_session: + active_session_factory = session_factory or resolve_session_factory( + database_url, + settings=settings, + ) + async with active_session_factory() as owned_session: yield owned_session @@ -58,9 +69,11 @@ type SessionScopeDep = Annotated[AsyncSession, Depends(session_scope)] @asynccontextmanager async def transaction_scope( *, + settings: Settings | None = None, database_url: str | None = None, - session: AsyncSessionTransaction | None = None, -) -> AsyncGenerator[AsyncSessionTransaction]: + 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(): @@ -71,8 +84,11 @@ async def transaction_scope( yield async_transaction return - session_factory = resolve_session_factory(database_url) - async with session_factory().begin() as owned_session: + 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 diff --git a/src/transcription/services/base.py b/src/transcription/services/base.py index d8c42a9..693c062 100644 --- a/src/transcription/services/base.py +++ b/src/transcription/services/base.py @@ -31,7 +31,10 @@ class ServiceBase(ABC): @asynccontextmanager async def _session_scope(self, session: AsyncSession | None = None): """Provide a transactional scope around a series of operations.""" - async with session_scope(session=session) as active_session: + async with session_scope( + session_factory=self.session_factory, + session=session, + ) as active_session: yield active_session async def _finalize(