From ec6617a1c44ddb6c75335cd735af51c8457c0c69 Mon Sep 17 00:00:00 2001 From: John Lancaster <32917998+jsl12@users.noreply.github.com> Date: Thu, 30 Jul 2026 23:16:17 -0500 Subject: [PATCH] updates --- src/transcription/app.py | 11 ++-- src/transcription/app_state.py | 2 +- src/transcription/config.py | 31 +++++++++-- src/transcription/db/__init__.py | 11 +++- src/transcription/db/operations.py | 26 ++++----- src/transcription/db/runtime.py | 64 +++-------------------- src/transcription/services/base.py | 14 ++--- src/transcription/ui/pages/jobs_page.py | 9 ++-- src/transcription/ui/pages/upload_page.py | 6 +-- src/transcription/worker.py | 4 +- tests/conftest.py | 20 ++++--- tests/test_db.py | 9 ++-- tests/ui/conftest.py | 14 ++--- 13 files changed, 101 insertions(+), 120 deletions(-) diff --git a/src/transcription/app.py b/src/transcription/app.py index 9aeafe2..6a70841 100644 --- a/src/transcription/app.py +++ b/src/transcription/app.py @@ -20,8 +20,10 @@ from .config import Settings from .config import configure_logging from .config import get_settings from .db import create_all -from .db import dispose_database_runtime from .db import initialize_database_runtime +from .db.engine import get_database_url +from .db.engine import resolve_engine +from .db.session import dispose_session_factory from .services import ServiceBundle from .services.jobs import JobService from .ui import register_pages @@ -40,7 +42,7 @@ async def _lifespan(app: FastAPI): app.state.runtime = initialize_database_runtime(settings=settings) if settings.should_bootstrap_schema: - await create_all(engine=app.state.runtime.engine) + await create_all(engine=resolve_engine(settings=settings)) settings.upload_dir.mkdir(parents=True, exist_ok=True) settings.prompt_dir.mkdir(parents=True, exist_ok=True) @@ -48,7 +50,10 @@ async def _lifespan(app: FastAPI): await _recover_stale_processing_jobs(app) async with AsyncExitStack() as stack: - stack.push_async_callback(dispose_database_runtime) + stack.push_async_callback( + dispose_session_factory, + database_url=get_database_url(settings), + ) stop_event, worker_notifier = await stack.enter_async_context( worker_consumer_lifespan( session_factory=app.state.runtime.session_factory, diff --git a/src/transcription/app_state.py b/src/transcription/app_state.py index 110f641..37a000a 100644 --- a/src/transcription/app_state.py +++ b/src/transcription/app_state.py @@ -7,7 +7,7 @@ from sqlalchemy.ext.asyncio import async_sessionmaker from sqlmodel.ext.asyncio.session import AsyncSession from transcription.db.runtime import DatabaseRuntime -from transcription.db.runtime import get_session_factory +from transcription.db.session import get_session_factory from transcription.worker import WorkerNotifier from transcription.worker import resolve_worker_notifier diff --git a/src/transcription/config.py b/src/transcription/config.py index 9280c68..3bb6618 100644 --- a/src/transcription/config.py +++ b/src/transcription/config.py @@ -9,12 +9,14 @@ import logging.config from enum import StrEnum from functools import cache from pathlib import Path +from typing import Annotated from typing import Any from typing import Literal +from pydantic import BaseModel from pydantic import Field +from pydantic import SecretStr from pydantic_settings import BaseSettings -from pydantic_settings import CliImplicitFlag from pydantic_settings import SettingsConfigDict logger = logging.getLogger(__name__) @@ -24,19 +26,41 @@ class Provider(StrEnum): OPENROUTER = "openrouter" +class SqliteSettings(BaseModel): + driver: Literal["sqlite"] = "sqlite" + path: str = "app.db" + + +class PostgresSettings(BaseModel): + driver: Literal["postgres"] = "postgres" + host: str + port: int = 5432 + database: str + user: str + password: SecretStr + + +DatabaseSettings = Annotated[ + SqliteSettings | PostgresSettings, + Field(discriminator="driver"), +] + + class Settings(BaseSettings): model_config = SettingsConfigDict( env_file=".env", env_file_encoding="utf-8", extra="ignore", cli_parse_args=True, + cli_implicit_flags=True, + cli_kebab_case=True, ) # --- NiceGUI Server --- host: str = "0.0.0.0" port: int = 8000 log_level: Literal["critical", "error", "warning", "info", "debug", "trace"] = "info" - reload: CliImplicitFlag[bool] = False + reload: bool = False # --- AI provider --- provider: Provider = Provider.OPENROUTER @@ -49,8 +73,9 @@ class Settings(BaseSettings): environment: Literal["development", "test", "production"] = "development" # --- persistence --- + database: DatabaseSettings = Field(default_factory=SqliteSettings) database_url: str = "sqlite:///./transcription.db" - bootstrap_schema_on_startup: bool | None = None + bootstrap_schema_on_startup: bool = False sqlite_check_same_thread: bool = False # --- filesystem paths --- diff --git a/src/transcription/db/__init__.py b/src/transcription/db/__init__.py index ff53ed2..961fbd7 100644 --- a/src/transcription/db/__init__.py +++ b/src/transcription/db/__init__.py @@ -1,6 +1,13 @@ from .operations import create_all from .runtime import dispose_database_runtime -from .runtime import get_session from .runtime import initialize_database_runtime +from .session import session_scope +from .session import transaction_scope -__all__ = ["create_all", "dispose_database_runtime", "get_session", "initialize_database_runtime"] +__all__ = [ + "create_all", + "dispose_database_runtime", + "initialize_database_runtime", + "session_scope", + "transaction_scope", +] diff --git a/src/transcription/db/operations.py b/src/transcription/db/operations.py index 4be28da..6562448 100644 --- a/src/transcription/db/operations.py +++ b/src/transcription/db/operations.py @@ -10,13 +10,25 @@ from sqlmodel import SQLModel from sqlmodel import select from sqlmodel.ext.asyncio.session import AsyncSession +from .engine import resolve_engine from .models import Job from .models import JobStatus -from .runtime import get_engine logger = logging.getLogger(__name__) +async def create_all(*, engine: AsyncEngine | None = None) -> None: + """Create all tables on the selected engine.""" + # Import models so SQLModel metadata is fully registered before bootstrap. + from transcription.db import models as _models # noqa: F401 + + active_engine = engine or resolve_engine() + async with active_engine.begin() as connection: + await connection.run_sync(SQLModel.metadata.create_all) + await connection.run_sync(_ensure_sqlite_compat_columns) + logger.debug("Database schema bootstrap complete for database_url=%s", active_engine.url) + + async def get_next_queued_job(*, session: AsyncSession) -> Job | None: """Get the next queued job, if any.""" result = await session.exec( @@ -28,18 +40,6 @@ async def get_next_queued_job(*, session: AsyncSession) -> Job | None: return result.first() -async def create_all(*, engine: AsyncEngine | None = None) -> None: - """Create all tables on the selected engine.""" - # Import models so SQLModel metadata is fully registered before bootstrap. - from transcription.db import models as _models # noqa: F401 - - active_engine = engine or get_engine() - async with active_engine.begin() as connection: - await connection.run_sync(SQLModel.metadata.create_all) - await connection.run_sync(_ensure_sqlite_compat_columns) - logger.debug("Database schema bootstrap complete for database_url=%s", active_engine.url) - - def _ensure_sqlite_compat_columns(connection: Connection) -> None: """Apply lightweight dev/test SQLite compatibility column patches. diff --git a/src/transcription/db/runtime.py b/src/transcription/db/runtime.py index cf6f6c0..7bcc10f 100644 --- a/src/transcription/db/runtime.py +++ b/src/transcription/db/runtime.py @@ -1,18 +1,16 @@ import logging -from collections.abc import AsyncGenerator -from contextlib import asynccontextmanager from contextvars import ContextVar from dataclasses import dataclass -from functools import partial from sqlalchemy.ext.asyncio import AsyncEngine from sqlalchemy.ext.asyncio import async_sessionmaker -from sqlalchemy.ext.asyncio import create_async_engine from sqlmodel.ext.asyncio.session import AsyncSession -from sqlmodel.pool import StaticPool from ..config import Settings from ..config import get_settings +from .engine import get_database_url +from .engine import get_engine +from .session import get_session_factory logger = logging.getLogger(__name__) @@ -37,33 +35,6 @@ async def dispose_database_runtime() -> None: _runtime.set(None) -def _to_async_database_url(database_url: str) -> str: - """Normalize configured database URL to an async SQLAlchemy driver URL.""" - if database_url.startswith("sqlite://") and not database_url.startswith("sqlite+aiosqlite://"): - return database_url.replace("sqlite://", "sqlite+aiosqlite://", 1) - if database_url.startswith("postgresql://") and not database_url.startswith("postgresql+asyncpg://"): - return database_url.replace("postgresql://", "postgresql+asyncpg://", 1) - return database_url - - -def _build_engine(settings: Settings) -> AsyncEngine: - database_url = _to_async_database_url(settings.database_url) - engine_factory = partial( - create_async_engine, - url=database_url, - echo=False, - pool_pre_ping=True, - ) - - if database_url.startswith("sqlite"): - sqlite_connect_settings = {"check_same_thread": settings.sqlite_check_same_thread} - engine_factory = partial(engine_factory, connect_args=sqlite_connect_settings) - if ":memory:" in database_url: - engine_factory = partial(engine_factory, poolclass=StaticPool) - - return engine_factory() - - def initialize_database_runtime(*, settings: Settings | None = None) -> DatabaseRuntime: """Initialize lifespan-owned async DB resources once per process.""" runtime = _runtime.get() @@ -71,33 +42,10 @@ def initialize_database_runtime(*, settings: Settings | None = None) -> Database return runtime active_settings = settings or get_settings() - engine = _build_engine(active_settings) - session_factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) + database_url = get_database_url(active_settings) + engine = get_engine(database_url) + session_factory = get_session_factory(database_url) runtime = DatabaseRuntime(engine=engine, session_factory=session_factory) _runtime.set(runtime) logger.debug("Initialized async database runtime for database_url=%s", engine.url) return runtime - - -def get_engine(settings: Settings | None = None) -> AsyncEngine: - """Return the current async SQLAlchemy engine.""" - runtime = _runtime.get() or initialize_database_runtime(settings=settings) - return runtime.engine - - -def get_session_factory(settings: Settings | None = None) -> async_sessionmaker[AsyncSession]: - """Return the shared async session factory.""" - runtime = _runtime.get() or initialize_database_runtime(settings=settings) - return runtime.session_factory - - -@asynccontextmanager -async def get_session( - *, - settings: Settings | None = None, - session_factory: async_sessionmaker[AsyncSession] | None = None, -) -> AsyncGenerator[AsyncSession]: - """Yield a database session and ensure cleanup.""" - active_session_factory = session_factory or get_session_factory(settings) - async with active_session_factory() as session: - yield session diff --git a/src/transcription/services/base.py b/src/transcription/services/base.py index ff15ad1..d8c42a9 100644 --- a/src/transcription/services/base.py +++ b/src/transcription/services/base.py @@ -8,7 +8,8 @@ from sqlmodel.ext.asyncio.session import AsyncSession from ..config import Settings from ..config import get_settings -from ..db.runtime import get_session_factory +from ..db.session import resolve_session_factory +from ..db.session import session_scope class ServiceBase(ABC): @@ -24,19 +25,14 @@ class ServiceBase(ABC): queue: asyncio.Queue | None = None, ): self.settings = get_settings() - self.session_factory = session_factory or get_session_factory() + self.session_factory = session_factory or resolve_session_factory() self.queue = queue or asyncio.Queue() @asynccontextmanager async def _session_scope(self, session: AsyncSession | None = None): """Provide a transactional scope around a series of operations.""" - if session is not None: - # Reuse the provided session if one is passed in - yield session - else: - # Otherwise, create a new session for this scope - async with self.session_factory() as new_session: - yield new_session + async with session_scope(session=session) as active_session: + yield active_session async def _finalize( self, diff --git a/src/transcription/ui/pages/jobs_page.py b/src/transcription/ui/pages/jobs_page.py index 3d222b2..0891501 100644 --- a/src/transcription/ui/pages/jobs_page.py +++ b/src/transcription/ui/pages/jobs_page.py @@ -4,10 +4,8 @@ from __future__ import annotations from uuid import UUID -from fastapi import Request from nicegui import ui -from transcription.app_state import resolve_session_factory from transcription.db.models import Job from transcription.db.models import JobStatus from transcription.db.models import Source @@ -17,6 +15,7 @@ from transcription.ui.components.app_shell import render_navigation_header from transcription.ui.components.error_presenter import show_error from transcription.ui.components.table.jobs import render_jobs_table +from ...db.session import SessionFactoryDep from ..components.document_panzoom import render_document_panzoom from ..components.table.jobs import JobTableRow from ..components.transcript import render_original_transcription_card @@ -27,8 +26,7 @@ def register_page() -> None: # noqa: PLR0915 """Register jobs list and detail routes.""" @ui.page("/jobs") - async def jobs_page(request: Request) -> None: - session_factory = resolve_session_factory(request.app.state) + async def jobs_page(session_factory: SessionFactoryDep) -> None: jobs_service = JobService(session_factory=session_factory) render_navigation_header(current_path="/jobs") @@ -51,8 +49,7 @@ def register_page() -> None: # noqa: PLR0915 await render_table() @ui.page("/jobs/{job_id}") - async def job_detail_page(job_id: str, request: Request) -> None: # noqa: PLR0915 - session_factory = resolve_session_factory(request.app.state) + async def job_detail_page(job_id: str, session_factory: SessionFactoryDep) -> None: # noqa: PLR0915 jobs_service = JobService(session_factory=session_factory) transcription_service = TranscriptionService(session_factory=session_factory) render_navigation_header(current_path="/jobs") diff --git a/src/transcription/ui/pages/upload_page.py b/src/transcription/ui/pages/upload_page.py index cc94645..6819a07 100644 --- a/src/transcription/ui/pages/upload_page.py +++ b/src/transcription/ui/pages/upload_page.py @@ -5,8 +5,7 @@ from __future__ import annotations from fastapi import Request from nicegui import ui -from transcription.app_state import resolve_session_factory -from transcription.db import get_session +from transcription.db import session_scope from transcription.services.store import create_upload_job from transcription.ui.components.app_shell import render_navigation_header from transcription.ui.components.upload import render_upload_widget @@ -19,10 +18,9 @@ def register_page() -> None: @ui.page("/upload", title="Upload Document") def upload_page(request: Request) -> None: render_navigation_header(current_path="/upload") - session_factory = resolve_session_factory(request.app.state) async def submit_upload(filename: str, file_bytes: bytes): - async with get_session(session_factory=session_factory) as session: + async with session_scope() as session: return await create_upload_job( filename=filename, file_bytes=file_bytes, diff --git a/src/transcription/worker.py b/src/transcription/worker.py index c5a39de..e8465ae 100644 --- a/src/transcription/worker.py +++ b/src/transcription/worker.py @@ -14,7 +14,7 @@ from uuid import UUID from sqlalchemy.ext.asyncio import async_sessionmaker from sqlmodel.ext.asyncio.session import AsyncSession -from transcription.db import get_session +from transcription.db import session_scope from transcription.errors import AppError from transcription.errors import classify_unexpected_error @@ -178,7 +178,7 @@ async def process_next_queued_job( ) if session is None: - async with get_session(session_factory=session_factory) as local_session: + async with session_scope(session_factory=session_factory) as local_session: return await process_next_queued_job_workflow(services=services, session=local_session) return await process_next_queued_job_workflow(services=services, session=session) diff --git a/tests/conftest.py b/tests/conftest.py index c65490b..d2edb0b 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -13,11 +13,12 @@ from sqlmodel.pool import StaticPool from transcription.config import Settings from transcription.config import get_settings +from transcription.db.engine import get_database_url +from transcription.db.engine import get_engine from transcription.db.operations import create_all -from transcription.db.runtime import dispose_database_runtime -from transcription.db.runtime import get_engine -from transcription.db.runtime import get_session -from transcription.db.runtime import get_session_factory +from transcription.db.session import dispose_session_factory +from transcription.db.session import get_session_factory +from transcription.db.session import session_scope from transcription.services.documents import DocumentService from transcription.services.jobs import JobService @@ -39,23 +40,26 @@ def session(): async def default_settings(): """Provide default settings for tests.""" settings = get_settings(database_url="sqlite:///:memory:") - await create_all(engine=get_engine(settings=settings)) + db_url = get_database_url(settings) + await create_all(engine=get_engine(database_url=db_url)) return settings @pytest_asyncio.fixture async def async_session(default_settings: Settings): """Provide a clean asynchronous database session for async tests.""" - async with get_session(settings=default_settings) as async_session: + db_url = get_database_url(default_settings) + async with session_scope(database_url=db_url) as async_session: yield async_session - await dispose_database_runtime() + await dispose_session_factory(db_url) @pytest.fixture def default_session_factory(default_settings: Settings): """Provide a base fixture for tests that require database access.""" - session_factory = get_session_factory(settings=default_settings) + db_url = get_database_url(default_settings) + session_factory = get_session_factory(database_url=db_url) return session_factory diff --git a/tests/test_db.py b/tests/test_db.py index 1b5c6aa..41c1f42 100644 --- a/tests/test_db.py +++ b/tests/test_db.py @@ -4,17 +4,18 @@ import pytest from sqlalchemy import inspect from transcription.config import Settings +from transcription.config import SqliteSettings from transcription.db import create_all from transcription.db import dispose_database_runtime -from transcription.db import get_session from transcription.db import initialize_database_runtime +from transcription.db import session_scope @pytest.mark.asyncio async def test_create_all_creates_expected_tables(tmp_path): settings = Settings( openrouter_api_key="test-key", - database_url=f"sqlite:///{tmp_path / 'schema.db'}", + database=SqliteSettings(path=str(tmp_path / "schema.db")), environment="test", ) runtime = initialize_database_runtime(settings=settings) @@ -36,13 +37,13 @@ async def test_create_all_creates_expected_tables(tmp_path): async def test_get_session_yields_async_session(tmp_path): settings = Settings( openrouter_api_key="test-key", - database_url=f"sqlite:///{tmp_path / 'session.db'}", + database=SqliteSettings(path=str(tmp_path / "session.db")), environment="test", ) initialize_database_runtime(settings=settings) try: - async with get_session(settings=settings) as session: + async with session_scope(settings=settings) as session: assert session is not None finally: await dispose_database_runtime() diff --git a/tests/ui/conftest.py b/tests/ui/conftest.py index 5b288b0..71bdf8e 100644 --- a/tests/ui/conftest.py +++ b/tests/ui/conftest.py @@ -4,6 +4,7 @@ from __future__ import annotations import asyncio from collections.abc import Callable +from collections.abc import Generator from pathlib import Path from uuid import UUID @@ -14,10 +15,10 @@ from sqlmodel import delete from transcription.app import create_app from transcription.config import Settings -from transcription.config import _settings +from transcription.config import SqliteSettings from transcription.db import create_all -from transcription.db import get_session from transcription.db import initialize_database_runtime +from transcription.db import session_scope from transcription.db.models import Document from transcription.db.models import Job from transcription.db.models import JobStatus @@ -28,18 +29,17 @@ RevisionSeed = str @pytest.fixture(scope="session") -def app_client(tmp_path_factory: pytest.TempPathFactory) -> tuple[FastAPI, TestClient]: +def app_client(tmp_path_factory: pytest.TempPathFactory) -> Generator[tuple[FastAPI, TestClient]]: """Provide a real application and test client backed by in-memory SQLite.""" tmp_path = tmp_path_factory.mktemp("ui") settings = Settings( openrouter_api_key="test-key", - database_url="sqlite:///:memory:", + database=SqliteSettings(path=":memory:"), environment="test", bootstrap_schema_on_startup=True, upload_dir=tmp_path / "uploads", prompt_dir=tmp_path / "prompts", ) - _settings.set(settings) app = create_app() app.state.runtime = initialize_database_runtime(settings=settings) @@ -54,7 +54,7 @@ def clear_ui_database(app_client: tuple[FastAPI, TestClient]) -> None: app, _ = app_client async def _clear() -> None: - async with get_session(session_factory=app.state.runtime.session_factory) as session: + async with session_scope() as session: await session.exec(delete(Revision)) await session.exec(delete(Source)) await session.exec(delete(Job)) @@ -80,7 +80,7 @@ def seed_job(app_client: tuple[FastAPI, TestClient]) -> Callable[..., UUID]: source_file: Path | None = None, ) -> UUID: async def _insert() -> UUID: - async with get_session(session_factory=app.state.runtime.session_factory) as session: + async with session_scope() as session: stored_path = app.state.settings.upload_dir / filename stored_path.parent.mkdir(parents=True, exist_ok=True) source_path = source_file or fixtures_dir / "small_png.png"