From 90669738ccb9ba70e6e5544b276aa33f3c1a3839 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 15:43:51 +0200 Subject: [PATCH 01/14] docs(api/adr): add ADR-002 no global mutable state --- .../adr/000-vertical-slice-architecture.md | 2 +- api/docs/adr/002-no-global-mutable-state.md | 54 +++++++++++++++++++ api/docs/architecture.md | 2 +- 3 files changed, 56 insertions(+), 2 deletions(-) create mode 100644 api/docs/adr/002-no-global-mutable-state.md diff --git a/api/docs/adr/000-vertical-slice-architecture.md b/api/docs/adr/000-vertical-slice-architecture.md index 3c67d6ba..c5ad179a 100644 --- a/api/docs/adr/000-vertical-slice-architecture.md +++ b/api/docs/adr/000-vertical-slice-architecture.md @@ -60,7 +60,7 @@ damnit_api/ ├── main.py # entrypoint: env/args → Settings → create_app ├── app.py # composition root: create_app(settings), lifespan, │ # DI wiring, exception handlers, middleware -├── state.py # AppState + factories (no domain classes) +├── state.py # AppState + factories (no domain classes); see ADR-002 ├── settings.py # Settings models only ├── logging.py # structlog configuration + request-logging middleware │ diff --git a/api/docs/adr/002-no-global-mutable-state.md b/api/docs/adr/002-no-global-mutable-state.md new file mode 100644 index 00000000..b716ab02 --- /dev/null +++ b/api/docs/adr/002-no-global-mutable-state.md @@ -0,0 +1,54 @@ +--- +date: 2026-07-07 +--- + +# ADR-002 - No Global Mutable State: `AppState`, Factories, One Composition Root + +## Context and Problem Statement + +This service has a few long-lived runtime dependencies: + +- application database engine and session factory +- (authenticated) clients (MyMdC, OAuth) +- per-proposal DAMNIT database accessors +- OAuth token store +- subscription cursors + +These are currently implemented as module-level singletons bootstrapped at startup (`_db.__ENGINE`, `_mymdc.CLIENT`, `auth.__CLIENT`, `TOKEN_STORE`, a registry metaclass). + +This is problematic as any function can reach any dependency, initialisation order becomes critical and is quite opaque, tests must mutate or clear process-wide state between cases, multiple app instances with different configurations cannot coexist in one process, in-process state looks shareable but silently breaks under multiple workers, etc... + +The aim of this ADR is to consider different options for managing these dependencies. + +## Considered Options + +- Keep module-level singletons, initialised by startup hooks +- A single typed state container built in the application lifespan, with dependencies injected everywhere else + +## Decision Outcome + +Chosen option: a single frozen `AppState` dataclass built once in the lifespan. This makes initialisation typed and order-explicit (the constructor is the startup contract), lets tests inject doubles by constructing state rather than patching modules, and makes deliberately process-local state visible instead of hidden in module scope. + +### Consequences + +- Good: initialisation order and the full dependency set are explicit in one place + - Good: missing dependency results in construction error at startup, not a `None` at request time. +- Good: tests build `AppState` with fakes (a mock MyMdC client, an in-memory token store) + - Good: removes/reduces need to monkeypatch modules and clear cache between tests. +- Bad: (ish?) dependencies must be specified through signatures instead of imported where needed + - This is kind of the whole point, but it does mean there is more code to do the same thing (importing a global is easier/shorter) + +## Details + +### The rules + +1. All runtime dependencies stored in a single frozen `AppState` dataclass (`state.py`) which is constructed once in the application lifespan and attached to `app.state`. This currently contains: + - App DB engine/sessionmaker, MyMdC client, OAuth client (`None` when auth is disabled), DAMNIT DB registry, token store, subscription cursors. +2. Each field is built by a pure factory function (`create_*`) taking `Settings` as explicit arguments. No factory reads module state or has side effects beyond constructing its object. +3. There is exactly one composition/setup root: the app entrypoint and its lifespan (target shape: `create_app(settings)`, see the ADR-000 layout). + - This is the only place that reads settings to select implementations. + - Handlers, resolvers, and services receive dependencies via DI or plain parameters, they must **never** import them. + +4. Caches must be treated as state. + - Any cache must be owned by an object that is itself created by a factory and reachable from `AppState`. + - Module-level and class-level cache decorators on application code are banned. diff --git a/api/docs/architecture.md b/api/docs/architecture.md index 2eb79af2..421db38c 100644 --- a/api/docs/architecture.md +++ b/api/docs/architecture.md @@ -26,7 +26,7 @@ For more information, see [ADR-000](adr/000-vertical-slice-architecture.md). | `appdb/` | The app's own database (infrastructure) | Models, engine/session plumbing for `dw_api.sqlite` | Planned | `_db/` | | `mymdc/` | MyMdC client (infrastructure) | Ports, clients, vendored models | Planned | `_mymdc/` | | `core/` | Cross-cutting, framework-free | Shared error classes (see [ADR-001](adr/001-error-classes.md)), `DamnitType`, value types, converters | Planned | `shared/` + `utils.py` | -| `main.py` / `app.py` / `state.py` | Composition root | `AppState`, `create_*` factories, `create_app()` - the only place that may import everything and read settings | Partial | `main.py` only | +| `main.py` / `app.py` / `state.py` | Composition root | `AppState`, `create_*` factories (see [ADR-002](adr/002-no-global-mutable-state.md)), `create_app()` - the only place that may import everything and read settings | Partial | `main.py` only | Where new code goes: From 4f77d3c9c24576a6a2a1edfdea3ee90a895d2540 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 15:45:05 +0200 Subject: [PATCH 02/14] docs(api/adr): add ADR-003 injected settings --- .../adr/000-vertical-slice-architecture.md | 2 +- api/docs/adr/003-injected-settings.md | 36 +++++++++++++++++++ api/docs/architecture.md | 2 +- 3 files changed, 38 insertions(+), 2 deletions(-) create mode 100644 api/docs/adr/003-injected-settings.md diff --git a/api/docs/adr/000-vertical-slice-architecture.md b/api/docs/adr/000-vertical-slice-architecture.md index c5ad179a..43981105 100644 --- a/api/docs/adr/000-vertical-slice-architecture.md +++ b/api/docs/adr/000-vertical-slice-architecture.md @@ -61,7 +61,7 @@ damnit_api/ ├── app.py # composition root: create_app(settings), lifespan, │ # DI wiring, exception handlers, middleware ├── state.py # AppState + factories (no domain classes); see ADR-002 -├── settings.py # Settings models only +├── settings.py # Settings models only; see ADR-003 ├── logging.py # structlog configuration + request-logging middleware │ ├── core/ # framework-free, imports nothing app-specific: diff --git a/api/docs/adr/003-injected-settings.md b/api/docs/adr/003-injected-settings.md new file mode 100644 index 00000000..73405c56 --- /dev/null +++ b/api/docs/adr/003-injected-settings.md @@ -0,0 +1,36 @@ +--- +date: 2026-07-07 +--- + +# ADR-003 - Settings: Injected Configuration, No Import-Time Singleton + +## Context and Problem Statement + +Configuration (auth credentials, database paths, MyMdC endpoints, local-mode selection) must be available throughout the application. There are two ways to provide it: a module-level singleton importable from anywhere, or an object constructed once and passed explicitly. + +An import-time singleton has structural costs. Importing any module transitively requires a valid environment. Validation errors then fire at import, which breaks tooling, tests, and scripting contexts that never run the app. Configuration access is invisible in signatures, so nothing documents what depends on what. The application factory cannot be called twice with different configurations in one process, which blocks table-driven app tests. Modules dodge import-time failures with function-body imports, which then ossify into circular-import workarounds. + +## Considered Options + +- Keep the module-level `settings = Settings()` singleton, imported wherever configuration is needed +- Settings models only in `settings.py`; one instance constructed at the entrypoint and threaded explicitly through the composition root + +## Decision Outcome + +Chosen option: settings models only, constructed once at the entrypoint and injected, because it makes configuration dependencies visible in signatures, keeps imports environment-free, and allows multiple differently-configured app instances in one process. + +### Consequences + +- Good: tests build `Settings(...)` directly (pydantic-settings accepts init kwargs) and get components wired for that config - no env patching, no module reloads. +- Good: mode-dependent behaviour is forced up to the composition root, because nothing deeper can consult configuration without it showing up in a signature. +- Bad: signatures grow explicit parameters; that visibility is the point, but it is more ceremony than importing a global. + +## Details + +### The rules + +1. `settings.py` defines models only (`Settings` and its nested models). The target state has no module-level instance. +2. `Settings` is constructed exactly once, at the entrypoint, and passed into the composition root. The composition root threads it into factories; everything else receives either the settings object or - preferably - the specific values it needs as plain parameters. +3. Environment handling: `DW_API_` prefix, `__` nested delimiter, `.env` support. Local mode is a derived property read only in the composition root. +4. Defaults must be production-safe: no paths into `tests/`, no writes into the source tree. Development conveniences belong in `.env` files and documentation, not in field defaults. +5. Verification: the `settings` instance is imported only by the composition root and tests; importing any other module with a bare environment succeeds. diff --git a/api/docs/architecture.md b/api/docs/architecture.md index 421db38c..17b4c171 100644 --- a/api/docs/architecture.md +++ b/api/docs/architecture.md @@ -55,7 +55,7 @@ The key rules are: - Importing another package's `_underscore` name is always wrong. 2. **Composition root is the top:** it may import everything, but nothing is allowed to import it. - If importing a slice from the composition root forces a function-body import to avoid cycles, the type probably belongs in `core/`. -3. **Composition root reads settings:** everything else receives configuration as parameters. +3. **Composition root reads settings:** everything else receives configuration as parameters (see [ADR-003](adr/003-injected-settings.md)). 4. **Authorisation applied at the edge:** routes and resolvers use dependencies and permission classes. - This means that services should not apply authorisation rules themselves. 5. **No `if settings.is_local:` outside the composition root:** Local mode is selected by composition, not conditionals throughout the codebase. From d158f771ec66e12039dd23d3d9707ce7573371ee Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 15:49:40 +0200 Subject: [PATCH 03/14] refactor(api/state): add AppState container and factory functions --- api/src/damnit_api/state.py | 58 +++++++++++++++++++++++++++++++++++++ 1 file changed, 58 insertions(+) create mode 100644 api/src/damnit_api/state.py diff --git a/api/src/damnit_api/state.py b/api/src/damnit_api/state.py new file mode 100644 index 00000000..d998277c --- /dev/null +++ b/api/src/damnit_api/state.py @@ -0,0 +1,58 @@ +"""Typed application state and pure factory functions. + +All long-lived runtime dependencies live on the frozen :class:`AppState`, +built once in the application lifespan and attached to ``app.state``. Each +field is produced by a pure ``create_*`` factory taking :class:`Settings` +(or already-built collaborators) as explicit arguments. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING + +from fastapi import ( + Request, # noqa: TC002 - FastAPI DI inspects annotations at runtime +) +from sqlalchemy.ext.asyncio import AsyncEngine, async_sessionmaker, create_async_engine +from sqlmodel.ext.asyncio.session import AsyncSession + +if TYPE_CHECKING: + from ._mymdc.clients import MyMdCClient + from .shared.settings import Settings + + +@dataclass(frozen=True) +class AppState: + db_engine: AsyncEngine + db_sessionmaker: async_sessionmaker[AsyncSession] + mymdc_client: MyMdCClient + + +def create_db_engine(settings: Settings) -> AsyncEngine: + db_url = f"sqlite+aiosqlite:///{settings.db_path}" + return create_async_engine(db_url, echo=False, future=True) + + +def create_db_sessionmaker(engine: AsyncEngine) -> async_sessionmaker[AsyncSession]: + return async_sessionmaker(bind=engine, class_=AsyncSession, expire_on_commit=False) + + +def create_mymdc_client(settings: Settings) -> MyMdCClient: + from ._mymdc import clients + from ._mymdc.settings import MyMdCHTTPSettings, MyMdCMockSettings + + match settings.mymdc: + case MyMdCHTTPSettings(): + auth = clients.MyMdCAuth.model_validate(settings.mymdc.model_dump()) + return clients.MyMdCClientAsync(auth) + case MyMdCMockSettings(): + return clients.MyMdCClientMock.model_validate(settings.mymdc.model_dump()) + case _: + msg = "Invalid MyMdC configuration" + raise ValueError(msg) + + +def get_app_state(request: Request) -> AppState: + """FastAPI dependency: the application's :class:`AppState`.""" + return request.app.state.app_state From 6df0ae6e161cc1f0aeeaf0dd0d2b244d50aaa6db Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 15:52:04 +0200 Subject: [PATCH 04/14] refactor(api): build AppState in lifespan alongside legacy bootstraps --- api/src/damnit_api/main.py | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/api/src/damnit_api/main.py b/api/src/damnit_api/main.py index 7ad1f5cd..8af87f1a 100644 --- a/api/src/damnit_api/main.py +++ b/api/src/damnit_api/main.py @@ -16,6 +16,12 @@ def create_app(): from . import _db, _logging, _mymdc, auth, contextfile, get_logger, metadata from .shared import errors, gql from .shared.settings import settings + from .state import ( + AppState, + create_db_engine, + create_db_sessionmaker, + create_mymdc_client, + ) logger = get_logger("lifespan") @@ -34,6 +40,13 @@ async def lifespan(app: FastAPI): for bs in bootstraps: tg.create_task(bs(settings)) + db_engine = create_db_engine(settings) + app.state.app_state = AppState( + db_engine=db_engine, + db_sessionmaker=create_db_sessionmaker(db_engine), + mymdc_client=create_mymdc_client(settings), + ) + if settings.is_local: app.router.include_router(auth.noauth_router) else: From 38a9e132d13b36deb331e4ee3da3e7a15374d5ef Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 15:53:39 +0200 Subject: [PATCH 05/14] refactor(api/db): serve app DB sessions from AppState, drop bootstrap --- api/src/damnit_api/_db/__init__.py | 16 ------- api/src/damnit_api/_db/bootstrap.py | 65 -------------------------- api/src/damnit_api/_db/dependencies.py | 10 ++-- api/src/damnit_api/main.py | 4 +- 4 files changed, 7 insertions(+), 88 deletions(-) delete mode 100644 api/src/damnit_api/_db/bootstrap.py diff --git a/api/src/damnit_api/_db/__init__.py b/api/src/damnit_api/_db/__init__.py index 33d359cf..c999fd9d 100644 --- a/api/src/damnit_api/_db/__init__.py +++ b/api/src/damnit_api/_db/__init__.py @@ -1,17 +1 @@ """Database package exports for damnit_api._db.""" - -from sqlalchemy.ext.asyncio import ( - AsyncEngine, - async_sessionmaker, -) -from sqlmodel.ext.asyncio.session import AsyncSession - -from .bootstrap import bootstrap - -global __ENGINE, __SESSION_LOCAL - -__ENGINE: AsyncEngine = None # pyright: ignore[reportAssignmentType] - -__SESSION_LOCAL: async_sessionmaker[AsyncSession] = None # pyright: ignore[reportAssignmentType] - -__all__ = ["bootstrap"] diff --git a/api/src/damnit_api/_db/bootstrap.py b/api/src/damnit_api/_db/bootstrap.py deleted file mode 100644 index f3215f60..00000000 --- a/api/src/damnit_api/_db/bootstrap.py +++ /dev/null @@ -1,65 +0,0 @@ -"""Bootstrapping for the database layer. - -This module configures an async SQLModel engine and sessionmaker and provides helpers -and dependencies for FastAPI integration. -""" - -from typing import TYPE_CHECKING - -from sqlalchemy.ext.asyncio import ( - async_sessionmaker, - create_async_engine, -) -from sqlmodel.ext.asyncio.session import AsyncSession - -from .. import get_logger - -logger = get_logger() - -if TYPE_CHECKING: - from ..shared.settings import Settings - - -async def bootstrap(settings: "Settings") -> None: - """Initialize the async engine and sessionmaker. - - This function is intended to be called during application startup. - """ - import damnit_api._db - - if damnit_api._db.__ENGINE is not None: - await logger.awarning("Database engine already initialized") - return - - db_url = f"sqlite+aiosqlite:///{settings.db_path}" - await logger.ainfo("Configuring database engine", db_url=db_url) - damnit_api._db.__ENGINE = create_async_engine(str(db_url), echo=False, future=True) - damnit_api._db.__SESSION_LOCAL = async_sessionmaker( - bind=damnit_api._db.__ENGINE, class_=AsyncSession, expire_on_commit=False - ) - - -def init_db() -> None: - """Create database tables from SQLModel metadata.""" - from sqlalchemy import create_engine - from sqlmodel import SQLModel - - engine = create_engine(f"sqlite:///{settings.db_path}", echo=True, future=True) - - SQLModel.metadata.create_all(engine) - - -if __name__ == "__main__": - import asyncio - - from ..metadata import models as _md # noqa: F401 - from ..shared.settings import settings - - asyncio.run(bootstrap(settings)) - - if settings.db_path.exists(): - logger.warning("Database file already exists.", db_path=settings.db_path) - if input("Delete table? [y/n]: ").lower() != "y": - exit(0) - - init_db() diff --git a/api/src/damnit_api/_db/dependencies.py b/api/src/damnit_api/_db/dependencies.py index 02adb4f8..54e5f029 100644 --- a/api/src/damnit_api/_db/dependencies.py +++ b/api/src/damnit_api/_db/dependencies.py @@ -3,16 +3,16 @@ from collections.abc import AsyncIterator from typing import Annotated -from fastapi import Depends +from fastapi import Depends, Request from sqlmodel.ext.asyncio.session import AsyncSession -import damnit_api._db +from ..state import get_app_state -async def get_session() -> AsyncIterator[AsyncSession]: - """Provide a database session for FastAPI dependencies.""" +async def get_session(request: Request) -> AsyncIterator[AsyncSession]: + """Provide a database session from the application state.""" - async with damnit_api._db.__SESSION_LOCAL() as session: + async with get_app_state(request).db_sessionmaker() as session: yield session diff --git a/api/src/damnit_api/main.py b/api/src/damnit_api/main.py index 8af87f1a..75a8f015 100644 --- a/api/src/damnit_api/main.py +++ b/api/src/damnit_api/main.py @@ -13,7 +13,7 @@ def create_app(): - from . import _db, _logging, _mymdc, auth, contextfile, get_logger, metadata + from . import _logging, _mymdc, auth, contextfile, get_logger, metadata from .shared import errors, gql from .shared.settings import settings from .state import ( @@ -35,7 +35,7 @@ async def lifespan(app: FastAPI): logger.info("Starting application lifespan") - bootstraps = [_mymdc.bootstrap, auth.bootstrap, _db.bootstrap] + bootstraps = [_mymdc.bootstrap, auth.bootstrap] async with TaskGroup() as tg: for bs in bootstraps: tg.create_task(bs(settings)) From 7e44176595472581e8a52dabd553b57833ba3d45 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 15:56:24 +0200 Subject: [PATCH 06/14] refactor(api/mymdc): inject MyMdC client from AppState --- api/src/damnit_api/_mymdc/__init__.py | 15 -------- api/src/damnit_api/_mymdc/bootstrap.py | 42 ----------------------- api/src/damnit_api/_mymdc/dependencies.py | 13 +++++-- api/src/damnit_api/_mymdc/ports.py | 14 -------- api/src/damnit_api/main.py | 8 ++--- 5 files changed, 12 insertions(+), 80 deletions(-) delete mode 100644 api/src/damnit_api/_mymdc/bootstrap.py diff --git a/api/src/damnit_api/_mymdc/__init__.py b/api/src/damnit_api/_mymdc/__init__.py index 797c7b45..226cab5b 100644 --- a/api/src/damnit_api/_mymdc/__init__.py +++ b/api/src/damnit_api/_mymdc/__init__.py @@ -5,18 +5,3 @@ This module should **not** directly expose any routes, it should **only** be used internally by other modules to interact with the MyMdC API. """ - -from typing import TYPE_CHECKING - -from .bootstrap import bootstrap as bootstrap - -if TYPE_CHECKING: - from . import clients - -global CLIENT - -CLIENT: "clients.MyMdCClientAsync" = None # pyright: ignore[reportAssignmentType] -"""Global/singleton MyMdC client instance, configured by [`.bootstrap`]""" - - -__all__ = ["bootstrap"] diff --git a/api/src/damnit_api/_mymdc/bootstrap.py b/api/src/damnit_api/_mymdc/bootstrap.py deleted file mode 100644 index 7287b132..00000000 --- a/api/src/damnit_api/_mymdc/bootstrap.py +++ /dev/null @@ -1,42 +0,0 @@ -"""Bootstrapping code for MyMdC client instantiation.""" - -from typing import TYPE_CHECKING - -from .. import get_logger -from . import clients -from .settings import MyMdCHTTPSettings, MyMdCMockSettings - -if TYPE_CHECKING: - from ..shared.settings import Settings - -logger = get_logger() - - -async def bootstrap(settings: "Settings"): - """Bootstrap MyMdC client based on the current settings, saving the instance to the - global [`damnit._mymdc.__CLIENT`] variable. - - Calling this multiple times will have no effect after the first call, but will log a - warning.""" - import damnit_api._mymdc - - if damnit_api._mymdc.CLIENT is None: - await logger.ainfo("Initialising MyMdC client") - damnit_api._mymdc.CLIENT = await _init(settings) - else: - await logger.awarning("MyMdC client already initialised") - - -async def _init(settings: "Settings") -> clients.MyMdCClient: - """Create MyMdC client based on settings.""" - match settings.mymdc: - case MyMdCHTTPSettings(): - await logger.ainfo("Creating MyMdC client from credentials") - auth = clients.MyMdCAuth.model_validate(settings.mymdc.model_dump()) - return clients.MyMdCClientAsync(auth) - case MyMdCMockSettings(): - await logger.ainfo("Creating Mock MyMdC client") - return clients.MyMdCClientMock.model_validate(settings.mymdc.model_dump()) - case _: - msg = "Invalid MyMdC configuration" - raise ValueError(msg) diff --git a/api/src/damnit_api/_mymdc/dependencies.py b/api/src/damnit_api/_mymdc/dependencies.py index a5b610de..a9ad18e1 100644 --- a/api/src/damnit_api/_mymdc/dependencies.py +++ b/api/src/damnit_api/_mymdc/dependencies.py @@ -1,8 +1,15 @@ from typing import Annotated -from fastapi import Depends +from fastapi import Depends, Request -from . import clients, ports +from ..state import get_app_state +from . import clients -MyMdCClient = Annotated[clients.MyMdCClient, Depends(ports.MyMdCPort.from_global)] + +def get_mymdc_client(request: Request) -> "clients.MyMdCClient": + """Provide the MyMdC client from the application state.""" + return get_app_state(request).mymdc_client + + +MyMdCClient = Annotated[clients.MyMdCClient, Depends(get_mymdc_client)] """Type alias for the MyMdC client dependency.""" diff --git a/api/src/damnit_api/_mymdc/ports.py b/api/src/damnit_api/_mymdc/ports.py index 66e55b62..5a518f31 100644 --- a/api/src/damnit_api/_mymdc/ports.py +++ b/api/src/damnit_api/_mymdc/ports.py @@ -1,7 +1,6 @@ """MyMdC Ports (Interfaces) definitions.""" from abc import ABC, abstractmethod -from typing import TYPE_CHECKING import async_lru @@ -15,9 +14,6 @@ UserProposals, ) -if TYPE_CHECKING: - from . import clients - logger = get_logger() @@ -30,16 +26,6 @@ class MyMdCPort(ABC): added for main metadata module. """ - @classmethod - def from_global(cls) -> "clients.MyMdCClient": - """Create a MyMdCPort from the global client.""" - from damnit_api import _mymdc - - if _mymdc.CLIENT is None: - msg = "MyMdC client has not been initialized. Call bootstrap() first." - raise RuntimeError(msg) - return _mymdc.CLIENT - @abstractmethod async def _get_proposal_by_number(self, no: ProposalNumber) -> dict: ... diff --git a/api/src/damnit_api/main.py b/api/src/damnit_api/main.py index 75a8f015..9017b3d2 100644 --- a/api/src/damnit_api/main.py +++ b/api/src/damnit_api/main.py @@ -1,4 +1,3 @@ -from asyncio import TaskGroup from contextlib import asynccontextmanager from fastapi import FastAPI, HTTPException, Request, status @@ -13,7 +12,7 @@ def create_app(): - from . import _logging, _mymdc, auth, contextfile, get_logger, metadata + from . import _logging, auth, contextfile, get_logger, metadata from .shared import errors, gql from .shared.settings import settings from .state import ( @@ -35,10 +34,7 @@ async def lifespan(app: FastAPI): logger.info("Starting application lifespan") - bootstraps = [_mymdc.bootstrap, auth.bootstrap] - async with TaskGroup() as tg: - for bs in bootstraps: - tg.create_task(bs(settings)) + await auth.bootstrap(settings) db_engine = create_db_engine(settings) app.state.app_state = AppState( From bc7fe968ffc7eae7748c4aa08c2ee73c444110d9 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 16:00:31 +0200 Subject: [PATCH 07/14] refactor(api/auth): build OAuth client via factory, drop globals --- api/src/damnit_api/auth/__init__.py | 13 +---- api/src/damnit_api/auth/bootstrap.py | 70 ------------------------- api/src/damnit_api/auth/dependencies.py | 17 +++++- api/src/damnit_api/main.py | 6 ++- api/src/damnit_api/state.py | 20 +++++++ api/tests/refactor/e2e/conftest.py | 26 --------- api/tests/test_contextfile.py | 7 ++- api/tests/test_errors.py | 7 ++- 8 files changed, 47 insertions(+), 119 deletions(-) delete mode 100644 api/src/damnit_api/auth/bootstrap.py diff --git a/api/src/damnit_api/auth/__init__.py b/api/src/damnit_api/auth/__init__.py index 7f035770..caa65331 100644 --- a/api/src/damnit_api/auth/__init__.py +++ b/api/src/damnit_api/auth/__init__.py @@ -1,14 +1,3 @@ -from authlib.integrations.starlette_client import ( # type: ignore[import-untyped] - StarletteOAuth2App, -) - -from .bootstrap import bootstrap as bootstrap from .routers import noauth_router, router -global __CLIENT - -__CLIENT: StarletteOAuth2App = None # type: ignore[assignment] -"""Global/singleton OAuth client instance.""" - - -__all__ = ["bootstrap", "noauth_router", "router"] +__all__ = ["noauth_router", "router"] diff --git a/api/src/damnit_api/auth/bootstrap.py b/api/src/damnit_api/auth/bootstrap.py deleted file mode 100644 index 4e399228..00000000 --- a/api/src/damnit_api/auth/bootstrap.py +++ /dev/null @@ -1,70 +0,0 @@ -"""Bootstrapping for the auth module.""" - -from typing import TYPE_CHECKING - -from authlib.integrations.starlette_client import OAuth, StarletteOAuth2App - -from .. import get_logger - -logger = get_logger() - -if TYPE_CHECKING: - from ..shared.settings import AuthSettings, Settings - - -global __OAUTH - -__OAUTH: OAuth = None # type: ignore[assignment] - - -async def bootstrap(settings: "Settings"): - """Bootstrap auth module - registers client defined in settings to [`.__OAUTH`] as - `damnit_web` and sets [`damnit_api.auth.__CLIENT`] to [`.__OAUTH.damnit_web`].""" - if settings.auth is None: - await logger.awarning("Auth is disabled, skipping OAuth bootstrap") - return - - import damnit_api.auth - - global __OAUTH - __OAUTH = OAuth() - - if damnit_api.auth.__CLIENT is None: - await logger.ainfo("Configuring OAuth client") - _register(settings.auth) - damnit_api.auth.__CLIENT = __OAUTH.damnit_web # type: ignore[assignment, no-redef] - await damnit_api.auth.__CLIENT.load_server_metadata() # pyright: ignore[reportOptionalMemberAccess] - else: - await logger.awarning("OAuth client already configured") - - -def _register(auth: "AuthSettings"): - """Register the OAuth client defined in settings to [`.__OAUTH`] as `damnit_web`.""" - global __OAUTH - __OAUTH.register( - name="damnit_web", - client_id=auth.client_id, - client_secret=auth.client_secret.get_secret_value(), - server_metadata_url=str(auth.server_metadata_url), - client_kwargs={"scope": "openid email groups"}, - ) - - -def get_oauth_client() -> StarletteOAuth2App: - """Get the global OAuth client instance - for use with fastapi dependencies. - - Returns: - The global OAuth client. - - Raises: - RuntimeError: If the OAuth client has not been initialized. - """ - from damnit_api import auth - - if auth.__CLIENT is None: - msg = ( - "OAuth client has not been initialized. Call " - "[`damnit_api.auth.bootstrap.bootstrap()`] first." - ) - raise RuntimeError(msg) - return auth.__CLIENT diff --git a/api/src/damnit_api/auth/dependencies.py b/api/src/damnit_api/auth/dependencies.py index f64b2de8..20c72248 100644 --- a/api/src/damnit_api/auth/dependencies.py +++ b/api/src/damnit_api/auth/dependencies.py @@ -3,9 +3,9 @@ from typing import Annotated from authlib.integrations.starlette_client import StarletteOAuth2App -from fastapi import Depends +from fastapi import Depends, Request -from .bootstrap import get_oauth_client +from ..state import get_app_state from .models import OAuthUserInfo as _OAuthUserInfo from .models import User as _User @@ -15,6 +15,19 @@ def _get_default_redirect_login_uri() -> str: return "/app/home" +def get_oauth_client(request: Request) -> StarletteOAuth2App: + """Provide the OAuth client from the application state. + + Raises: + RuntimeError: If auth is disabled and no client was built. + """ + client = get_app_state(request).oauth_client + if client is None: + msg = "OAuth client is not configured (auth is disabled)." + raise RuntimeError(msg) + return client + + RedirectURI = Annotated[str, Depends(_get_default_redirect_login_uri)] """Type alias for the redirect URI dependency.""" diff --git a/api/src/damnit_api/main.py b/api/src/damnit_api/main.py index 9017b3d2..d98d47ff 100644 --- a/api/src/damnit_api/main.py +++ b/api/src/damnit_api/main.py @@ -20,6 +20,7 @@ def create_app(): create_db_engine, create_db_sessionmaker, create_mymdc_client, + create_oauth_client, ) logger = get_logger("lifespan") @@ -34,13 +35,16 @@ async def lifespan(app: FastAPI): logger.info("Starting application lifespan") - await auth.bootstrap(settings) + oauth_client = create_oauth_client(settings) + if oauth_client is not None: + await oauth_client.load_server_metadata() db_engine = create_db_engine(settings) app.state.app_state = AppState( db_engine=db_engine, db_sessionmaker=create_db_sessionmaker(db_engine), mymdc_client=create_mymdc_client(settings), + oauth_client=oauth_client, ) if settings.is_local: diff --git a/api/src/damnit_api/state.py b/api/src/damnit_api/state.py index d998277c..2150b1a9 100644 --- a/api/src/damnit_api/state.py +++ b/api/src/damnit_api/state.py @@ -18,6 +18,8 @@ from sqlmodel.ext.asyncio.session import AsyncSession if TYPE_CHECKING: + from authlib.integrations.starlette_client import StarletteOAuth2App + from ._mymdc.clients import MyMdCClient from .shared.settings import Settings @@ -27,6 +29,7 @@ class AppState: db_engine: AsyncEngine db_sessionmaker: async_sessionmaker[AsyncSession] mymdc_client: MyMdCClient + oauth_client: StarletteOAuth2App | None # None when auth is disabled def create_db_engine(settings: Settings) -> AsyncEngine: @@ -53,6 +56,23 @@ def create_mymdc_client(settings: Settings) -> MyMdCClient: raise ValueError(msg) +def create_oauth_client(settings: Settings) -> StarletteOAuth2App | None: + if settings.auth is None: + return None + + from authlib.integrations.starlette_client import OAuth + + oauth = OAuth() + oauth.register( + name="damnit_web", + client_id=settings.auth.client_id, + client_secret=settings.auth.client_secret.get_secret_value(), + server_metadata_url=str(settings.auth.server_metadata_url), + client_kwargs={"scope": "openid email groups"}, + ) + return oauth.damnit_web # pyright: ignore[reportReturnType] + + def get_app_state(request: Request) -> AppState: """FastAPI dependency: the application's :class:`AppState`.""" return request.app.state.app_state diff --git a/api/tests/refactor/e2e/conftest.py b/api/tests/refactor/e2e/conftest.py index 7b52ab23..cc4cc64a 100644 --- a/api/tests/refactor/e2e/conftest.py +++ b/api/tests/refactor/e2e/conftest.py @@ -53,32 +53,6 @@ def vcr_config(): } -@pytest.fixture(autouse=True) -def _reset_bootstrap_globals(): - """Reset the clients/engines cached in module-level globals between tests. - - !!! warning - - Without the auth reset, only the first test's startup performs the OIDC - discovery fetch and later cassettes cannot replay standalone. Without the db - reset, later tests keep the first test's engine and ignore their own - `settings.db_path`. - """ - import damnit_api._db - import damnit_api._mymdc - import damnit_api.auth - - def _reset(): - damnit_api.auth.__CLIENT = None - damnit_api._mymdc.CLIENT = None - damnit_api._db.__ENGINE = None - damnit_api._db.__SESSION_LOCAL = None - - _reset() - yield - _reset() - - DATA_ROOT = Path(__file__).parents[2] / "mock" / "data" / "gpfs" / "exfel" / "exp" MYMDC_CASSETTE = Path(__file__).parents[2] / "mock" / "mymdc" / "mymdc.yaml" diff --git a/api/tests/test_contextfile.py b/api/tests/test_contextfile.py index 7782f161..fe5fa51d 100644 --- a/api/tests/test_contextfile.py +++ b/api/tests/test_contextfile.py @@ -19,10 +19,9 @@ def app(monkeypatch): ) monkeypatch.setenv("DW_API_SESSION_SECRET", "test") - async def noop_bootstrap(settings): - pass - - monkeypatch.setattr("damnit_api.auth.bootstrap", noop_bootstrap) + # No OAuth client: skips the OIDC metadata fetch at startup; these tests + # never exercise the oauth routes. + monkeypatch.setattr("damnit_api.state.create_oauth_client", lambda settings: None) app = create_app() yield app diff --git a/api/tests/test_errors.py b/api/tests/test_errors.py index c3ae8c00..05b0e908 100644 --- a/api/tests/test_errors.py +++ b/api/tests/test_errors.py @@ -87,10 +87,9 @@ def app(monkeypatch): ) monkeypatch.setenv("DW_API_SESSION_SECRET", "test") - async def noop_bootstrap(settings): - pass - - monkeypatch.setattr("damnit_api.auth.bootstrap", noop_bootstrap) + # No OAuth client: skips the OIDC metadata fetch at startup; these tests + # never exercise the oauth routes. + monkeypatch.setattr("damnit_api.state.create_oauth_client", lambda settings: None) app = create_app() yield app From 84bbc5699ddca2216616432fa94d3f8457d202ed Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 16:02:02 +0200 Subject: [PATCH 08/14] refactor(api/auth): replace TOKEN_STORE dict with TokenStore --- api/src/damnit_api/auth/dependencies.py | 9 +++++++++ api/src/damnit_api/auth/routers.py | 10 +++++----- api/src/damnit_api/auth/token_store.py | 19 +++++++++++++++++++ api/src/damnit_api/main.py | 2 ++ api/src/damnit_api/state.py | 8 ++++++++ 5 files changed, 43 insertions(+), 5 deletions(-) create mode 100644 api/src/damnit_api/auth/token_store.py diff --git a/api/src/damnit_api/auth/dependencies.py b/api/src/damnit_api/auth/dependencies.py index 20c72248..0ab22862 100644 --- a/api/src/damnit_api/auth/dependencies.py +++ b/api/src/damnit_api/auth/dependencies.py @@ -8,6 +8,7 @@ from ..state import get_app_state from .models import OAuthUserInfo as _OAuthUserInfo from .models import User as _User +from .token_store import TokenStore # TODO: Get from settings @@ -28,9 +29,17 @@ def get_oauth_client(request: Request) -> StarletteOAuth2App: return client +def get_token_store(request: Request) -> TokenStore: + """Provide the token store from the application state.""" + return get_app_state(request).token_store + + RedirectURI = Annotated[str, Depends(_get_default_redirect_login_uri)] """Type alias for the redirect URI dependency.""" +TokenStoreDep = Annotated[TokenStore, Depends(get_token_store)] +"""Type alias for the token store dependency.""" + Client = Annotated[StarletteOAuth2App, Depends(get_oauth_client)] """Type alias for the OAuth client dependency.""" diff --git a/api/src/damnit_api/auth/routers.py b/api/src/damnit_api/auth/routers.py index b3cd7d52..40044c73 100644 --- a/api/src/damnit_api/auth/routers.py +++ b/api/src/damnit_api/auth/routers.py @@ -13,8 +13,6 @@ router = APIRouter(prefix="/oauth", tags=["auth"]) -TOKEN_STORE = {} - @router.get("/login", status_code=307) async def auth( @@ -47,6 +45,7 @@ async def callback( request: Request, redirect_uri: dependencies.RedirectURI, client: dependencies.Client, + token_store: dependencies.TokenStoreDep, ) -> RedirectResponse: """OAuth2 callback endpoint to handle the response from the OAuth provider.""" try: @@ -63,7 +62,7 @@ async def callback( # TODO: could (should?) be stored in db for persistence across server restarts # NOTE: required for revoking tokens on logout - TOKEN_STORE[user["sub"]] = token + token_store.store(str(user["sub"]), token) return RedirectResponse(url=unquote(str(redirect_uri))) @@ -72,6 +71,7 @@ async def callback( async def logout( request: Request, client: dependencies.Client, + token_store: dependencies.TokenStoreDep, ) -> JSONResponse: """Fully logout the user and revoke tokens. @@ -87,7 +87,7 @@ async def logout( token = await client.fetch_access_token() for k in ("refresh_token", "access_token"): - if user_token := TOKEN_STORE.get(user_sub, {}).pop(k, None): + if user_token := token_store.pop_token_field(user_sub, k): await client.post( revocation_endpoint, token=token, @@ -97,7 +97,7 @@ async def logout( end_session_endpoint = client.server_metadata.get("end_session_endpoint") - token_id = TOKEN_STORE.get(user_sub, {}).pop("id_token", None) + token_id = token_store.pop_token_field(user_sub, "id_token") logout_url = None if token_id and end_session_endpoint: diff --git a/api/src/damnit_api/auth/token_store.py b/api/src/damnit_api/auth/token_store.py new file mode 100644 index 00000000..cccab754 --- /dev/null +++ b/api/src/damnit_api/auth/token_store.py @@ -0,0 +1,19 @@ +"""TokenStore protocol and implementations.""" + +from typing import Any, Protocol + + +class TokenStore(Protocol): + def store(self, sub: str, token: dict[str, Any]) -> None: ... + def pop_token_field(self, sub: str, field: str) -> Any: ... + + +class InMemoryTokenStore: + def __init__(self) -> None: + self._tokens: dict[str, dict[str, Any]] = {} + + def store(self, sub: str, token: dict[str, Any]) -> None: + self._tokens[sub] = token + + def pop_token_field(self, sub: str, field: str) -> Any: + return self._tokens.get(sub, {}).pop(field, None) diff --git a/api/src/damnit_api/main.py b/api/src/damnit_api/main.py index d98d47ff..0ee513b0 100644 --- a/api/src/damnit_api/main.py +++ b/api/src/damnit_api/main.py @@ -21,6 +21,7 @@ def create_app(): create_db_sessionmaker, create_mymdc_client, create_oauth_client, + create_token_store, ) logger = get_logger("lifespan") @@ -45,6 +46,7 @@ async def lifespan(app: FastAPI): db_sessionmaker=create_db_sessionmaker(db_engine), mymdc_client=create_mymdc_client(settings), oauth_client=oauth_client, + token_store=create_token_store(), ) if settings.is_local: diff --git a/api/src/damnit_api/state.py b/api/src/damnit_api/state.py index 2150b1a9..b16cba43 100644 --- a/api/src/damnit_api/state.py +++ b/api/src/damnit_api/state.py @@ -21,6 +21,7 @@ from authlib.integrations.starlette_client import StarletteOAuth2App from ._mymdc.clients import MyMdCClient + from .auth.token_store import TokenStore from .shared.settings import Settings @@ -30,6 +31,7 @@ class AppState: db_sessionmaker: async_sessionmaker[AsyncSession] mymdc_client: MyMdCClient oauth_client: StarletteOAuth2App | None # None when auth is disabled + token_store: TokenStore def create_db_engine(settings: Settings) -> AsyncEngine: @@ -73,6 +75,12 @@ def create_oauth_client(settings: Settings) -> StarletteOAuth2App | None: return oauth.damnit_web # pyright: ignore[reportReturnType] +def create_token_store() -> TokenStore: + from .auth.token_store import InMemoryTokenStore + + return InMemoryTokenStore() + + def get_app_state(request: Request) -> AppState: """FastAPI dependency: the application's :class:`AppState`.""" return request.app.state.app_state From 20f7072f2fd9b938302c77891cd5c99caf3cb998 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 16:08:13 +0200 Subject: [PATCH 09/14] refactor(api/db): replace Registry metaclass with DamnitDBRegistry --- api/src/damnit_api/runs/sqlite/session.py | 38 ++++++++++++++++++----- api/src/damnit_api/utils.py | 36 +-------------------- api/tests/conftest.py | 9 +++--- api/tests/graphql/test_queries.py | 6 ++-- 4 files changed, 40 insertions(+), 49 deletions(-) diff --git a/api/src/damnit_api/runs/sqlite/session.py b/api/src/damnit_api/runs/sqlite/session.py index a75d276a..16815f0d 100644 --- a/api/src/damnit_api/runs/sqlite/session.py +++ b/api/src/damnit_api/runs/sqlite/session.py @@ -1,5 +1,5 @@ from collections.abc import AsyncIterator -from contextlib import asynccontextmanager +from contextlib import AbstractAsyncContextManager, asynccontextmanager from pathlib import Path from sqlalchemy.ext.asyncio import ( @@ -11,7 +11,7 @@ from sqlalchemy.pool import NullPool from ...shared.const import DEFAULT_PROPOSAL -from ...utils import Registry, find_proposal +from ...utils import find_proposal DAMNIT_PATH = "usr/Shared/amore/" @@ -20,7 +20,7 @@ # Asynchronous -class DatabaseSessionManager(metaclass=Registry): +class DatabaseSessionManager: def __init__(self, proposal: str = DEFAULT_PROPOSAL): self.proposal = proposal self.root_path = get_damnit_path(proposal) @@ -73,12 +73,36 @@ async def session(self) -> AsyncIterator[AsyncSession]: await session.close() -def get_session(proposal) -> AsyncSession: - return DatabaseSessionManager(proposal).session() # FIX: # pyright: ignore[reportReturnType] +class DamnitDBRegistry: + """Per-proposal DAMNIT database registry.""" + def __init__(self) -> None: + self._managers: dict[str, DatabaseSessionManager] = {} -def get_connection(proposal) -> AsyncConnection: - return DatabaseSessionManager(proposal).connect() # FIX: # pyright: ignore[reportReturnType] + def get(self, proposal: str) -> DatabaseSessionManager: + if proposal not in self._managers: + self._managers[proposal] = DatabaseSessionManager(proposal) + return self._managers[proposal] + + def pop( + self, proposal: str, default: DatabaseSessionManager | None = None + ) -> DatabaseSessionManager | None: + return self._managers.pop(proposal, default) + + def clear(self) -> None: + self._managers.clear() + + +# TODO: remove, replace with app state (next commit) +damnit_registry = DamnitDBRegistry() + + +def get_session(proposal: str) -> AbstractAsyncContextManager[AsyncSession]: + return damnit_registry.get(proposal).session() + + +def get_connection(proposal: str) -> AbstractAsyncContextManager[AsyncConnection]: + return damnit_registry.get(proposal).connect() # ----------------------------------------------------------------------------- diff --git a/api/src/damnit_api/utils.py b/api/src/damnit_api/utils.py index ce0e2b55..39bf0235 100644 --- a/api/src/damnit_api/utils.py +++ b/api/src/damnit_api/utils.py @@ -1,10 +1,9 @@ import io import os.path as osp -from abc import ABCMeta from base64 import b64encode from glob import iglob from types import UnionType -from typing import Any, ClassVar, Union, get_args, get_origin +from typing import Union, get_args, get_origin import numpy as np @@ -88,39 +87,6 @@ def find_proposal(propno): return "" -# ----------------------------------------------------------------------------- -# Metaclasses - - -class Singleton(ABCMeta): - _instances: ClassVar[dict[type, Any]] = {} - - def __call__(cls, *args, **kwargs): - instance = cls._instances.get(cls) - if not instance: - instance = super(type(cls), cls).__call__(*args, **kwargs) - cls._instances[cls] = instance - return instance - - -class Registry(ABCMeta): - def __call__(cls, proposal, *args, **kwargs): - instance = cls.registry.get( # FIX: # pyright: ignore[reportAttributeAccessIssue] - proposal - ) - if instance is None: - instance = super().__call__(proposal, *args, **kwargs) - cls.registry[ # FIX: # pyright: ignore[reportAttributeAccessIssue] - proposal - ] = instance - return instance - - def __new__(cls, name, bases, attrs): - new_class = super().__new__(cls, name, bases, attrs) - new_class.registry = {} # FIX: # pyright: ignore[reportAttributeAccessIssue] - return new_class - - # ----------------------------------------------------------------------------- # Etc. diff --git a/api/tests/conftest.py b/api/tests/conftest.py index 3f130b3d..c6b4e465 100644 --- a/api/tests/conftest.py +++ b/api/tests/conftest.py @@ -1,12 +1,13 @@ import pytest -from damnit_api.runs.sqlite import DatabaseSessionManager, async_table +from damnit_api.runs.sqlite import async_table +from damnit_api.runs.sqlite.session import damnit_registry @pytest.fixture(autouse=True) -def _clear_db_session_manager_registry(): - DatabaseSessionManager.registry.clear() +def _clear_damnit_registry(): + damnit_registry.clear() async_table.cache_clear() yield - DatabaseSessionManager.registry.clear() + damnit_registry.clear() async_table.cache_clear() diff --git a/api/tests/graphql/test_queries.py b/api/tests/graphql/test_queries.py index 6042471f..e2795da4 100644 --- a/api/tests/graphql/test_queries.py +++ b/api/tests/graphql/test_queries.py @@ -3,6 +3,7 @@ from sqlalchemy import text from damnit_api.runs.sqlite import DAMNIT_PATH, DatabaseSessionManager +from damnit_api.runs.sqlite.session import damnit_registry from damnit_api.runs.types import DamnitRun from .const import ( @@ -131,8 +132,7 @@ async def real_damnit_db(mocker, tmp_path): "damnit_api.runs.sqlite.session.find_proposal", return_value=str(proposal_root), ) - # `registry` is injected by the `Registry` metaclass at class creation. - DatabaseSessionManager.registry.pop(proposal, None) # pyright: ignore[reportAttributeAccessIssue] + damnit_registry.pop(proposal, None) manager = DatabaseSessionManager(proposal) async with manager.connect() as conn: @@ -182,7 +182,7 @@ async def real_damnit_db(mocker, tmp_path): yield proposal await manager.close() - DatabaseSessionManager.registry.pop(proposal, None) # pyright: ignore[reportAttributeAccessIssue] + damnit_registry.pop(proposal, None) @pytest.mark.asyncio From b62913856d5fdd212308a9d9a6045d8e2035b06a Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 16:16:58 +0200 Subject: [PATCH 10/14] refactor(api): route DAMNIT registry through AppState --- api/src/damnit_api/auth/gql.py | 4 +- api/src/damnit_api/auth/routers.py | 7 ++- api/src/damnit_api/graphql/metadata.py | 12 ++--- api/src/damnit_api/graphql/queries.py | 13 +++-- api/src/damnit_api/graphql/subscriptions.py | 21 +++++--- api/src/damnit_api/graphql/utils.py | 8 +-- api/src/damnit_api/main.py | 2 + api/src/damnit_api/metadata/services.py | 7 +-- api/src/damnit_api/runs/sqlite/__init__.py | 2 + .../damnit_api/runs/sqlite/dependencies.py | 17 ++++++ api/src/damnit_api/runs/sqlite/repository.py | 53 ++++++++----------- api/src/damnit_api/runs/sqlite/session.py | 36 ++++++++++--- api/src/damnit_api/shared/gql.py | 14 ++++- api/src/damnit_api/state.py | 8 +++ api/tests/conftest.py | 14 ++--- api/tests/graphql/conftest.py | 28 +++++++++- api/tests/graphql/test_queries.py | 3 -- api/tests/refactor/conftest.py | 1 + api/tests/test_db.py | 6 +-- 19 files changed, 171 insertions(+), 85 deletions(-) create mode 100644 api/src/damnit_api/runs/sqlite/dependencies.py diff --git a/api/src/damnit_api/auth/gql.py b/api/src/damnit_api/auth/gql.py index f1be50e3..2f8b3bde 100644 --- a/api/src/damnit_api/auth/gql.py +++ b/api/src/damnit_api/auth/gql.py @@ -34,7 +34,9 @@ async def proposals( if settings.is_local: from ..metadata.services import _local_proposal_meta, _local_proposal_number - proposal_number = await _local_proposal_number() + proposal_number = await _local_proposal_number( + info.context.damnit_registry + ) if proposal_number is None: return [] return [ProposalMeta.from_pydantic(_local_proposal_meta(proposal_number))] diff --git a/api/src/damnit_api/auth/routers.py b/api/src/damnit_api/auth/routers.py index 40044c73..81ee43d4 100644 --- a/api/src/damnit_api/auth/routers.py +++ b/api/src/damnit_api/auth/routers.py @@ -142,11 +142,14 @@ async def userinfo( @noauth_router.get("/userinfo") -async def noauth_userinfo(): +async def noauth_userinfo(request: Request): from ..metadata.services import LOCAL_CYCLE, _local_proposal_number + from ..state import get_app_state proposals = {} - proposal_number = await _local_proposal_number() + proposal_number = await _local_proposal_number( + get_app_state(request).damnit_registry + ) if proposal_number: proposals = {LOCAL_CYCLE: [proposal_number]} diff --git a/api/src/damnit_api/graphql/metadata.py b/api/src/damnit_api/graphql/metadata.py index 0cad4039..d83278db 100644 --- a/api/src/damnit_api/graphql/metadata.py +++ b/api/src/damnit_api/graphql/metadata.py @@ -8,7 +8,7 @@ @alru_cache(ttl=10) -async def fetch_metadata(proposal=db.DEFAULT_PROPOSAL): +async def fetch_metadata(registry: db.DamnitDBRegistry, proposal=db.DEFAULT_PROPOSAL): """Fetch the per-proposal metadata snapshot from SQLite. Returns a dict with `runs`, `variables`, `tags`, and `timestamp`. Result @@ -16,11 +16,11 @@ async def fetch_metadata(proposal=db.DEFAULT_PROPOSAL): it observes new data so subsequent reads stay fresh. """ tags, variables, variable_tags, runs, max_timestamp = await asyncio.gather( - db.async_all_tags(proposal), - db.async_variables(proposal), - db.async_variable_tags(proposal), - db.async_column(proposal, table="run_info", name="run"), - db.async_max(proposal, table="run_variables", column="timestamp"), + db.async_all_tags(registry, proposal), + db.async_variables(registry, proposal), + db.async_variable_tags(registry, proposal), + db.async_column(registry, proposal, table="run_info", name="run"), + db.async_max(registry, proposal, table="run_variables", column="timestamp"), ) for name, var in variables.items(): diff --git a/api/src/damnit_api/graphql/queries.py b/api/src/damnit_api/graphql/queries.py index b6b68ebf..faa6aad7 100644 --- a/api/src/damnit_api/graphql/queries.py +++ b/api/src/damnit_api/graphql/queries.py @@ -66,8 +66,8 @@ def group_by_run(record): return list(grouped.values()) -async def fetch_variables(proposal, *, limit, offset, names=None): - table = await async_table(proposal, name="run_variables") +async def fetch_variables(registry, proposal, *, limit, offset, names=None): + table = await async_table(registry, proposal, name="run_variables") if table is None: return [] @@ -122,7 +122,7 @@ async def fetch_variables(proposal, *, limit, offset, names=None): .order_by(runs_subquery.c.run) ) - async with get_session(proposal) as session: + async with get_session(registry, proposal) as session: result = await session.execute(query) if not result: raise ValueError # TODO: Better error handling @@ -179,6 +179,7 @@ async def runs( names = _selected_variable_names(info) variables = await fetch_variables( + info.context.damnit_registry, proposal, limit=per_page, offset=(page - 1) * per_page, @@ -190,7 +191,9 @@ async def runs( if _wants_run_info(names): info_rows = await fetch_info( - proposal, runs=[v["run"]["value"] for v in variables] + info.context.damnit_registry, + proposal, + runs=[v["run"]["value"] for v in variables], ) else: info_rows = [{} for _ in variables] @@ -214,7 +217,7 @@ async def metadata( await _ensure_damnit_path(info, proposal) - snapshot = await fetch_metadata(proposal) + snapshot = await fetch_metadata(info.context.damnit_registry, proposal) return { **snapshot, "timestamp": snapshot["timestamp"] * 1000, # ms for JS diff --git a/api/src/damnit_api/graphql/subscriptions.py b/api/src/damnit_api/graphql/subscriptions.py index 3e3de997..996a0906 100644 --- a/api/src/damnit_api/graphql/subscriptions.py +++ b/api/src/damnit_api/graphql/subscriptions.py @@ -4,6 +4,7 @@ import strawberry from async_lru import alru_cache from strawberry.scalars import JSON +from strawberry.types import Info from ..auth.permissions import PROPOSAL_PERMISSIONS from ..runs.sqlite import async_latest_rows, async_max, async_table @@ -22,18 +23,19 @@ # Per-client cursor is deliberately omitted from the cache key so that # concurrent subscribers coalesce into a single DB read per tick. @alru_cache(maxsize=32, ttl=POLLING_INTERVAL) -async def poll_proposal(proposal): - table = await async_table(proposal, name="run_variables") +async def poll_proposal(registry, proposal): + table = await async_table(registry, proposal, name="run_variables") if table is None: return None if proposal not in _last_seen_timestamp: max_timestamp = await async_max( - proposal, table="run_variables", column="timestamp" + registry, proposal, table="run_variables", column="timestamp" ) _last_seen_timestamp[proposal] = max_timestamp or 0 rows = await async_latest_rows( + registry, proposal, table=table, by="timestamp", @@ -44,11 +46,13 @@ async def poll_proposal(proposal): latest_data = LatestData.from_list(rows) - latest_runs = await fetch_info(proposal, runs=list(latest_data.runs.keys())) + latest_runs = await fetch_info( + registry, proposal, runs=list(latest_data.runs.keys()) + ) latest_runs = create_map(latest_runs, key="run") - fetch_metadata.cache_invalidate(proposal) - metadata = await fetch_metadata(proposal) + fetch_metadata.cache_invalidate(registry, proposal) + metadata = await fetch_metadata(registry, proposal) runs = {} run_timestamps = {} @@ -109,13 +113,16 @@ class Subscription: @strawberry.subscription(permission_classes=PROPOSAL_PERMISSIONS) async def latest_data( self, + info: Info, database: DatabaseInput, timestamp: Timestamp, ) -> AsyncGenerator[JSON]: # FIX: # pyright: ignore[reportInvalidTypeForm] while True: await asyncio.sleep(POLLING_INTERVAL) - snapshot = await poll_proposal(proposal=database.proposal) + snapshot = await poll_proposal( + info.context.damnit_registry, proposal=database.proposal + ) result = filter_for_client(snapshot, timestamp) if result is not None: yield result # FIX: # pyright: ignore[reportReturnType] diff --git a/api/src/damnit_api/graphql/utils.py b/api/src/damnit_api/graphql/utils.py index aa390941..6f3738bc 100644 --- a/api/src/damnit_api/graphql/utils.py +++ b/api/src/damnit_api/graphql/utils.py @@ -5,7 +5,7 @@ import strawberry from sqlalchemy import or_, select -from ..runs.sqlite import async_table, get_session +from ..runs.sqlite import DamnitDBRegistry, async_table, get_session from ..shared.const import DEFAULT_PROPOSAL @@ -64,14 +64,14 @@ def from_list(cls, sequence): return instance -async def fetch_info(proposal, *, runs): - table = await async_table(proposal, name="run_info") +async def fetch_info(registry: DamnitDBRegistry, proposal, *, runs): + table = await async_table(registry, proposal, name="run_info") if table is None: return [] conditions = [table.c.run == run for run in runs] query = select(table).where(or_(*conditions)).order_by(table.c.run) - async with get_session(proposal) as session: + async with get_session(registry, proposal) as session: result = await session.execute(query) if not result: raise ValueError # TODO: Better error handling diff --git a/api/src/damnit_api/main.py b/api/src/damnit_api/main.py index 0ee513b0..08bc8405 100644 --- a/api/src/damnit_api/main.py +++ b/api/src/damnit_api/main.py @@ -17,6 +17,7 @@ def create_app(): from .shared.settings import settings from .state import ( AppState, + create_damnit_registry, create_db_engine, create_db_sessionmaker, create_mymdc_client, @@ -47,6 +48,7 @@ async def lifespan(app: FastAPI): mymdc_client=create_mymdc_client(settings), oauth_client=oauth_client, token_store=create_token_store(), + damnit_registry=create_damnit_registry(), ) if settings.is_local: diff --git a/api/src/damnit_api/metadata/services.py b/api/src/damnit_api/metadata/services.py index 525c5d7a..3610886d 100644 --- a/api/src/damnit_api/metadata/services.py +++ b/api/src/damnit_api/metadata/services.py @@ -19,6 +19,7 @@ from .._db.dependencies import DBSession from .._mymdc.clients import MyMdCClient from ..auth.dependencies import User + from ..runs.sqlite import DamnitDBRegistry LOCAL_CYCLE = "197001" @@ -42,17 +43,17 @@ def _local_proposal_meta(proposal_number: ProposalNumber) -> ProposalMeta: ) -async def _local_proposal_number() -> int | None: +async def _local_proposal_number(registry: "DamnitDBRegistry") -> int | None: from sqlalchemy import select from ..runs.sqlite import async_table, get_session from ..shared.const import DEFAULT_PROPOSAL - table = await async_table(DEFAULT_PROPOSAL, name="metameta") + table = await async_table(registry, DEFAULT_PROPOSAL, name="metameta") if table is None: return None - async with get_session(DEFAULT_PROPOSAL) as session: + async with get_session(registry, DEFAULT_PROPOSAL) as session: result = await session.execute( select(table.c.value).where(table.c.key == "proposal") ) diff --git a/api/src/damnit_api/runs/sqlite/__init__.py b/api/src/damnit_api/runs/sqlite/__init__.py index 48729c29..97cc9b13 100644 --- a/api/src/damnit_api/runs/sqlite/__init__.py +++ b/api/src/damnit_api/runs/sqlite/__init__.py @@ -10,6 +10,7 @@ ) from .session import ( DAMNIT_PATH, + DamnitDBRegistry, DatabaseSessionManager, get_connection, get_damnit_path, @@ -19,6 +20,7 @@ __all__ = [ "DAMNIT_PATH", "DEFAULT_PROPOSAL", + "DamnitDBRegistry", "DatabaseSessionManager", "async_all_tags", "async_column", diff --git a/api/src/damnit_api/runs/sqlite/dependencies.py b/api/src/damnit_api/runs/sqlite/dependencies.py new file mode 100644 index 00000000..277e9fd7 --- /dev/null +++ b/api/src/damnit_api/runs/sqlite/dependencies.py @@ -0,0 +1,17 @@ +"""FastAPI dependency helpers for the DAMNIT database registry.""" + +from typing import Annotated + +from fastapi import Depends, Request + +from ...state import get_app_state +from .session import DamnitDBRegistry + + +def get_damnit_registry(request: Request) -> DamnitDBRegistry: + """Provide the DAMNIT database registry from the application state.""" + return get_app_state(request).damnit_registry + + +DamnitRegistry = Annotated[DamnitDBRegistry, Depends(get_damnit_registry)] +"""Type alias for the DAMNIT database registry dependency.""" diff --git a/api/src/damnit_api/runs/sqlite/repository.py b/api/src/damnit_api/runs/sqlite/repository.py index 8116336a..0a9bfe14 100644 --- a/api/src/damnit_api/runs/sqlite/repository.py +++ b/api/src/damnit_api/runs/sqlite/repository.py @@ -1,39 +1,29 @@ from collections import defaultdict from datetime import datetime -from async_lru import alru_cache from sqlalchemy import ( - MetaData, Table, desc, func, select, ) -from sqlalchemy.exc import NoSuchTableError from ...utils import create_map -from .session import get_connection, get_session +from .session import DamnitDBRegistry, get_session -@alru_cache(ttl=300) -async def async_table(proposal, name: str = "runs") -> Table | None: - async with get_connection(proposal) as conn: - try: - return await conn.run_sync( - lambda conn: Table(name, MetaData(), autoload_with=conn) - ) - except NoSuchTableError: - # Don't cache misses; the table may appear shortly. - async_table.cache_invalidate(proposal, name) - return None +async def async_table( + registry: DamnitDBRegistry, proposal, name: str = "runs" +) -> Table | None: + return await registry.get(proposal).get_table(name) -async def async_variables(proposal): - variables = await async_table(proposal, name="variables") +async def async_variables(registry: DamnitDBRegistry, proposal): + variables = await async_table(registry, proposal, name="variables") if variables is None: return {} selection_variables = select(variables.c.name, variables.c.title) - async with get_session(proposal) as session: + async with get_session(registry, proposal) as session: result = await session.execute(selection_variables) variable_rows = result.mappings().all() @@ -42,6 +32,7 @@ async def async_variables(proposal): async def async_latest_rows( + registry: DamnitDBRegistry, proposal, *, table: Table, @@ -55,56 +46,56 @@ async def async_latest_rows( selection = select(table).where(table.c.get(by) > start_at).order_by(order_by) - async with get_session(proposal) as session: + async with get_session(registry, proposal) as session: result = await session.execute(selection) return result.mappings().all() # FIX: # pyright: ignore[reportReturnType] -async def async_max(proposal, *, table: str, column: str): - table = await async_table(proposal, name=table) +async def async_max(registry: DamnitDBRegistry, proposal, *, table: str, column: str): + table = await async_table(registry, proposal, name=table) if table is None: return None selection = select(func.max(table.c.get(column))) - async with get_session(proposal) as session: + async with get_session(registry, proposal) as session: result = await session.execute(selection) return result.scalar() -async def async_column(proposal, *, table: str, name: str): - table = await async_table(proposal, name=table) +async def async_column(registry: DamnitDBRegistry, proposal, *, table: str, name: str): + table = await async_table(registry, proposal, name=table) if table is None: return [] selection = select(table.c.get(name)) - async with get_session(proposal) as session: + async with get_session(registry, proposal) as session: result = await session.execute(selection) return result.scalars().all() -async def async_all_tags(proposal): - tags_table = await async_table(proposal, name="tags") +async def async_all_tags(registry: DamnitDBRegistry, proposal): + tags_table = await async_table(registry, proposal, name="tags") if tags_table is None: return {} selection = select( tags_table.c.id, tags_table.c.name, ) - async with get_session(proposal) as session: + async with get_session(registry, proposal) as session: result = await session.execute(selection) return create_map(result.mappings().all(), key="id") -async def async_variable_tags(proposal): - variable_tags_table = await async_table(proposal, name="variable_tags") +async def async_variable_tags(registry: DamnitDBRegistry, proposal): + variable_tags_table = await async_table(registry, proposal, name="variable_tags") if variable_tags_table is None: return {} selection = select( variable_tags_table.c.variable_name, variable_tags_table.c.tag_id ) - async with get_session(proposal) as session: + async with get_session(registry, proposal) as session: result = await session.execute(selection) variable_tags: dict[str, list[int]] = defaultdict(list) diff --git a/api/src/damnit_api/runs/sqlite/session.py b/api/src/damnit_api/runs/sqlite/session.py index 16815f0d..49502887 100644 --- a/api/src/damnit_api/runs/sqlite/session.py +++ b/api/src/damnit_api/runs/sqlite/session.py @@ -2,6 +2,9 @@ from contextlib import AbstractAsyncContextManager, asynccontextmanager from pathlib import Path +from async_lru import alru_cache +from sqlalchemy import MetaData, Table +from sqlalchemy.exc import NoSuchTableError from sqlalchemy.ext.asyncio import ( AsyncConnection, AsyncSession, @@ -30,6 +33,9 @@ def __init__(self, proposal: str = DEFAULT_PROPOSAL): poolclass=NullPool, ) self._sessionmaker = async_sessionmaker(autocommit=False, bind=self._engine) + # Owned by this manager, not the module: + # one reflection cache per proposal, cleared when the manager is. + self._table_cache = alru_cache(ttl=300)(self._reflect_table) @property def db_path(self): @@ -72,6 +78,20 @@ async def session(self) -> AsyncIterator[AsyncSession]: finally: await session.close() + async def _reflect_table(self, name: str) -> Table | None: + async with self.connect() as conn: + try: + return await conn.run_sync( + lambda conn: Table(name, MetaData(), autoload_with=conn) + ) + except NoSuchTableError: + # Don't cache misses; the table may appear shortly. + self._table_cache.cache_invalidate(name) + return None + + async def get_table(self, name: str = "runs") -> Table | None: + return await self._table_cache(name) + class DamnitDBRegistry: """Per-proposal DAMNIT database registry.""" @@ -93,16 +113,16 @@ def clear(self) -> None: self._managers.clear() -# TODO: remove, replace with app state (next commit) -damnit_registry = DamnitDBRegistry() - - -def get_session(proposal: str) -> AbstractAsyncContextManager[AsyncSession]: - return damnit_registry.get(proposal).session() +def get_session( + registry: DamnitDBRegistry, proposal: str +) -> AbstractAsyncContextManager[AsyncSession]: + return registry.get(proposal).session() -def get_connection(proposal: str) -> AbstractAsyncContextManager[AsyncConnection]: - return damnit_registry.get(proposal).connect() +def get_connection( + registry: DamnitDBRegistry, proposal: str +) -> AbstractAsyncContextManager[AsyncConnection]: + return registry.get(proposal).connect() # ----------------------------------------------------------------------------- diff --git a/api/src/damnit_api/shared/gql.py b/api/src/damnit_api/shared/gql.py index d010b7d6..754d4ab1 100644 --- a/api/src/damnit_api/shared/gql.py +++ b/api/src/damnit_api/shared/gql.py @@ -16,6 +16,7 @@ from ..auth.models import User from ..metadata import gql as metadata from ..runs import types as run_types +from ..runs.sqlite.dependencies import DamnitRegistry SUBSCRIPTION_PROTOCOLS = [ GRAPHQL_TRANSPORT_WS_PROTOCOL, @@ -55,6 +56,7 @@ class Context(BaseContext): mymdc: MyMdCClient oauth_user: OAuthUserInfo session: DBSession + damnit_registry: DamnitRegistry _user: User | None = None async def get_user(self) -> User: @@ -67,9 +69,17 @@ async def get_user(self) -> User: async def get_context( # noqa: RUF029 - oauth_user: OAuthUserInfo, mymdc: MyMdCClient, session: DBSession + oauth_user: OAuthUserInfo, + mymdc: MyMdCClient, + session: DBSession, + damnit_registry: DamnitRegistry, ): - return Context(oauth_user=oauth_user, mymdc=mymdc, session=session) + return Context( + oauth_user=oauth_user, + mymdc=mymdc, + session=session, + damnit_registry=damnit_registry, + ) def get_gql_app(): diff --git a/api/src/damnit_api/state.py b/api/src/damnit_api/state.py index b16cba43..2e63bc53 100644 --- a/api/src/damnit_api/state.py +++ b/api/src/damnit_api/state.py @@ -22,6 +22,7 @@ from ._mymdc.clients import MyMdCClient from .auth.token_store import TokenStore + from .runs.sqlite.session import DamnitDBRegistry from .shared.settings import Settings @@ -32,6 +33,7 @@ class AppState: mymdc_client: MyMdCClient oauth_client: StarletteOAuth2App | None # None when auth is disabled token_store: TokenStore + damnit_registry: DamnitDBRegistry def create_db_engine(settings: Settings) -> AsyncEngine: @@ -81,6 +83,12 @@ def create_token_store() -> TokenStore: return InMemoryTokenStore() +def create_damnit_registry() -> DamnitDBRegistry: + from .runs.sqlite.session import DamnitDBRegistry + + return DamnitDBRegistry() + + def get_app_state(request: Request) -> AppState: """FastAPI dependency: the application's :class:`AppState`.""" return request.app.state.app_state diff --git a/api/tests/conftest.py b/api/tests/conftest.py index c6b4e465..9c374349 100644 --- a/api/tests/conftest.py +++ b/api/tests/conftest.py @@ -1,13 +1,9 @@ import pytest -from damnit_api.runs.sqlite import async_table -from damnit_api.runs.sqlite.session import damnit_registry +from damnit_api.runs.sqlite import DamnitDBRegistry -@pytest.fixture(autouse=True) -def _clear_damnit_registry(): - damnit_registry.clear() - async_table.cache_clear() - yield - damnit_registry.clear() - async_table.cache_clear() +@pytest.fixture +def damnit_registry() -> DamnitDBRegistry: + """A fresh per-test registry; nothing module-level to clear between tests.""" + return DamnitDBRegistry() diff --git a/api/tests/graphql/conftest.py b/api/tests/graphql/conftest.py index ed4a34a8..5fe2734e 100644 --- a/api/tests/graphql/conftest.py +++ b/api/tests/graphql/conftest.py @@ -1,3 +1,5 @@ +from types import SimpleNamespace + import pytest import strawberry from strawberry.schema.config import StrawberryConfig @@ -111,16 +113,40 @@ def graphql_schema_no_auth( ) +class _SchemaWithDefaultContext: + """Wraps a strawberry `Schema` so resolvers see a default context (with + a fresh `damnit_registry`) without every test passing `context_value`; + an explicit `context_value` at the call site still wins.""" + + def __init__(self, schema, context): + self._schema = schema + self._context = context + + def execute(self, query, **kwargs): + kwargs.setdefault("context_value", self._context) + return self._schema.execute(query, **kwargs) + + def subscribe(self, query, **kwargs): + kwargs.setdefault("context_value", self._context) + return self._schema.subscribe(query, **kwargs) + + +@pytest.fixture +def graphql_context(damnit_registry): + return SimpleNamespace(damnit_registry=damnit_registry) + + @pytest.fixture def graphql_schema( bypass_proposal_permission, mocked_ensure_damnit_path, mocked_metadata_max, graphql_schema_no_auth, + graphql_context, ): """Same schema as graphql_schema_no_auth, with permission and damnit-path checks bypassed so tests exercise resolver logic only.""" - return graphql_schema_no_auth + return _SchemaWithDefaultContext(graphql_schema_no_auth, graphql_context) @pytest.fixture diff --git a/api/tests/graphql/test_queries.py b/api/tests/graphql/test_queries.py index e2795da4..c8a827bb 100644 --- a/api/tests/graphql/test_queries.py +++ b/api/tests/graphql/test_queries.py @@ -3,7 +3,6 @@ from sqlalchemy import text from damnit_api.runs.sqlite import DAMNIT_PATH, DatabaseSessionManager -from damnit_api.runs.sqlite.session import damnit_registry from damnit_api.runs.types import DamnitRun from .const import ( @@ -132,7 +131,6 @@ async def real_damnit_db(mocker, tmp_path): "damnit_api.runs.sqlite.session.find_proposal", return_value=str(proposal_root), ) - damnit_registry.pop(proposal, None) manager = DatabaseSessionManager(proposal) async with manager.connect() as conn: @@ -182,7 +180,6 @@ async def real_damnit_db(mocker, tmp_path): yield proposal await manager.close() - damnit_registry.pop(proposal, None) @pytest.mark.asyncio diff --git a/api/tests/refactor/conftest.py b/api/tests/refactor/conftest.py index 850c2924..cdf5dd2f 100644 --- a/api/tests/refactor/conftest.py +++ b/api/tests/refactor/conftest.py @@ -6,6 +6,7 @@ from ..graphql.conftest import ( # noqa: F401 bypass_proposal_permission, + graphql_context, graphql_schema, graphql_schema_no_auth, mocked_ensure_damnit_path, diff --git a/api/tests/test_db.py b/api/tests/test_db.py index 603a03bd..f0b49069 100644 --- a/api/tests/test_db.py +++ b/api/tests/test_db.py @@ -81,12 +81,12 @@ def test_engine_uses_nullpool_and_autocommit(damnit_db): # this test verifies (no file descriptors leak after the loop dies). # alru_cached async_table sees that loop change; warning is intrinsic. @pytest.mark.filterwarnings("ignore::async_lru.AlruCacheLoopResetWarning") -def test_no_lingering_file_descriptor_after_read(damnit_db): +def test_no_lingering_file_descriptor_after_read(damnit_db, damnit_registry): db_file = Path(damnit_db) / DAMNIT_PATH / "runs.sqlite" async def do_read(): - table = await async_table(damnit_db, name="runs") - async with get_session(damnit_db) as session: + table = await async_table(damnit_registry, damnit_db, name="runs") + async with get_session(damnit_registry, damnit_db) as session: await session.execute(table.select()) assert _open_file_descriptors_to(db_file) == [] From 6c68d2686dcbee2c0e176183ffeab5fd6b22a9b1 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 16:20:00 +0200 Subject: [PATCH 11/14] refactor(api/graphql): move subscription cursors into AppState --- api/src/damnit_api/auth/gql.py | 4 +-- api/src/damnit_api/graphql/dependencies.py | 19 +++++++++++ api/src/damnit_api/graphql/subscriptions.py | 37 ++++++++++++++++----- api/src/damnit_api/main.py | 2 ++ api/src/damnit_api/shared/gql.py | 4 +++ api/src/damnit_api/state.py | 8 +++++ api/tests/graphql/conftest.py | 13 +++++--- 7 files changed, 71 insertions(+), 16 deletions(-) create mode 100644 api/src/damnit_api/graphql/dependencies.py diff --git a/api/src/damnit_api/auth/gql.py b/api/src/damnit_api/auth/gql.py index 2f8b3bde..aff3ad38 100644 --- a/api/src/damnit_api/auth/gql.py +++ b/api/src/damnit_api/auth/gql.py @@ -34,9 +34,7 @@ async def proposals( if settings.is_local: from ..metadata.services import _local_proposal_meta, _local_proposal_number - proposal_number = await _local_proposal_number( - info.context.damnit_registry - ) + proposal_number = await _local_proposal_number(info.context.damnit_registry) if proposal_number is None: return [] return [ProposalMeta.from_pydantic(_local_proposal_meta(proposal_number))] diff --git a/api/src/damnit_api/graphql/dependencies.py b/api/src/damnit_api/graphql/dependencies.py new file mode 100644 index 00000000..f6fcc03d --- /dev/null +++ b/api/src/damnit_api/graphql/dependencies.py @@ -0,0 +1,19 @@ +"""FastAPI dependency helpers for GraphQL subscription state.""" + +from typing import Annotated + +from fastapi import Depends, Request + +from ..state import get_app_state +from .subscriptions import SubscriptionCursors + + +def get_subscription_cursors(request: Request) -> SubscriptionCursors: + """Provide the subscription cursors from the application state.""" + return get_app_state(request).subscription_cursors + + +SubscriptionCursorsDep = Annotated[ + SubscriptionCursors, Depends(get_subscription_cursors) +] +"""Type alias for the subscription cursors dependency.""" diff --git a/api/src/damnit_api/graphql/subscriptions.py b/api/src/damnit_api/graphql/subscriptions.py index 996a0906..3c47dd5a 100644 --- a/api/src/damnit_api/graphql/subscriptions.py +++ b/api/src/damnit_api/graphql/subscriptions.py @@ -15,31 +15,48 @@ POLLING_INTERVAL = 1 # seconds -# Server-side high-water mark per proposal so each tick only fetches rows -# newer than what the previous tick already shipped. -_last_seen_timestamp: dict[str, float] = {} + +class SubscriptionCursors: + """Server-side high-water mark per proposal so each tick only fetches rows + newer than what the previous tick already shipped. Hashable by identity + for alru_cache.""" + + def __init__(self) -> None: + self._data: dict[str, float] = {} + + def __contains__(self, proposal: str) -> bool: + return proposal in self._data + + def __getitem__(self, proposal: str) -> float: + return self._data[proposal] + + def __setitem__(self, proposal: str, value: float) -> None: + self._data[proposal] = value + + def clear(self) -> None: + self._data.clear() # Per-client cursor is deliberately omitted from the cache key so that # concurrent subscribers coalesce into a single DB read per tick. @alru_cache(maxsize=32, ttl=POLLING_INTERVAL) -async def poll_proposal(registry, proposal): +async def poll_proposal(registry, proposal, cursors: SubscriptionCursors): table = await async_table(registry, proposal, name="run_variables") if table is None: return None - if proposal not in _last_seen_timestamp: + if proposal not in cursors: max_timestamp = await async_max( registry, proposal, table="run_variables", column="timestamp" ) - _last_seen_timestamp[proposal] = max_timestamp or 0 + cursors[proposal] = max_timestamp or 0 rows = await async_latest_rows( registry, proposal, table=table, by="timestamp", - start_at=_last_seen_timestamp[proposal], + start_at=cursors[proposal], ) if not rows: return None @@ -80,7 +97,7 @@ async def poll_proposal(registry, proposal): msg = "Latest data has no timestamp." raise ValueError(msg) - _last_seen_timestamp[proposal] = latest_data.timestamp + cursors[proposal] = latest_data.timestamp metadata = { "runs": sorted(set(metadata["runs"]) | set(runs.keys())), @@ -121,7 +138,9 @@ async def latest_data( await asyncio.sleep(POLLING_INTERVAL) snapshot = await poll_proposal( - info.context.damnit_registry, proposal=database.proposal + info.context.damnit_registry, + proposal=database.proposal, + cursors=info.context.subscription_cursors, ) result = filter_for_client(snapshot, timestamp) if result is not None: diff --git a/api/src/damnit_api/main.py b/api/src/damnit_api/main.py index 08bc8405..9e1fa233 100644 --- a/api/src/damnit_api/main.py +++ b/api/src/damnit_api/main.py @@ -22,6 +22,7 @@ def create_app(): create_db_sessionmaker, create_mymdc_client, create_oauth_client, + create_subscription_cursors, create_token_store, ) @@ -49,6 +50,7 @@ async def lifespan(app: FastAPI): oauth_client=oauth_client, token_store=create_token_store(), damnit_registry=create_damnit_registry(), + subscription_cursors=create_subscription_cursors(), ) if settings.is_local: diff --git a/api/src/damnit_api/shared/gql.py b/api/src/damnit_api/shared/gql.py index 754d4ab1..1ff184c8 100644 --- a/api/src/damnit_api/shared/gql.py +++ b/api/src/damnit_api/shared/gql.py @@ -14,6 +14,7 @@ from ..auth import gql as auth from ..auth.dependencies import OAuthUserInfo from ..auth.models import User +from ..graphql.dependencies import SubscriptionCursorsDep from ..metadata import gql as metadata from ..runs import types as run_types from ..runs.sqlite.dependencies import DamnitRegistry @@ -57,6 +58,7 @@ class Context(BaseContext): oauth_user: OAuthUserInfo session: DBSession damnit_registry: DamnitRegistry + subscription_cursors: SubscriptionCursorsDep _user: User | None = None async def get_user(self) -> User: @@ -73,12 +75,14 @@ async def get_context( # noqa: RUF029 mymdc: MyMdCClient, session: DBSession, damnit_registry: DamnitRegistry, + subscription_cursors: SubscriptionCursorsDep, ): return Context( oauth_user=oauth_user, mymdc=mymdc, session=session, damnit_registry=damnit_registry, + subscription_cursors=subscription_cursors, ) diff --git a/api/src/damnit_api/state.py b/api/src/damnit_api/state.py index 2e63bc53..92039ceb 100644 --- a/api/src/damnit_api/state.py +++ b/api/src/damnit_api/state.py @@ -22,6 +22,7 @@ from ._mymdc.clients import MyMdCClient from .auth.token_store import TokenStore + from .graphql.subscriptions import SubscriptionCursors from .runs.sqlite.session import DamnitDBRegistry from .shared.settings import Settings @@ -34,6 +35,7 @@ class AppState: oauth_client: StarletteOAuth2App | None # None when auth is disabled token_store: TokenStore damnit_registry: DamnitDBRegistry + subscription_cursors: SubscriptionCursors def create_db_engine(settings: Settings) -> AsyncEngine: @@ -89,6 +91,12 @@ def create_damnit_registry() -> DamnitDBRegistry: return DamnitDBRegistry() +def create_subscription_cursors() -> SubscriptionCursors: + from .graphql.subscriptions import SubscriptionCursors + + return SubscriptionCursors() + + def get_app_state(request: Request) -> AppState: """FastAPI dependency: the application's :class:`AppState`.""" return request.app.state.app_state diff --git a/api/tests/graphql/conftest.py b/api/tests/graphql/conftest.py index 5fe2734e..6380d78f 100644 --- a/api/tests/graphql/conftest.py +++ b/api/tests/graphql/conftest.py @@ -4,11 +4,14 @@ import strawberry from strawberry.schema.config import StrawberryConfig -from damnit_api.graphql import subscriptions from damnit_api.graphql.directives import lightweight from damnit_api.graphql.metadata import fetch_metadata from damnit_api.graphql.queries import Query -from damnit_api.graphql.subscriptions import Subscription, poll_proposal +from damnit_api.graphql.subscriptions import ( + Subscription, + SubscriptionCursors, + poll_proposal, +) from damnit_api.runs.types import SCALAR_MAP, DamnitVariable from .const import ( @@ -23,7 +26,6 @@ def reset_caches(): fetch_metadata.cache_clear() poll_proposal.cache_clear() - subscriptions._last_seen_timestamp.clear() return @@ -133,7 +135,10 @@ def subscribe(self, query, **kwargs): @pytest.fixture def graphql_context(damnit_registry): - return SimpleNamespace(damnit_registry=damnit_registry) + return SimpleNamespace( + damnit_registry=damnit_registry, + subscription_cursors=SubscriptionCursors(), + ) @pytest.fixture From a8c3ff4a7f758760da87828dc88b6aa87d124915 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 16:21:27 +0200 Subject: [PATCH 12/14] test(api/state): cover AppState DI surface, factories, stores --- api/tests/test_state.py | 70 +++++++++++++++++++++++++++++++++++++++++ 1 file changed, 70 insertions(+) create mode 100644 api/tests/test_state.py diff --git a/api/tests/test_state.py b/api/tests/test_state.py new file mode 100644 index 00000000..ff883c30 --- /dev/null +++ b/api/tests/test_state.py @@ -0,0 +1,70 @@ +"""Tests for the AppState container and its DI surface (ADR-002/ADR-003).""" + +import ast +from pathlib import Path + +from damnit_api.auth.token_store import InMemoryTokenStore +from damnit_api.runs.sqlite import DamnitDBRegistry +from damnit_api.shared.settings import Settings +from damnit_api.state import create_oauth_client + + +def test_appstate_only_imported_by_composition_root(): + """Handlers and dependencies must depend on specific objects, never the + whole `AppState`; only the composition root may import it (ADR-002).""" + src_root = Path("src/damnit_api") + allowed = {src_root / "state.py", src_root / "main.py"} + + def imports_app_state(path: Path) -> bool: + tree = ast.parse(path.read_text()) + return any( + isinstance(node, ast.ImportFrom) + and any(alias.name == "AppState" for alias in node.names) + for node in ast.walk(tree) + ) + + offenders = [ + str(path) + for path in src_root.rglob("*.py") + if path not in allowed and imports_app_state(path) + ] + + assert offenders == [] + + +def test_create_oauth_client_returns_none_when_auth_disabled(tmp_path): + settings = Settings(damnit_path=tmp_path) # local mode: auth is None + assert settings.auth is None + assert create_oauth_client(settings) is None + + +def test_registry_memoizes_managers_per_proposal(monkeypatch): + created = [] + + class DummyManager: + def __init__(self, proposal): + self.proposal = proposal + created.append(proposal) + + monkeypatch.setattr( + "damnit_api.runs.sqlite.session.DatabaseSessionManager", DummyManager + ) + + registry = DamnitDBRegistry() + first = registry.get("1234") + assert registry.get("1234") is first + assert registry.get("5678") is not first + assert created == ["1234", "5678"] + + assert registry.pop("1234") is first + assert registry.pop("1234") is None # already removed + assert registry.get("1234") is not first # rebuilt after pop + + +def test_token_store_stores_and_pops_fields(): + store = InMemoryTokenStore() + store.store("sub-1", {"access_token": "a", "id_token": "i"}) + + assert store.pop_token_field("sub-1", "access_token") == "a" + assert store.pop_token_field("sub-1", "access_token") is None # popped + assert store.pop_token_field("unknown-sub", "id_token") is None From a797c9f0e8b56e858ba68f540ed6f57c63ed01d2 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 16:22:18 +0200 Subject: [PATCH 13/14] docs(api): mark state.py landed in the composition-root row --- api/docs/architecture.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/api/docs/architecture.md b/api/docs/architecture.md index 17b4c171..18a87b12 100644 --- a/api/docs/architecture.md +++ b/api/docs/architecture.md @@ -26,7 +26,7 @@ For more information, see [ADR-000](adr/000-vertical-slice-architecture.md). | `appdb/` | The app's own database (infrastructure) | Models, engine/session plumbing for `dw_api.sqlite` | Planned | `_db/` | | `mymdc/` | MyMdC client (infrastructure) | Ports, clients, vendored models | Planned | `_mymdc/` | | `core/` | Cross-cutting, framework-free | Shared error classes (see [ADR-001](adr/001-error-classes.md)), `DamnitType`, value types, converters | Planned | `shared/` + `utils.py` | -| `main.py` / `app.py` / `state.py` | Composition root | `AppState`, `create_*` factories (see [ADR-002](adr/002-no-global-mutable-state.md)), `create_app()` - the only place that may import everything and read settings | Partial | `main.py` only | +| `main.py` / `app.py` / `state.py` | Composition root | `AppState`, `create_*` factories (see [ADR-002](adr/002-no-global-mutable-state.md)), `create_app()` - the only place that may import everything and read settings | Partial | `main.py` + `state.py` | Where new code goes: From 1d8be88d627cbd36f547f785d42e308c199c5829 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Tue, 7 Jul 2026 16:27:27 +0200 Subject: [PATCH 14/14] refactor(api/db): drop dead registry eviction and get_connection --- api/src/damnit_api/graphql/subscriptions.py | 3 --- api/src/damnit_api/runs/sqlite/__init__.py | 2 -- api/src/damnit_api/runs/sqlite/session.py | 14 -------------- api/tests/test_state.py | 4 ---- 4 files changed, 23 deletions(-) diff --git a/api/src/damnit_api/graphql/subscriptions.py b/api/src/damnit_api/graphql/subscriptions.py index 3c47dd5a..d16e4555 100644 --- a/api/src/damnit_api/graphql/subscriptions.py +++ b/api/src/damnit_api/graphql/subscriptions.py @@ -33,9 +33,6 @@ def __getitem__(self, proposal: str) -> float: def __setitem__(self, proposal: str, value: float) -> None: self._data[proposal] = value - def clear(self) -> None: - self._data.clear() - # Per-client cursor is deliberately omitted from the cache key so that # concurrent subscribers coalesce into a single DB read per tick. diff --git a/api/src/damnit_api/runs/sqlite/__init__.py b/api/src/damnit_api/runs/sqlite/__init__.py index 97cc9b13..ecbb9438 100644 --- a/api/src/damnit_api/runs/sqlite/__init__.py +++ b/api/src/damnit_api/runs/sqlite/__init__.py @@ -12,7 +12,6 @@ DAMNIT_PATH, DamnitDBRegistry, DatabaseSessionManager, - get_connection, get_damnit_path, get_session, ) @@ -29,7 +28,6 @@ "async_table", "async_variable_tags", "async_variables", - "get_connection", "get_damnit_path", "get_session", ] diff --git a/api/src/damnit_api/runs/sqlite/session.py b/api/src/damnit_api/runs/sqlite/session.py index 49502887..e690cae3 100644 --- a/api/src/damnit_api/runs/sqlite/session.py +++ b/api/src/damnit_api/runs/sqlite/session.py @@ -104,14 +104,6 @@ def get(self, proposal: str) -> DatabaseSessionManager: self._managers[proposal] = DatabaseSessionManager(proposal) return self._managers[proposal] - def pop( - self, proposal: str, default: DatabaseSessionManager | None = None - ) -> DatabaseSessionManager | None: - return self._managers.pop(proposal, default) - - def clear(self) -> None: - self._managers.clear() - def get_session( registry: DamnitDBRegistry, proposal: str @@ -119,12 +111,6 @@ def get_session( return registry.get(proposal).session() -def get_connection( - registry: DamnitDBRegistry, proposal: str -) -> AbstractAsyncContextManager[AsyncConnection]: - return registry.get(proposal).connect() - - # ----------------------------------------------------------------------------- # Etc. diff --git a/api/tests/test_state.py b/api/tests/test_state.py index ff883c30..8e43dcb8 100644 --- a/api/tests/test_state.py +++ b/api/tests/test_state.py @@ -56,10 +56,6 @@ def __init__(self, proposal): assert registry.get("5678") is not first assert created == ["1234", "5678"] - assert registry.pop("1234") is first - assert registry.pop("1234") is None # already removed - assert registry.get("1234") is not first # rebuilt after pop - def test_token_store_stores_and_pops_fields(): store = InMemoryTokenStore()