From 67c8f634c205b685a513c14bb2fbfd26c5a23114 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 04:58:57 +0200 Subject: [PATCH 01/12] build(api): add litestar alongside fastapi for the swap --- api/pyproject.toml | 2 +- uv.lock | 110 ++++++++++++++++++++++++++++++++++++++++----- 2 files changed, 100 insertions(+), 12 deletions(-) diff --git a/api/pyproject.toml b/api/pyproject.toml index ff1f62df..c8e57501 100644 --- a/api/pyproject.toml +++ b/api/pyproject.toml @@ -17,7 +17,7 @@ dependencies = [ "numpy~=2.3", "matplotlib~=3.7", "h5py~=3.9", - "strawberry-graphql[fastapi]>=0.283.3", + "strawberry-graphql[fastapi,litestar]>=0.283.3", "uvicorn[standard]~=0.29", "aiosqlite~=0.19", "scipy~=1.11", diff --git a/uv.lock b/uv.lock index 1f981486..e30ad501 100644 --- a/uv.lock +++ b/uv.lock @@ -412,7 +412,7 @@ dependencies = [ { name = "scipy" }, { name = "sqlalchemy", extra = ["asyncio"] }, { name = "sqlmodel" }, - { name = "strawberry-graphql", extra = ["fastapi"] }, + { name = "strawberry-graphql", extra = ["fastapi", "litestar"] }, { name = "structlog" }, { name = "uvicorn", extra = ["standard"] }, { name = "xarray" }, @@ -489,7 +489,7 @@ requires-dist = [ { name = "scipy", specifier = "~=1.11" }, { name = "sqlalchemy", extras = ["asyncio"], specifier = "~=2.0" }, { name = "sqlmodel", specifier = ">=0.0.31" }, - { name = "strawberry-graphql", extras = ["fastapi"], specifier = ">=0.283.3" }, + { name = "strawberry-graphql", extras = ["fastapi", "litestar"], specifier = ">=0.283.3" }, { name = "structlog", specifier = "~=24.4" }, { name = "uvicorn", extras = ["standard"], specifier = "~=0.29" }, { name = "xarray", specifier = ">=2024.11" }, @@ -610,7 +610,7 @@ wheels = [ [[package]] name = "fastapi" -version = "0.135.3" +version = "0.139.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "annotated-doc" }, @@ -619,9 +619,9 @@ dependencies = [ { name = "typing-extensions" }, { name = "typing-inspection" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/f7/e6/7adb4c5fa231e82c35b8f5741a9f2d055f520c29af5546fd70d3e8e1cd2e/fastapi-0.135.3.tar.gz", hash = "sha256:bd6d7caf1a2bdd8d676843cdcd2287729572a1ef524fc4d65c17ae002a1be654", size = 396524, upload-time = "2026-04-01T16:23:58.188Z" } +sdist = { url = "https://files.pythonhosted.org/packages/d3/af/a5f50ccfa659ec1802cb4ca842c23f06d906a8cc9aef6016a2caeea3d4ed/fastapi-0.139.0.tar.gz", hash = "sha256:99ab7b2d92223c76d6cf10757ab3f89d45b38267fc20b2a136cf02f6beac3145", size = 423016, upload-time = "2026-07-01T16:35:33.436Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/84/a4/5caa2de7f917a04ada20018eccf60d6cc6145b0199d55ca3711b0fc08312/fastapi-0.135.3-py3-none-any.whl", hash = "sha256:9b0f590c813acd13d0ab43dd8494138eb58e484bfac405db1f3187cfc5810d98", size = 117734, upload-time = "2026-04-01T16:23:59.328Z" }, + { url = "https://files.pythonhosted.org/packages/9e/7c/8e3c6ad324ea5cb36604fc3f968554887891c316d9dfde57761611d907ad/fastapi-0.139.0-py3-none-any.whl", hash = "sha256:cf15e1e9e667ddb0ad63811e60bd11390d1aac838ca4a7a23f421807b2308189", size = 130339, upload-time = "2026-07-01T16:35:32.19Z" }, ] [[package]] @@ -950,6 +950,39 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/4e/f6/71d6ec9f18da0b2201287ce9db6afb1a1f637dedb3f0703409558981c723/ldap3-2.9.1-py2.py3-none-any.whl", hash = "sha256:5869596fc4948797020d3f03b7939da938778a0f9e2009f7a072ccf92b8e8d70", size = 432192, upload-time = "2021-07-18T06:34:12.905Z" }, ] +[[package]] +name = "litestar" +version = "2.24.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "anyio" }, + { name = "click" }, + { name = "httpx" }, + { name = "litestar-htmx" }, + { name = "msgspec" }, + { name = "multidict" }, + { name = "multipart" }, + { name = "polyfactory" }, + { name = "pyyaml" }, + { name = "rich" }, + { name = "rich-click" }, + { name = "sniffio" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/42/f6/33619a1562c6732dab211a967b1e7af3285de29d655c0b37b3f379aa2700/litestar-2.24.0.tar.gz", hash = "sha256:8f4b137cb115554b9fbc12bd01d398d5bb40ba85f2d333dd73f4b747308bbb46", size = 383070, upload-time = "2026-06-11T11:27:00.519Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/c8/ce/4697b547790f23d2fcc37a928d48c38140cf86360e4041ca8177e3367c14/litestar-2.24.0-py3-none-any.whl", hash = "sha256:0ef13630173ea147847363f03f0459877ab14a572e68b2af7a25796055513a31", size = 581807, upload-time = "2026-06-11T11:26:58.509Z" }, +] + +[[package]] +name = "litestar-htmx" +version = "0.5.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/3f/b9/7e296aa1adada25cce8e5f89a996b0e38d852d93b1b656a2058226c542a2/litestar_htmx-0.5.0.tar.gz", hash = "sha256:e02d1a3a92172c874835fa3e6749d65ae9fc626d0df46719490a16293e2146fb", size = 119755, upload-time = "2025-06-11T21:19:45.573Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/f2/24/8d99982f0aa9c1cd82073c6232b54a0dbe6797c7d63c0583a6c68ee3ddf2/litestar_htmx-0.5.0-py3-none-any.whl", hash = "sha256:92833aa47e0d0e868d2a7dbfab75261f124f4b83d4f9ad12b57b9a68f86c50e6", size = 9970, upload-time = "2025-06-11T21:19:44.465Z" }, +] + [[package]] name = "markdown" version = "3.10.2" @@ -1144,6 +1177,22 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/a4/8e/469e5a4a2f5855992e425f3cb33804cc07bf18d48f2db061aec61ce50270/more_itertools-10.8.0-py3-none-any.whl", hash = "sha256:52d4362373dcf7c52546bc4af9a86ee7c4579df9a8dc268be0a2f949d376cc9b", size = 69667, upload-time = "2025-09-02T15:23:09.635Z" }, ] +[[package]] +name = "msgspec" +version = "0.21.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/e3/60/f79b9b013a16fa3a58350c9295ddc6789f2e335f36ea61ed10a21b215364/msgspec-0.21.1.tar.gz", hash = "sha256:2313508e394b0d208f8f56892ca9b2799e2561329de9763b19619595a6c0f72c", size = 319193, upload-time = "2026-04-12T21:44:50.394Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7e/74/f11ede02839b19ff459f88e3145df5d711626ca84da4e23520cebf819367/msgspec-0.21.1-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:764173717a01743f007e9f74520ed281f24672c604514f7d76c1c3a10e8edb66", size = 196176, upload-time = "2026-04-12T21:44:17.613Z" }, + { url = "https://files.pythonhosted.org/packages/bb/40/4476c1bd341418a046c4955aff632ec769315d1e3cb94e6acf86d461f9ed/msgspec-0.21.1-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:344c7cd0eaed1fb81d7959f99100ef71ec9b536881a376f11b9a6c4803365697", size = 188524, upload-time = "2026-04-12T21:44:18.815Z" }, + { url = "https://files.pythonhosted.org/packages/ca/d9/9e9d7d7e5061b47540d03d640fab9b3965ba7ae49c1b2154861c8f007518/msgspec-0.21.1-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:48943e278b3854c2f89f955ddc6f9f430d3f0784b16e47d10604ee0463cd21f5", size = 218880, upload-time = "2026-04-12T21:44:20.028Z" }, + { url = "https://files.pythonhosted.org/packages/74/66/2bb344f34abb4b57e60c7c9c761994e0417b9718ec1460bf00c296f2a7ea/msgspec-0.21.1-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:a9aa659ebb0101b1cbc31461212b87e341d961f0ab0772aaf068a99e001ec4aa", size = 225050, upload-time = "2026-04-12T21:44:21.577Z" }, + { url = "https://files.pythonhosted.org/packages/1a/84/7c1e412f76092277bf760cef12b7979d03314d259ab5b5cafde5d0c1722d/msgspec-0.21.1-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:f7b27d1a8ead2b6f5b0c4f2d07b8be1ccfcc041c8a0e704781edebe3ae13c484", size = 222713, upload-time = "2026-04-12T21:44:22.83Z" }, + { url = "https://files.pythonhosted.org/packages/4e/27/0bba04b2b4ef05f3d068429410bc71d2cea925f1596a8f41152cccd5edb8/msgspec-0.21.1-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:38fe93e86b61328fe544cb7fd871fad5a27c8734bfda90f65e5dbe288ae50f61", size = 227259, upload-time = "2026-04-12T21:44:24.11Z" }, + { url = "https://files.pythonhosted.org/packages/b0/2d/09574b0eea02fed2c2c1383dbaae2c7f79dc16dcd6487a886000afb5d7c4/msgspec-0.21.1-cp313-cp313-win_amd64.whl", hash = "sha256:8bc666331c35fcce05a7cd2d6221adbe0f6058f8e750711413d22793c080ac6a", size = 189857, upload-time = "2026-04-12T21:44:25.359Z" }, + { url = "https://files.pythonhosted.org/packages/46/34/105b1576ad182879914f0c821f17ee1d13abb165cb060448f96fe2aff078/msgspec-0.21.1-cp313-cp313-win_arm64.whl", hash = "sha256:42bb1241e0750c1a4346f2aa84db26c5ffd99a4eb3a954927d9f149ff2f42898", size = 175403, upload-time = "2026-04-12T21:44:26.608Z" }, +] + [[package]] name = "multidict" version = "6.7.1" @@ -1189,6 +1238,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/81/08/7036c080d7117f28a4af526d794aab6a84463126db031b007717c1a6676e/multidict-6.7.1-py3-none-any.whl", hash = "sha256:55d97cc6dae627efa6a6e548885712d4864b81110ac76fa4e534c03819fa4a56", size = 12319, upload-time = "2026-01-26T02:46:44.004Z" }, ] +[[package]] +name = "multipart" +version = "1.3.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/8e/d6/9c4f366d6f9bb8f8fb5eae3acac471335c39510c42b537fd515213d7d8c3/multipart-1.3.1.tar.gz", hash = "sha256:211d7cfc1a7a43e75c4d24ee0e8e0f4f61d522f1a21575303ae85333dea687bf", size = 38929, upload-time = "2026-02-27T10:17:13.7Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/19/ed/e1f03200ee1f0bf4a2b9b72709afefbf5319b68df654e0b84b35c65613ee/multipart-1.3.1-py3-none-any.whl", hash = "sha256:a82b59e1befe74d3d30b3d3f70efd5a2eba4d938f845dcff9faace968888ff29", size = 15061, upload-time = "2026-02-27T10:17:11.943Z" }, +] + [[package]] name = "mypy-extensions" version = "1.1.0" @@ -1377,6 +1435,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/54/20/4d324d65cc6d9205fabedc306948156824eb9f0ee1633355a8f7ec5c66bf/pluggy-1.6.0-py3-none-any.whl", hash = "sha256:e920276dd6813095e9377c0bc5566d94c932c33b27a3e3945d8389c374dd4746", size = 20538, upload-time = "2025-05-15T12:30:06.134Z" }, ] +[[package]] +name = "polyfactory" +version = "3.3.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "faker" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/85/68/7717bd9e63ed254617a7d3dc9260904fb736d6ea203e58ffddcb186c64e4/polyfactory-3.3.0.tar.gz", hash = "sha256:237258b6ff43edf362ffd1f68086bb796466f786adfa002b0ac256dbf2246e9a", size = 348668, upload-time = "2026-02-22T09:46:28.01Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/dd/34/b6f19941adcdaf415b5e8a8d577499f5b6a76b59cbae37f9b125a9ffe9f2/polyfactory-3.3.0-py3-none-any.whl", hash = "sha256:686abcaa761930d3df87b91e95b26b8d8cb9fdbbbe0b03d5f918acff5c72606e", size = 62707, upload-time = "2026-02-22T09:46:25.985Z" }, +] + [[package]] name = "pre-commit" version = "4.5.1" @@ -1637,11 +1708,11 @@ wheels = [ [[package]] name = "python-multipart" -version = "0.0.22" +version = "0.0.32" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/94/01/979e98d542a70714b0cb2b6728ed0b7c46792b695e3eaec3e20711271ca3/python_multipart-0.0.22.tar.gz", hash = "sha256:7340bef99a7e0032613f56dc36027b959fd3b30a787ed62d310e951f7c3a3a58", size = 37612, upload-time = "2026-01-25T10:15:56.219Z" } +sdist = { url = "https://files.pythonhosted.org/packages/5b/42/55c32bb9b12693c092ad250a0e82edb5b31ddeda6eb772de5f308b3804ad/python_multipart-0.0.32.tar.gz", hash = "sha256:be54b7f3fa167bb83e4fcd936b887b708f4e57fe75911c02aebf53efaf8d938e", size = 46881, upload-time = "2026-06-04T16:18:58.647Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/1b/d0/397f9626e711ff749a95d96b7af99b9c566a9bb5129b8e4c10fc4d100304/python_multipart-0.0.22-py3-none-any.whl", hash = "sha256:2b2cd894c83d21bf49d702499531c7bafd057d730c201782048f7945d82de155", size = 24579, upload-time = "2026-01-25T10:15:54.811Z" }, + { url = "https://files.pythonhosted.org/packages/e1/04/e8135ebd1ad02c56ec633277529b2602ff99ff634be76cdba5744cf554fd/python_multipart-0.0.32-py3-none-any.whl", hash = "sha256:ff6d3f776f16878c894e52e107296ffc890e913c611b1a4ec6c44e2821fe2e23", size = 30042, upload-time = "2026-06-04T16:18:57.319Z" }, ] [[package]] @@ -1739,6 +1810,20 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/14/25/b208c5683343959b670dc001595f2f3737e051da617f66c31f7c4fa93abc/rich-14.3.3-py3-none-any.whl", hash = "sha256:793431c1f8619afa7d3b52b2cdec859562b950ea0d4b6b505397612db8d5362d", size = 310458, upload-time = "2026-02-19T17:23:13.732Z" }, ] +[[package]] +name = "rich-click" +version = "1.9.8" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "click" }, + { name = "colorama", marker = "sys_platform == 'win32'" }, + { name = "rich" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/f7/ea/21e4867ea0ef881ffd4c0550fc21a061435e50d6324bcd034396633cbc18/rich_click-1.9.8.tar.gz", hash = "sha256:4008f921da88b5d91646c134ec881c1500e5a6b3f093e90e8f29400e09608371", size = 75363, upload-time = "2026-05-28T19:54:59.144Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/6d/97/a87901aef6b7e7e4a34c6dd6cc17dca8594a592ef9d9dd765fca2b7facf7/rich_click-1.9.8-py3-none-any.whl", hash = "sha256:12873865396e6927835d4eabb1cc3996edcd65b7ac9b2391a29eca4f335a2f93", size = 72189, upload-time = "2026-05-28T19:54:57.867Z" }, +] + [[package]] name = "ruff" version = "0.15.8" @@ -1859,14 +1944,14 @@ wheels = [ [[package]] name = "starlette" -version = "1.0.0" +version = "1.3.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "anyio" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/81/69/17425771797c36cded50b7fe44e850315d039f28b15901ab44839e70b593/starlette-1.0.0.tar.gz", hash = "sha256:6a4beaf1f81bb472fd19ea9b918b50dc3a77a6f2e190a12954b25e6ed5eea149", size = 2655289, upload-time = "2026-03-22T18:29:46.779Z" } +sdist = { url = "https://files.pythonhosted.org/packages/eb/e3/7c1dc7381d9f8ab7d854328ebfa884e62cb3f3d8549ddfd37c7814f42afa/starlette-1.3.1.tar.gz", hash = "sha256:05d0213193f2fbaae60e2ecb593b4add4262ad4e46536b54abe36f11a71724e0", size = 2703240, upload-time = "2026-06-12T09:23:11.602Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/0b/c9/584bc9651441b4ba60cc4d557d8a547b5aff901af35bda3a4ee30c819b82/starlette-1.0.0-py3-none-any.whl", hash = "sha256:d3ec55e0bb321692d275455ddfd3df75fff145d009685eb40dc91fc66b03d38b", size = 72651, upload-time = "2026-03-22T18:29:45.111Z" }, + { url = "https://files.pythonhosted.org/packages/ec/bb/2799cc2ede3ed41131f8975621e7213dfc7ef4acbbaadfa440f32500c370/starlette-1.3.1-py3-none-any.whl", hash = "sha256:c7372aae11c3c3f26a42df7bd626cec2f47d03483d261d369516a615a53714c6", size = 73632, upload-time = "2026-06-12T09:23:10.017Z" }, ] [[package]] @@ -1890,6 +1975,9 @@ fastapi = [ { name = "fastapi" }, { name = "python-multipart" }, ] +litestar = [ + { name = "litestar" }, +] [[package]] name = "structlog" From d92088086c62e8e7c88c8d240ac7b0b5648b02ae Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 05:39:02 +0200 Subject: [PATCH 02/12] refactor(api/state): replace authlib client and FastAPI state accessor --- api/src/damnit_api/state.py | 50 ++++++++++++------- ...aphql_query_without_session_rejected.yaml} | 0 2 files changed, 33 insertions(+), 17 deletions(-) rename api/tests/refactor/e2e/cassettes/test_authz_parity/{test_graphql_query_without_session_unchanged.yaml => test_graphql_query_without_session_rejected.yaml} (100%) diff --git a/api/src/damnit_api/state.py b/api/src/damnit_api/state.py index e362a803..4917c5e3 100644 --- a/api/src/damnit_api/state.py +++ b/api/src/damnit_api/state.py @@ -8,31 +8,52 @@ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import TYPE_CHECKING -from fastapi import ( - Request, # noqa: TC002 - FastAPI DI inspects annotations at runtime +from litestar.datastructures import ( + State as LitestarState, # noqa: TC002 - Litestar inspects annotations at runtime via get_type_hints ) from sqlalchemy.ext.asyncio import AsyncEngine, async_sessionmaker, create_async_engine from sqlmodel.ext.asyncio.session import AsyncSession if TYPE_CHECKING: - from authlib.integrations.starlette_client import StarletteOAuth2App - from ._mymdc.clients import MyMdCClient from .auth.token_store import TokenStore from .graphql.subscriptions import SubscriptionCursors from .runs.repository import DamnitRepositoryRegistry from .shared.settings import Settings +# Session cookie name, shared by `main.py`'s `CookieBackendConfig` and the +# logout handlers in `auth/routers.py` so the two cannot drift. +SESSION_COOKIE_KEY = "session" + + +@dataclass +class OAuthClient: + """OAuth2/OIDC client configuration with lazily loaded server metadata.""" + + client_id: str + client_secret: str + scope: str + server_metadata_url: str + server_metadata: dict = field(default_factory=dict) + + async def load_server_metadata(self) -> None: + import httpx + + async with httpx.AsyncClient() as http: + resp = await http.get(self.server_metadata_url) + resp.raise_for_status() + self.server_metadata = resp.json() + @dataclass(frozen=True) class AppState: db_engine: AsyncEngine db_sessionmaker: async_sessionmaker[AsyncSession] mymdc_client: MyMdCClient - oauth_client: StarletteOAuth2App | None # None when auth is disabled + oauth_client: OAuthClient | None # None when auth is disabled token_store: TokenStore repositories: DamnitRepositoryRegistry subscription_cursors: SubscriptionCursors @@ -62,21 +83,16 @@ def create_mymdc_client(settings: Settings) -> MyMdCClient: raise ValueError(msg) -def create_oauth_client(settings: Settings) -> StarletteOAuth2App | None: +def create_oauth_client(settings: Settings) -> OAuthClient | None: if settings.auth is None: return None - from authlib.integrations.starlette_client import OAuth - - oauth = OAuth() - oauth.register( - name="damnit_web", + return OAuthClient( client_id=settings.auth.client_id, client_secret=settings.auth.client_secret.get_secret_value(), + scope="openid email groups", server_metadata_url=str(settings.auth.server_metadata_url), - client_kwargs={"scope": "openid email groups"}, ) - return oauth.damnit_web # pyright: ignore[reportReturnType] def create_token_store() -> TokenStore: @@ -103,6 +119,6 @@ def create_subscription_cursors() -> SubscriptionCursors: return SubscriptionCursors() -def get_app_state(request: Request) -> AppState: - """FastAPI dependency: the application's :class:`AppState`.""" - return request.app.state.app_state +def provide_app_state(state: LitestarState) -> AppState: + """Litestar dependency: the application's :class:`AppState`.""" + return state.app_state # type: ignore[attr-defined] diff --git a/api/tests/refactor/e2e/cassettes/test_authz_parity/test_graphql_query_without_session_unchanged.yaml b/api/tests/refactor/e2e/cassettes/test_authz_parity/test_graphql_query_without_session_rejected.yaml similarity index 100% rename from api/tests/refactor/e2e/cassettes/test_authz_parity/test_graphql_query_without_session_unchanged.yaml rename to api/tests/refactor/e2e/cassettes/test_authz_parity/test_graphql_query_without_session_rejected.yaml From eb9866805ed9969777cfa14a119532df87fd6376 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 05:39:26 +0200 Subject: [PATCH 03/12] refactor(api): port per-slice dependency providers to litestar --- api/src/damnit_api/_db/dependencies.py | 14 ++++++-------- api/src/damnit_api/_mymdc/dependencies.py | 15 ++++++++------- api/src/damnit_api/graphql/dependencies.py | 16 ++++++---------- api/src/damnit_api/runs/dependencies.py | 14 ++++++-------- 4 files changed, 26 insertions(+), 33 deletions(-) diff --git a/api/src/damnit_api/_db/dependencies.py b/api/src/damnit_api/_db/dependencies.py index 54e5f029..0a7f7873 100644 --- a/api/src/damnit_api/_db/dependencies.py +++ b/api/src/damnit_api/_db/dependencies.py @@ -1,19 +1,17 @@ -"""FastAPI dependency helpers for database sessions.""" +"""Litestar dependency helpers for database sessions.""" from collections.abc import AsyncIterator -from typing import Annotated -from fastapi import Depends, Request +from litestar.datastructures import State from sqlmodel.ext.asyncio.session import AsyncSession -from ..state import get_app_state - -async def get_session(request: Request) -> AsyncIterator[AsyncSession]: +async def get_session(state: State) -> AsyncIterator[AsyncSession]: """Provide a database session from the application state.""" - async with get_app_state(request).db_sessionmaker() as session: + async with state.app_state.db_sessionmaker() as session: # type: ignore[attr-defined] yield session -DBSession = Annotated[AsyncSession, Depends(get_session)] +# Plain type alias; Litestar injects by the parameter name `session`. +DBSession = AsyncSession diff --git a/api/src/damnit_api/_mymdc/dependencies.py b/api/src/damnit_api/_mymdc/dependencies.py index a9ad18e1..bcb500dd 100644 --- a/api/src/damnit_api/_mymdc/dependencies.py +++ b/api/src/damnit_api/_mymdc/dependencies.py @@ -1,15 +1,16 @@ -from typing import Annotated +from litestar.datastructures import State +from litestar.params import SkipValidation -from fastapi import Depends, Request - -from ..state import get_app_state from . import clients -def get_mymdc_client(request: Request) -> "clients.MyMdCClient": +def get_mymdc_client(state: State) -> "clients.MyMdCClient": """Provide the MyMdC client from the application state.""" - return get_app_state(request).mymdc_client + return state.app_state.mymdc_client # type: ignore[attr-defined] -MyMdCClient = Annotated[clients.MyMdCClient, Depends(get_mymdc_client)] +# `MyMdCClient` is a union of two concrete clients; Litestar's msgspec-based +# signature validation cannot build a decoder for a union of custom types, so +# injection sites skip validation of this app-provided collaborator. +MyMdCClient = SkipValidation[clients.MyMdCClient] """Type alias for the MyMdC client dependency.""" diff --git a/api/src/damnit_api/graphql/dependencies.py b/api/src/damnit_api/graphql/dependencies.py index f6fcc03d..7a0c69c8 100644 --- a/api/src/damnit_api/graphql/dependencies.py +++ b/api/src/damnit_api/graphql/dependencies.py @@ -1,19 +1,15 @@ -"""FastAPI dependency helpers for GraphQL subscription state.""" +"""Litestar dependency helpers for GraphQL subscription state.""" -from typing import Annotated +from litestar.datastructures import State -from fastapi import Depends, Request - -from ..state import get_app_state from .subscriptions import SubscriptionCursors -def get_subscription_cursors(request: Request) -> SubscriptionCursors: +def get_subscription_cursors(state: State) -> SubscriptionCursors: """Provide the subscription cursors from the application state.""" - return get_app_state(request).subscription_cursors + return state.app_state.subscription_cursors # type: ignore[attr-defined] -SubscriptionCursorsDep = Annotated[ - SubscriptionCursors, Depends(get_subscription_cursors) -] +# Plain type alias; Litestar injects by the parameter name `subscription_cursors`. +SubscriptionCursorsDep = SubscriptionCursors """Type alias for the subscription cursors dependency.""" diff --git a/api/src/damnit_api/runs/dependencies.py b/api/src/damnit_api/runs/dependencies.py index 4f24a86c..537dac78 100644 --- a/api/src/damnit_api/runs/dependencies.py +++ b/api/src/damnit_api/runs/dependencies.py @@ -1,17 +1,15 @@ -"""FastAPI dependency helpers for the runs repository registry.""" +"""Litestar dependency helpers for the runs repository registry.""" -from typing import Annotated +from litestar.datastructures import State -from fastapi import Depends, Request - -from ..state import get_app_state from .repository import DamnitRepositoryRegistry -def get_repositories(request: Request) -> DamnitRepositoryRegistry: +def get_repositories(state: State) -> DamnitRepositoryRegistry: """Provide the per-proposal repository registry from the application state.""" - return get_app_state(request).repositories + return state.app_state.repositories # type: ignore[attr-defined] -Repositories = Annotated[DamnitRepositoryRegistry, Depends(get_repositories)] +# Plain type alias; Litestar injects by the parameter name `repositories`. +Repositories = DamnitRepositoryRegistry """Type alias for the DAMNIT repository registry dependency.""" From 63c35d1d6a7a542db5496bdad3af4bbe3c5c5772 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 05:39:59 +0200 Subject: [PATCH 04/12] refactor(api/auth): port auth di, token store and models to litestar --- api/src/damnit_api/auth/__init__.py | 5 +- api/src/damnit_api/auth/dependencies.py | 76 +++++++++++++++---------- api/src/damnit_api/auth/models.py | 10 ++-- api/src/damnit_api/auth/token_store.py | 18 ++++-- 4 files changed, 67 insertions(+), 42 deletions(-) diff --git a/api/src/damnit_api/auth/__init__.py b/api/src/damnit_api/auth/__init__.py index caa65331..a831d473 100644 --- a/api/src/damnit_api/auth/__init__.py +++ b/api/src/damnit_api/auth/__init__.py @@ -1,3 +1,4 @@ -from .routers import noauth_router, router +from . import dependencies +from .routers import NoAuthOAuthController, OAuthController -__all__ = ["noauth_router", "router"] +__all__ = ["NoAuthOAuthController", "OAuthController", "dependencies"] diff --git a/api/src/damnit_api/auth/dependencies.py b/api/src/damnit_api/auth/dependencies.py index 0ab22862..bca3b2e9 100644 --- a/api/src/damnit_api/auth/dependencies.py +++ b/api/src/damnit_api/auth/dependencies.py @@ -1,50 +1,64 @@ -"""Dependency type aliases for the auth module.""" +"""Dependency functions and type aliases for the auth module.""" -from typing import Annotated +from collections.abc import AsyncIterator -from authlib.integrations.starlette_client import StarletteOAuth2App -from fastapi import Depends, Request +from authlib.integrations import httpx_client +from authlib.integrations.httpx_client import AsyncOAuth2Client +from litestar import Request +from litestar.datastructures import State +from sqlmodel.ext.asyncio.session import AsyncSession -from ..state import get_app_state +from .._mymdc.dependencies import MyMdCClient +from ..state import OAuthClient from .models import OAuthUserInfo as _OAuthUserInfo from .models import User as _User from .token_store import TokenStore -# TODO: Get from settings -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 +def get_oauth_client(state: State) -> OAuthClient: + client = state.app_state.oauth_client # type: ignore[attr-defined] if client is None: - msg = "OAuth client is not configured (auth is disabled)." + msg = ( + "OAuth client is not configured (settings.auth is None); " + "enable auth settings to use OAuth endpoints." + ) raise RuntimeError(msg) return client -def get_token_store(request: Request) -> TokenStore: - """Provide the token store from the application state.""" - return get_app_state(request).token_store +async def get_oauth_http_client( + oauth_config: OAuthClient, +) -> AsyncIterator[AsyncOAuth2Client]: + """Litestar dependency: a short-lived OAuth2 HTTP client, closed by DI.""" + client = httpx_client.AsyncOAuth2Client( + client_id=oauth_config.client_id, + client_secret=oauth_config.client_secret, + scope=oauth_config.scope, + ) + try: + yield client + finally: + await client.aclose() + + +def get_token_store(state: State) -> TokenStore: + return state.app_state.token_store # type: ignore[attr-defined] -RedirectURI = Annotated[str, Depends(_get_default_redirect_login_uri)] -"""Type alias for the redirect URI dependency.""" +def get_oauth_user_info(request: Request) -> _OAuthUserInfo: + """Litestar dependency: resolve OAuthUserInfo from the session.""" + return _OAuthUserInfo.from_connection(request) # type: ignore[arg-type] -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.""" +async def get_user( + request: Request, + mymdc: MyMdCClient, + session: AsyncSession, +) -> _User: + """Litestar dependency: resolve full User (with proposals) from session + DB.""" + return await _User.from_connection(request, mymdc, session) # type: ignore[arg-type] -OAuthUserInfo = Annotated[_OAuthUserInfo, Depends(_OAuthUserInfo.from_connection)] -"""Type alias for the OAuth user info dependency.""" -User = Annotated[_User, Depends(_User.from_connection)] -"""Type alias for the full User dependency.""" +# Plain type re-exports; consumed by other modules as annotations. +OAuthUserInfo = _OAuthUserInfo +User = _User diff --git a/api/src/damnit_api/auth/models.py b/api/src/damnit_api/auth/models.py index 2cfc226e..de2976c8 100644 --- a/api/src/damnit_api/auth/models.py +++ b/api/src/damnit_api/auth/models.py @@ -2,7 +2,7 @@ from typing import Self -from fastapi.requests import HTTPConnection +from litestar.connection import ASGIConnection from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, RootModel from .. import get_logger @@ -33,12 +33,12 @@ class OAuthUserInfo(BaseUserInfo): """Basic user information obtained from the OAuth provider.""" @classmethod - def from_connection(cls, connection: HTTPConnection) -> Self: + def from_connection(cls, connection: ASGIConnection) -> Self: """Create an OAuthUserInfo from the request session. !!! note - Dependency on `HTTPConnection` instead of `Request` to support websockets. + Dependency on `ASGIConnection` instead of `Request` to support websockets. """ user_dict = connection.session.get("user") if user_dict is None: @@ -82,7 +82,7 @@ def proposals(self) -> list[ProposalNumber]: @classmethod async def from_connection( cls, - connection: HTTPConnection, + connection: ASGIConnection, mymdc: MyMdCClient, session: DBSession, ) -> Self: @@ -91,7 +91,7 @@ async def from_connection( !!! note - Dependency on `HTTPConnection` instead of `Request` to support websockets. + Dependency on `ASGIConnection` instead of `Request` to support websockets. """ oauth = OAuthUserInfo.from_connection(connection) diff --git a/api/src/damnit_api/auth/token_store.py b/api/src/damnit_api/auth/token_store.py index cccab754..ddf64334 100644 --- a/api/src/damnit_api/auth/token_store.py +++ b/api/src/damnit_api/auth/token_store.py @@ -1,11 +1,15 @@ """TokenStore protocol and implementations.""" -from typing import Any, Protocol +from typing import Any, Protocol, runtime_checkable +# `runtime_checkable` is required so Litestar's msgspec-based signature +# validation (which does `isinstance(value, TokenStore)`) doesn't raise a +# TypeError on every request that injects this dependency. +@runtime_checkable class TokenStore(Protocol): def store(self, sub: str, token: dict[str, Any]) -> None: ... - def pop_token_field(self, sub: str, field: str) -> Any: ... + def pop_token_field(self, sub: str, field: str) -> str | None: ... class InMemoryTokenStore: @@ -15,5 +19,11 @@ def __init__(self) -> None: 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) + def pop_token_field(self, sub: str, field: str) -> str | None: + token = self._tokens.get(sub) + if token is None: + return None + value = token.pop(field, None) + if not token: + del self._tokens[sub] + return value From 05c849e1fb237b6961b549a57fee06f757fcc3ef Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 05:40:26 +0200 Subject: [PATCH 05/12] refactor(api): convert routers to litestar and repair oauth flow --- api/src/damnit_api/auth/routers.py | 395 +++++++++++++--------- api/src/damnit_api/contextfile/routers.py | 28 +- api/src/damnit_api/metadata/routers.py | 30 +- api/src/damnit_api/shared/settings.py | 7 + 4 files changed, 286 insertions(+), 174 deletions(-) diff --git a/api/src/damnit_api/auth/routers.py b/api/src/damnit_api/auth/routers.py index dd616f4c..d450a990 100644 --- a/api/src/damnit_api/auth/routers.py +++ b/api/src/damnit_api/auth/routers.py @@ -1,164 +1,257 @@ -from urllib.parse import parse_qs, unquote, urlencode +"""OAuth2 authentication route handlers (Litestar).""" -from authlib.integrations.starlette_client import OAuthError -from fastapi import APIRouter, HTTPException, Request -from fastapi.responses import JSONResponse, RedirectResponse +from typing import Annotated, ClassVar +from urllib.parse import urlencode, urlparse, urlunparse + +from authlib.integrations.httpx_client import AsyncOAuth2Client +from litestar import Controller, Request, get, post +from litestar.background_tasks import BackgroundTask +from litestar.datastructures import Cookie +from litestar.di import Provide +from litestar.exceptions import HTTPException +from litestar.params import Dependency +from litestar.response import Redirect, Response from .. import get_logger from .._db.dependencies import DBSession from .._mymdc.dependencies import MyMdCClient +from ..runs.dependencies import Repositories +from ..state import SESSION_COOKIE_KEY, OAuthClient from . import dependencies, models +from .token_store import TokenStore logger = get_logger() -router = APIRouter(prefix="/oauth", tags=["auth"]) - - -@router.get("/login", status_code=307) -async def auth( - request: Request, - redirect_uri: dependencies.RedirectURI, - client: dependencies.Client, -) -> RedirectResponse: - """Initiate the OAuth2 login flow.""" - # Note: session is managed by Starlette's SessionMiddleware, which handles - # expiration. - if request.session.get("user"): - return RedirectResponse(url=redirect_uri) - - callback_uri = request.url_for("callback") - - # TODO: error on non-HTTPS in production? - - if x_forwarded_host := request.headers.get("x-forwarded-host"): - callback_uri = callback_uri.replace(netloc=x_forwarded_host) - - res = await client.authorize_redirect(request, redirect_uri=str(callback_uri)) - - await logger.adebug("OAuth redirect response", headers=res.headers) - - return res - - -@router.get("/callback", status_code=307) -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.""" +DEFAULT_LOGIN_REDIRECT = "/app/home" + + +def _build_callback_uri(request: Request) -> str: + """Return the /oauth/callback URL, honouring x-forwarded-host if trusted.""" + uri = str(request.url_for("callback")) + + from ..shared.settings import settings + + if settings.trust_forwarded_host and ( + x_forwarded_host := request.headers.get("x-forwarded-host") + ): + parsed = urlparse(uri) + uri = urlunparse(parsed._replace(netloc=x_forwarded_host)) + return uri + + +def _sanitize_redirect_target(target: str | None) -> str: + """Allow-list post-login redirects: relative, same-origin paths only.""" + if not target: + return DEFAULT_LOGIN_REDIRECT + parsed = urlparse(target) + if ( + parsed.scheme + or parsed.netloc + or not target.startswith("/") + or target.startswith("//") + ): + return DEFAULT_LOGIN_REDIRECT + return target + + +async def _revoke_tokens( + oauth_config: OAuthClient, + revocation_endpoint: str, + tokens: list[tuple[str, str]], +) -> None: + """Best-effort token revocation; failures are logged and swallowed.""" + from authlib.integrations import httpx_client + + client = httpx_client.AsyncOAuth2Client( + client_id=oauth_config.client_id, + client_secret=oauth_config.client_secret, + ) try: - token = await client.authorize_access_token(request) - user = await client.userinfo(token=token) - except OAuthError as e: - raise HTTPException(status_code=401, detail=str(e)) from e - - state = request.query_params.get("state", "") - state_params = parse_qs(state) - redirect_uri = state_params.get("redirect_uri", [redirect_uri])[0] - - request.session["user"] = dict(user) - - # TODO: could (should?) be stored in db for persistence across server restarts - # NOTE: required for revoking tokens on logout - token_store.store(str(user["sub"]), token) - - return RedirectResponse(url=unquote(str(redirect_uri))) - - -@router.post("/logout") -async def logout( - request: Request, - client: dependencies.Client, - token_store: dependencies.TokenStoreDep, -) -> JSONResponse: - """Fully logout the user and revoke tokens. - - Endpoint (attempts to) revoke both access and refresh tokens via revocation - endpoint, then redirects to end session endpoint if available. - """ - user_sub = request.session.get("user", {}).get("sub") - - revocation_endpoint = client.server_metadata.get("revocation_endpoint", None) - - if client.client_id and client.client_secret: - auth = (client.client_id, client.client_secret) - token = await client.fetch_access_token() - - for k in ("refresh_token", "access_token"): - if user_token := token_store.pop_token_field(user_sub, k): - await client.post( + for token_type_hint, token in tokens: + try: + await client.post( # type: ignore[attr-defined] revocation_endpoint, - token=token, - auth=auth, - data={"token": user_token, "token_type_hint": f"{k}"}, + data={"token": token, "token_type_hint": token_type_hint}, ) - - end_session_endpoint = client.server_metadata.get("end_session_endpoint") - - token_id = token_store.pop_token_field(user_sub, "id_token") - - logout_url = None - if token_id and end_session_endpoint: - params = {} - if token_id: - params["id_token_hint"] = token_id - # TODO: request post_logout_redirect_uri added to keycloak client - # if logout_redirect := request.headers.get("x-forwarded-host"): - # params["post_logout_redirect_uri"] = logout_redirect - logout_url = f"{end_session_endpoint}?{urlencode(params)}" - - try: + except Exception: + await logger.adebug( + "Token revocation failed", token_type=token_type_hint + ) + finally: + await client.aclose() # type: ignore[attr-defined] + + +class OAuthController(Controller): + """OAuth2 login/callback/logout/userinfo endpoints.""" + + path = "/oauth" + tags: ClassVar[list[str]] = ["auth"] # ty: ignore[invalid-attribute-override] + dependencies: ClassVar[dict[str, Provide]] = { # ty: ignore[invalid-attribute-override] + "oauth_http_client": Provide(dependencies.get_oauth_http_client), + } + + @get("/login", status_code=302) + async def auth( + self, + request: Request, + oauth_config: OAuthClient, + oauth_http_client: Annotated[ + AsyncOAuth2Client, Dependency(skip_validation=True) + ], + redirect_uri: str | None = None, + ) -> Redirect: + """Initiate the OAuth2 login flow.""" + target = _sanitize_redirect_target(redirect_uri) + if request.session.get("user"): + return Redirect(path=target) + + callback_uri = _build_callback_uri(request) + + url, state = oauth_http_client.create_authorization_url( + oauth_config.server_metadata["authorization_endpoint"], + redirect_uri=callback_uri, + ) + + request.session["_oauth_state"] = state + request.session["_login_redirect"] = target + await logger.adebug("OAuth redirect", url=url) + return Redirect(path=url) + + @get("/callback", name="callback", status_code=302) + async def callback( + self, + request: Request, + oauth_config: OAuthClient, + token_store: TokenStore, + oauth_http_client: Annotated[ + AsyncOAuth2Client, Dependency(skip_validation=True) + ], + ) -> Redirect: + """OAuth2 callback: exchange code for token and set session.""" + saved_state = request.session.pop("_oauth_state", None) + received_state = request.query_params.get("state") + + if saved_state is None or saved_state != received_state: + msg = "OAuth state mismatch — possible CSRF attempt" + raise HTTPException(status_code=401, detail=msg) + + callback_uri = _build_callback_uri(request) + + try: + token = await oauth_http_client.fetch_token( + oauth_config.server_metadata["token_endpoint"], + authorization_response=str(request.url), + redirect_uri=callback_uri, + ) + userinfo_resp = await oauth_http_client.get( + oauth_config.server_metadata["userinfo_endpoint"], + token=token, # ty: ignore[unknown-argument] + ) + userinfo_resp.raise_for_status() + user = userinfo_resp.json() + except HTTPException: + raise + except Exception as e: + msg = str(e) + raise HTTPException(status_code=401, detail=msg) from e + + # Symmetric with /login: the post-login destination travels in the + # session and is re-validated against the relative-path allow-list. + target = _sanitize_redirect_target(request.session.pop("_login_redirect", None)) + + request.session["user"] = user + token_store.store(str(user["sub"]), token) + + return Redirect(path=target) + + @post("/logout") + async def logout( + self, + request: Request, + oauth_config: OAuthClient, + token_store: TokenStore, + ) -> Response: + """Clear the session; revoke tokens in the background.""" + user_sub = request.session.get("user", {}).get("sub") + revocation_endpoint = oauth_config.server_metadata.get("revocation_endpoint") + end_session_endpoint = oauth_config.server_metadata.get( + "end_session_endpoint" + ) + + tokens_to_revoke: list[tuple[str, str]] = [] + if ( + revocation_endpoint + and oauth_config.client_id + and oauth_config.client_secret + ): + tokens_to_revoke = [ + (k, token) + for k in ("refresh_token", "access_token") + if (token := token_store.pop_token_field(user_sub, k)) + ] + + token_id = token_store.pop_token_field(user_sub, "id_token") + logout_url = None + if token_id and end_session_endpoint: + params = {"id_token_hint": token_id} + logout_url = f"{end_session_endpoint}?{urlencode(params)}" + + try: + request.session.pop("user", None) + except Exception: + await logger.adebug("Failed clearing user session during logout") + + background = None + if tokens_to_revoke and revocation_endpoint: + background = BackgroundTask( + _revoke_tokens, oauth_config, revocation_endpoint, tokens_to_revoke + ) + + return Response( + content={"logout_url": logout_url}, + cookies=[Cookie(key=SESSION_COOKIE_KEY, value="", max_age=0, path="/")], + background=background, + ) + + @get("/userinfo") + async def userinfo( + self, + request: Request, + mymdc: MyMdCClient, + session: DBSession, + with_proposals: bool = True, + ) -> models.User | models.OAuthUserInfo: + """User information.""" + if with_proposals: + user = await models.User.from_connection(request, mymdc, session) + else: + user = models.OAuthUserInfo.from_connection(request) + return user + + +class NoAuthOAuthController(Controller): + """Local-mode (auth-disabled) equivalents of the OAuth endpoints.""" + + path = "/oauth" + tags: ClassVar[list[str]] = ["auth"] # ty: ignore[invalid-attribute-override] + + @get("/userinfo") + async def userinfo(self, repositories: Repositories) -> dict: + """User info for local (auth-disabled) mode.""" + from ..metadata.services import LOCAL_CYCLE, _local_proposal_number + + proposals = {} + proposal_number = await _local_proposal_number(repositories) + if proposal_number: + proposals = {LOCAL_CYCLE: [proposal_number]} + + return {**models.DEV_USER.model_dump(), "proposals_by_year_half": proposals} + + @post("/logout", sync_to_thread=False) + def logout(self, request: Request) -> Response: + """Logout for local (auth-disabled) mode.""" request.session.pop("user", None) - except Exception: - await logger.adebug("Failed clearing user session during logout") - - response = JSONResponse(status_code=200, content={"logout_url": logout_url}) - response.delete_cookie("session", path="/") - return response - - -@router.get("/userinfo") -async def userinfo( - request: Request, - mymdc: MyMdCClient, - session: DBSession, - with_proposals: bool = True, -) -> models.User | models.OAuthUserInfo: - """User information.""" - if with_proposals: - user = await models.User.from_connection(request, mymdc, session) - else: - user = models.OAuthUserInfo.from_connection(request) - - return user - - -# ----------------------------------------------------------------------------- -# No-auth mode - -noauth_router = APIRouter(prefix="/oauth", tags=["auth"]) - - -@noauth_router.get("/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( - get_app_state(request).repositories - ) - if proposal_number: - proposals = {LOCAL_CYCLE: [proposal_number]} - - return {**models.DEV_USER.model_dump(), "proposals_by_year_half": proposals} - - -@noauth_router.post("/logout") -async def noauth_logout(request: Request): - request.session.pop("user", None) - response = JSONResponse(status_code=200, content={"logout_url": None}) - response.delete_cookie("session", path="/") - return response + return Response( + content={"logout_url": None}, + cookies=[Cookie(key=SESSION_COOKIE_KEY, value="", max_age=0, path="/")], + ) diff --git a/api/src/damnit_api/contextfile/routers.py b/api/src/damnit_api/contextfile/routers.py index 127028ec..f8205b52 100644 --- a/api/src/damnit_api/contextfile/routers.py +++ b/api/src/damnit_api/contextfile/routers.py @@ -1,34 +1,32 @@ -from typing import Annotated - from anyio import Path as APath -from fastapi import APIRouter, Depends +from litestar import Router, get +from litestar.di import Provide from ..metadata.models import ProposalMeta from ..metadata.routers import get_proposal_meta from . import models -router = APIRouter(prefix="/contextfile") - -@router.get("/content") -async def get_content( - proposal: Annotated[ProposalMeta, Depends(get_proposal_meta)], -) -> models.ContextFile | None: +@get("/content") +async def get_content(proposal: ProposalMeta) -> models.ContextFile | None: if proposal.damnit_path is None: return None - return await models.ContextFile.from_file( APath(proposal.damnit_path) / "context.py" ) -@router.get("/last_modified") -async def get_modified( - proposal: Annotated[ProposalMeta, Depends(get_proposal_meta)], -) -> models.ModifiedTime | None: +@get("/last_modified") +async def get_modified(proposal: ProposalMeta) -> models.ModifiedTime | None: if proposal.damnit_path is None: return None - return await models.ModifiedTime.from_file( APath(proposal.damnit_path) / "context.py" ) + + +router = Router( + path="/contextfile", + route_handlers=[get_content, get_modified], + dependencies={"proposal": Provide(get_proposal_meta)}, +) diff --git a/api/src/damnit_api/metadata/routers.py b/api/src/damnit_api/metadata/routers.py index 115318d4..3f006b50 100644 --- a/api/src/damnit_api/metadata/routers.py +++ b/api/src/damnit_api/metadata/routers.py @@ -1,20 +1,34 @@ """Metadata routers.""" -from fastapi import APIRouter +from litestar import Router, get +from litestar.di import Provide +from sqlmodel.ext.asyncio.session import AsyncSession -from .._db.dependencies import DBSession from .._mymdc.dependencies import MyMdCClient -from ..auth.dependencies import User +from ..auth.models import User from ..shared.models import ProposalNumber from . import services from .models import ProposalMeta -router = APIRouter(prefix="/metadata", tags=["metadata"]) - -@router.get("/proposal/{proposal_number}") async def get_proposal_meta( - proposal_number: ProposalNumber, mymdc: MyMdCClient, user: User, session: DBSession + proposal_number: ProposalNumber, + mymdc: MyMdCClient, + user: User, + session: AsyncSession, ) -> ProposalMeta: - """Get proposal metadata by proposal number.""" + """Dependency: resolve ProposalMeta from path/query parameter.""" return await services.get_proposal_meta(mymdc, proposal_number, user, session) + + +@get("/proposal/{proposal_number:int}", sync_to_thread=False) +def get_proposal(proposal_meta: ProposalMeta) -> ProposalMeta: + return proposal_meta + + +router = Router( + path="/metadata", + route_handlers=[get_proposal], + dependencies={"proposal_meta": Provide(get_proposal_meta)}, + tags=["metadata"], +) diff --git a/api/src/damnit_api/shared/settings.py b/api/src/damnit_api/shared/settings.py index 5fecc8aa..9aeae7f0 100644 --- a/api/src/damnit_api/shared/settings.py +++ b/api/src/damnit_api/shared/settings.py @@ -58,6 +58,13 @@ class Settings(BaseSettings): session_secret: SecretStr | None = None + # Whether the OAuth callback URL may be built from the `x-forwarded-host` + # header. Only safe behind a reverse proxy that sets (and strips any + # client-supplied copy of) this header — otherwise a client could spoof + # it to redirect the OAuth callback. Off by default; a single flag is + # enough since the app either sits behind one trusted proxy layer or none. + trust_forwarded_host: bool = False + uvicorn: UvicornSettings = UvicornSettings() @property From b4e8f835e89ee1825e2fcc4763f6dcaf8608a5e5 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 05:40:55 +0200 Subject: [PATCH 06/12] refactor(api/graphql): bind strawberry via litestar controller --- api/src/damnit_api/metadata/gql.py | 3 -- api/src/damnit_api/shared/gql.py | 45 ++++++++++++++---------------- 2 files changed, 21 insertions(+), 27 deletions(-) diff --git a/api/src/damnit_api/metadata/gql.py b/api/src/damnit_api/metadata/gql.py index a17b1880..43ab2e98 100644 --- a/api/src/damnit_api/metadata/gql.py +++ b/api/src/damnit_api/metadata/gql.py @@ -51,9 +51,6 @@ async def proposal_metadata( """ from ..shared.settings import settings - if info.context.request is None: - return None - if settings.is_local: return [ ProposalMeta.from_pydantic(services._local_proposal_meta(n)) diff --git a/api/src/damnit_api/shared/gql.py b/api/src/damnit_api/shared/gql.py index 8786c791..edbbd139 100644 --- a/api/src/damnit_api/shared/gql.py +++ b/api/src/damnit_api/shared/gql.py @@ -1,10 +1,8 @@ -from dataclasses import dataclass - import numpy as np import orjson import strawberry -from strawberry.fastapi import BaseContext, GraphQLRouter from strawberry.http import GraphQLHTTPResponse +from strawberry.litestar import BaseContext, make_graphql_controller from strawberry.schema.config import StrawberryConfig from strawberry.subscriptions import GRAPHQL_TRANSPORT_WS_PROTOCOL @@ -29,21 +27,6 @@ class Query(auth.Query, gql_main.queries.Query, metadata.Query): pass -class Router(GraphQLRouter): - def encode_json(self, data: GraphQLHTTPResponse) -> str | bytes: # pyright: ignore[reportIncompatibleMethodOverride] - encoded = orjson.dumps( - data, - default=lambda x: None if isinstance(x, float) and np.isnan(x) else x, - option=orjson.OPT_SERIALIZE_NUMPY | orjson.OPT_NON_STR_KEYS, - ) - - # WebSocket protocol messages are strings - if isinstance(data, dict) and "type" in data: - return encoded.decode("utf-8") - - return encoded - - class Schema(strawberry.Schema): pass @@ -52,7 +35,6 @@ class Subscription(gql_main.subscriptions.Subscription): pass -@dataclass(slots=True) class Context(BaseContext): mymdc: MyMdCClient oauth_user: OAuthUserInfo @@ -76,7 +58,7 @@ async def get_context( # noqa: RUF029 session: DBSession, repositories: Repositories, subscription_cursors: SubscriptionCursorsDep, -): +) -> Context: return Context( oauth_user=oauth_user, mymdc=mymdc, @@ -86,7 +68,7 @@ async def get_context( # noqa: RUF029 ) -def get_gql_app(): +def get_gql_controller() -> type: schema = Schema( query=Query, subscription=Subscription, @@ -98,8 +80,23 @@ def get_gql_app(): ), ) - return Router( + base_controller = make_graphql_controller( schema=schema, - subscription_protocols=SUBSCRIPTION_PROTOCOLS, - context_getter=get_context, # pyright: ignore[reportArgumentType] + path="/graphql", + context_getter=get_context, + subscription_protocols=tuple(SUBSCRIPTION_PROTOCOLS), ) + + class GraphQLController(base_controller): # ty: ignore[unsupported-base] + def encode_json(self, data: GraphQLHTTPResponse) -> str | bytes: # type: ignore[override] + encoded = orjson.dumps( + data, + default=lambda x: None if isinstance(x, float) and np.isnan(x) else x, + option=orjson.OPT_SERIALIZE_NUMPY | orjson.OPT_NON_STR_KEYS, + ) + # WebSocket protocol messages are strings + if isinstance(data, dict) and "type" in data: + return encoded.decode("utf-8") + return encoded + + return GraphQLController From 95dd8735c0dabe867cf487c33b603dc2f94663b1 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 05:41:16 +0200 Subject: [PATCH 07/12] feat(api): litestar app factory with lifespan-managed appstate --- api/src/damnit_api/_logging.py | 70 ++++++------ api/src/damnit_api/main.py | 189 +++++++++++++++++++-------------- 2 files changed, 146 insertions(+), 113 deletions(-) diff --git a/api/src/damnit_api/_logging.py b/api/src/damnit_api/_logging.py index c47dfca6..29f018d7 100644 --- a/api/src/damnit_api/_logging.py +++ b/api/src/damnit_api/_logging.py @@ -1,20 +1,16 @@ import inspect import logging import sys -from typing import TYPE_CHECKING import colorama import structlog import structlog.typing import ulid -from starlette.middleware.base import BaseHTTPMiddleware +from litestar.middleware.base import MiddlewareProtocol +from litestar.types import ASGIApp, Message, Receive, Scope, Send from structlog.dev import RichTracebackFormatter from structlog.stdlib import ProcessorFormatter -if TYPE_CHECKING: # pragma: no cover - from starlette.requests import Request - from starlette.responses import Response - def get_logger(logger_name: str | None = None): if logger_name: @@ -181,52 +177,58 @@ def configure_uvicorn(renderer, shared_processors): uvicorn.config.LOGGING_CONFIG["loggers"]["uvicorn.access"]["propagate"] = False -class RequestLoggingMiddleware(BaseHTTPMiddleware): - _logger = None +class RequestLoggingMiddleware(MiddlewareProtocol): + """Log requests and responses via structlog, replacing uvicorn access logs.""" + + def __init__(self, app: ASGIApp) -> None: + self.app = app + self._logger = None @property def logger(self): if not self._logger: self._logger = structlog.get_logger(logger_name="damnit_api.access_log") - return self._logger - async def dispatch(self, request: "Request", call_next) -> "Response": - """Add a middleware to FastAPI that will log requests and responses, - this is used instead of the builtin Uvicorn access logging to better - integrate with structlog""" + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] != "http": + await self.app(scope, receive, send) + return structlog.contextvars.bind_contextvars(request_id=str(ulid.ULID())) - info = { - "method": request.method, - "path": request.scope["path"], - "client": request.client, + info: dict = { + "method": scope["method"], + "path": scope["path"], + "client": scope.get("client"), } - if request.query_params: - info["query_params"] = str(request.query_params) + if query_string := scope.get("query_string"): + info["query_params"] = query_string.decode() - if request.path_params: - info["path_params"] = str(request.path_params) + if path_params := scope.get("path_params"): + info["path_params"] = str(path_params) logger = self.logger.bind() - logger.info("Request", **info) - response = await call_next(request) + async def send_wrapper(message: Message) -> None: + if message["type"] == "http.response.start": + status_code: int = message["status"] + + if status_code < 400: + response_logger = logger.info + elif status_code < 500: + response_logger = logger.warn + else: + response_logger = logger.error - if response.status_code < 400: - response_logger = logger.info - elif response.status_code < 500: - response_logger = logger.warn - else: - response_logger = logger.error + # Health checks are noisy, so we downgrade their log level + if scope["path"].endswith("/health"): + response_logger = logger.debug - # Health checks are noisy, so we downgrade their log level - if request.url.path.endswith("/health"): - response_logger = logger.debug + response_logger("Response", status_code=status_code) - response_logger("Response", status_code=response.status_code) + await send(message) - return response + await self.app(scope, receive, send_wrapper) diff --git a/api/src/damnit_api/main.py b/api/src/damnit_api/main.py index ffe97b8b..b6b7d6c4 100644 --- a/api/src/damnit_api/main.py +++ b/api/src/damnit_api/main.py @@ -1,21 +1,34 @@ from contextlib import asynccontextmanager -from fastapi import FastAPI, HTTPException, Request, status -from fastapi.responses import JSONResponse, RedirectResponse -from starlette.middleware.sessions import SessionMiddleware +from litestar import Litestar +from litestar.di import Provide +from litestar.exceptions import HTTPException -from ._logging import RequestLoggingMiddleware +from . import contextfile, metadata +from ._db.dependencies import get_session +from ._mymdc.dependencies import get_mymdc_client +from .auth.dependencies import get_oauth_user_info, get_user -# Known paths are redirected to the login page and -# then back after successful authentication. +# Known paths are redirected to the login page after a 401. KNOWN_PATHS = ["/graphql"] def create_app(): - from . import _logging, auth, contextfile, get_logger, metadata - from .shared import errors, gql + import hashlib + + from litestar import Request, Response + from litestar.middleware.session.client_side import CookieBackendConfig + from litestar.openapi import OpenAPIConfig + from litestar.response import Redirect + + from . import _logging, auth, get_logger + from .graphql.dependencies import get_subscription_cursors + from .runs.dependencies import get_repositories + from .shared.errors import DamnitWebError + from .shared.gql import get_gql_controller from .shared.settings import settings from .state import ( + SESSION_COOKIE_KEY, AppState, create_db_engine, create_db_sessionmaker, @@ -24,12 +37,50 @@ def create_app(): create_repositories, create_subscription_cursors, create_token_store, + provide_app_state, ) logger = get_logger("lifespan") + # ── Session middleware ──────────────────────────────────────────────────── + # Derive a 32-byte AES key from the session secret via SHA-256. + session_secret = settings.session_secret + assert session_secret is not None # enforced by Settings validator # noqa: S101 + session_config = CookieBackendConfig( + secret=hashlib.sha256(session_secret.get_secret_value().encode()).digest(), + key=SESSION_COOKIE_KEY, + ) + + # ── Exception handlers ──────────────────────────────────────────────────── + def dw_error_handler(request: Request, exc: DamnitWebError) -> Response: + code = getattr(exc, "code", None) or 500 + return Response( + content={ + "message": exc.message, + "details": exc.details, + "request_id": exc.request_id, + }, + status_code=code, + ) + + def unauthorized_handler( + request: Request, exc: HTTPException + ) -> Response | Redirect: + if exc.status_code == 401 and request.url.path in KNOWN_PATHS: + from urllib.parse import urlencode + + redirect_to = ( + f"/oauth/login?{urlencode({'redirect_uri': request.url.path})}" + ) + return Redirect(path=redirect_to, status_code=307) + return Response(content={"detail": exc.detail}, status_code=exc.status_code) + + # ── OpenAPI config ──────────────────────────────────────────────────────── + openapi_config = OpenAPIConfig(title="DAMNIT Web API", version="1.0.0") + + # ── Lifespan ────────────────────────────────────────────────────────────── @asynccontextmanager - async def lifespan(app: FastAPI): + async def lifespan(app: Litestar): _logging.configure( level=settings.log_level, debug=settings.debug, @@ -38,14 +89,15 @@ async def lifespan(app: FastAPI): logger.info("Starting application lifespan") + engine = create_db_engine(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), + db_engine=engine, + db_sessionmaker=create_db_sessionmaker(engine), mymdc_client=create_mymdc_client(settings), oauth_client=oauth_client, token_store=create_token_store(), @@ -53,76 +105,55 @@ async def lifespan(app: FastAPI): subscription_cursors=create_subscription_cursors(), ) - if settings.is_local: - app.router.include_router(auth.noauth_router) - else: - app.router.include_router(auth.router) - app.router.include_router(metadata.router) - app.router.include_router(contextfile.router) - app.router.include_router(gql.get_gql_app(), prefix="/graphql") - yield - - swagger_oauth = ( - None - if settings.auth is None - else { - "usePkceWithAuthorizationCodeGrant": True, - "clientId": settings.auth.client_id, - } - ) - app = FastAPI(lifespan=lifespan, swagger_ui_init_oauth=swagger_oauth) - - @app.exception_handler(HTTPException) - async def http_exception_handler(request: Request, exc: HTTPException): # noqa: RUF029 - request_path = request.url.path - if ( - not settings.is_local - and exc.status_code == status.HTTP_401_UNAUTHORIZED - and request_path in KNOWN_PATHS - ): - return RedirectResponse(url=f"/oauth/login?redirect_uri={request_path}") - return JSONResponse( - status_code=exc.status_code, - content={"detail": exc.detail}, - ) - - @app.exception_handler(errors.DamnitWebError) - async def base_exception_handler(request: Request, exc: errors.DamnitWebError): # noqa: RUF029 - status_code = exc.code or status.HTTP_500_INTERNAL_SERVER_ERROR - - content: dict[str, str | int | dict] = { - "message": exc.message, - "status_code": status_code, - } - - if exc.details: - content["details"] = exc.details - - if exc.request_id: - content["request_id"] = exc.request_id - - return JSONResponse( - status_code=status_code, - content=content, - ) + try: + yield + finally: + await engine.dispose() - app.add_middleware( - SessionMiddleware, - secret_key=settings.session_secret.get_secret_value(), # pyright: ignore[reportOptionalMemberAccess] + # ── Auth controller (mode-dependent, ADR-008 composition) ─────────────── + auth_controller = ( + auth.NoAuthOAuthController if settings.is_local else auth.OAuthController ) - app.add_middleware(RequestLoggingMiddleware) - - try: - from uvicorn.middleware.proxy_headers import ProxyHeadersMiddleware - - app.add_middleware( - ProxyHeadersMiddleware, trusted_hosts=["localhost", "127.0.0.1"] - ) - except Exception: - logger.warning("Could not add proxy headers middleware") - - return app + # ── GraphQL controller ──────────────────────────────────────────────────── + gql_controller = get_gql_controller() + + return Litestar( + route_handlers=[ + metadata.router, + contextfile.router, + auth_controller, + gql_controller, + ], + lifespan=[lifespan], + dependencies={ + "app_state": Provide(provide_app_state, sync_to_thread=False), + # Shared across all routes; resolved on demand. + "oauth_config": Provide( + auth.dependencies.get_oauth_client, sync_to_thread=False + ), + "token_store": Provide( + auth.dependencies.get_token_store, sync_to_thread=False + ), + "session": Provide(get_session), + "mymdc": Provide(get_mymdc_client, sync_to_thread=False), + "user": Provide(get_user), + "oauth_user": Provide(get_oauth_user_info, sync_to_thread=False), + "subscription_cursors": Provide( + get_subscription_cursors, sync_to_thread=False + ), + "repositories": Provide(get_repositories, sync_to_thread=False), + }, + middleware=[ + session_config.middleware, + _logging.RequestLoggingMiddleware, + ], + exception_handlers={ # ty: ignore[invalid-argument-type] + DamnitWebError: dw_error_handler, + HTTPException: unauthorized_handler, + }, + openapi_config=openapi_config, + ) if __name__ == "__main__": From 1e3104772a22d00a2ed631a06f477852a32699d6 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 05:42:17 +0200 Subject: [PATCH 08/12] build(api): drop fastapi and pin litestar after the swap --- api/pyproject.toml | 4 ++-- uv.lock | 58 ++++------------------------------------------ 2 files changed, 6 insertions(+), 56 deletions(-) diff --git a/api/pyproject.toml b/api/pyproject.toml index c8e57501..b6db2bb3 100644 --- a/api/pyproject.toml +++ b/api/pyproject.toml @@ -12,12 +12,11 @@ maintainers = [ dependencies = [ "pandas~=2.0", "sqlalchemy[asyncio]~=2.0", - "fastapi~=0.115", "orjson~=3.8", "numpy~=2.3", "matplotlib~=3.7", "h5py~=3.9", - "strawberry-graphql[fastapi,litestar]>=0.283.3", + "strawberry-graphql[litestar]>=0.283.3", "uvicorn[standard]~=0.29", "aiosqlite~=0.19", "scipy~=1.11", @@ -39,6 +38,7 @@ dependencies = [ "python-ulid[pydantic]>=3.1.0", "sqlmodel>=0.0.31", "pyyaml>=6.0.3", + "litestar~=2.24.0", ] [dependency-groups] diff --git a/uv.lock b/uv.lock index e30ad501..fa2a4fcb 100644 --- a/uv.lock +++ b/uv.lock @@ -81,15 +81,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/00/b7/e3bf5133d697a08128598c8d0abc5e16377b51465a33756de24fa7dee953/aiosqlite-0.22.1-py3-none-any.whl", hash = "sha256:21c002eb13823fad740196c5a2e9d8e62f6243bd9e7e4a1f87fb5e44ecb4fceb", size = 17405, upload-time = "2025-12-23T19:25:42.139Z" }, ] -[[package]] -name = "annotated-doc" -version = "0.0.4" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/57/ba/046ceea27344560984e26a590f90bc7f4a75b06701f653222458922b558c/annotated_doc-0.0.4.tar.gz", hash = "sha256:fbcda96e87e9c92ad167c2e53839e57503ecfda18804ea28102353485033faa4", size = 7288, upload-time = "2025-11-10T22:07:42.062Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/1e/d3/26bf1008eb3d2daa8ef4cacc7f3bfdc11818d111f7e2d0201bc6e3b49d45/annotated_doc-0.0.4-py3-none-any.whl", hash = "sha256:571ac1dc6991c450b25a9c2d84a3705e2ae7a53467b5d111c24fa8baabbed320", size = 5303, upload-time = "2025-11-10T22:07:40.673Z" }, -] - [[package]] name = "annotated-types" version = "0.7.0" @@ -395,11 +386,11 @@ dependencies = [ { name = "authlib" }, { name = "colorama" }, { name = "damnit" }, - { name = "fastapi" }, { name = "h5py" }, { name = "httpx" }, { name = "itsdangerous" }, { name = "ldap3" }, + { name = "litestar" }, { name = "matplotlib" }, { name = "numpy" }, { name = "orjson" }, @@ -412,7 +403,7 @@ dependencies = [ { name = "scipy" }, { name = "sqlalchemy", extra = ["asyncio"] }, { name = "sqlmodel" }, - { name = "strawberry-graphql", extra = ["fastapi", "litestar"] }, + { name = "strawberry-graphql", extra = ["litestar"] }, { name = "structlog" }, { name = "uvicorn", extra = ["standard"] }, { name = "xarray" }, @@ -472,11 +463,11 @@ requires-dist = [ { name = "authlib", specifier = "~=1.3" }, { name = "colorama", specifier = "~=0.4" }, { name = "damnit", specifier = "~=0.2.1" }, - { name = "fastapi", specifier = "~=0.115" }, { name = "h5py", specifier = "~=3.9" }, { name = "httpx", specifier = "~=0.27" }, { name = "itsdangerous", specifier = "~=2.1" }, { name = "ldap3", specifier = "~=2.9" }, + { name = "litestar", specifier = "~=2.24.0" }, { name = "matplotlib", specifier = "~=3.7" }, { name = "numpy", specifier = "~=2.3" }, { name = "orjson", specifier = "~=3.8" }, @@ -489,7 +480,7 @@ requires-dist = [ { name = "scipy", specifier = "~=1.11" }, { name = "sqlalchemy", extras = ["asyncio"], specifier = "~=2.0" }, { name = "sqlmodel", specifier = ">=0.0.31" }, - { name = "strawberry-graphql", extras = ["fastapi", "litestar"], specifier = ">=0.283.3" }, + { name = "strawberry-graphql", extras = ["litestar"], specifier = ">=0.283.3" }, { name = "structlog", specifier = "~=24.4" }, { name = "uvicorn", extras = ["standard"], specifier = "~=0.29" }, { name = "xarray", specifier = ">=2024.11" }, @@ -608,22 +599,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/92/a6/6111b9f13c1564b0e2f5dbeeb611fd00dc731e6add2336f62b798598c73d/faker-40.28.1-py3-none-any.whl", hash = "sha256:e8d3f5c469100a553d246dce7937c291308068a5ec6c9c3a228d7878b50720be", size = 2061052, upload-time = "2026-07-01T22:23:41.946Z" }, ] -[[package]] -name = "fastapi" -version = "0.139.0" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "annotated-doc" }, - { name = "pydantic" }, - { name = "starlette" }, - { name = "typing-extensions" }, - { name = "typing-inspection" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/d3/af/a5f50ccfa659ec1802cb4ca842c23f06d906a8cc9aef6016a2caeea3d4ed/fastapi-0.139.0.tar.gz", hash = "sha256:99ab7b2d92223c76d6cf10757ab3f89d45b38267fc20b2a136cf02f6beac3145", size = 423016, upload-time = "2026-07-01T16:35:33.436Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/9e/7c/8e3c6ad324ea5cb36604fc3f968554887891c316d9dfde57761611d907ad/fastapi-0.139.0-py3-none-any.whl", hash = "sha256:cf15e1e9e667ddb0ad63811e60bd11390d1aac838ca4a7a23f421807b2308189", size = 130339, upload-time = "2026-07-01T16:35:32.19Z" }, -] - [[package]] name = "filelock" version = "3.25.2" @@ -1706,15 +1681,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/0b/d7/1959b9648791274998a9c3526f6d0ec8fd2233e4d4acce81bbae76b44b2a/python_dotenv-1.2.2-py3-none-any.whl", hash = "sha256:1d8214789a24de455a8b8bd8ae6fe3c6b69a5e3d64aa8a8e5d68e694bbcb285a", size = 22101, upload-time = "2026-03-01T16:00:25.09Z" }, ] -[[package]] -name = "python-multipart" -version = "0.0.32" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/5b/42/55c32bb9b12693c092ad250a0e82edb5b31ddeda6eb772de5f308b3804ad/python_multipart-0.0.32.tar.gz", hash = "sha256:be54b7f3fa167bb83e4fcd936b887b708f4e57fe75911c02aebf53efaf8d938e", size = 46881, upload-time = "2026-06-04T16:18:58.647Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/e1/04/e8135ebd1ad02c56ec633277529b2602ff99ff634be76cdba5744cf554fd/python_multipart-0.0.32-py3-none-any.whl", hash = "sha256:ff6d3f776f16878c894e52e107296ffc890e913c611b1a4ec6c44e2821fe2e23", size = 30042, upload-time = "2026-06-04T16:18:57.319Z" }, -] - [[package]] name = "python-ulid" version = "3.1.0" @@ -1942,18 +1908,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/b1/e1/7c8d18e737433f3b5bbe27b56a9072a9fcb36342b48f1bef34b6da1d61f2/sqlmodel-0.0.37-py3-none-any.whl", hash = "sha256:2137a4045ef3fd66a917a7717ada959a1ceb3630d95e1f6aaab39dd2c0aef278", size = 27224, upload-time = "2026-02-21T16:39:47.781Z" }, ] -[[package]] -name = "starlette" -version = "1.3.1" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "anyio" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/eb/e3/7c1dc7381d9f8ab7d854328ebfa884e62cb3f3d8549ddfd37c7814f42afa/starlette-1.3.1.tar.gz", hash = "sha256:05d0213193f2fbaae60e2ecb593b4add4262ad4e46536b54abe36f11a71724e0", size = 2703240, upload-time = "2026-06-12T09:23:11.602Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/ec/bb/2799cc2ede3ed41131f8975621e7213dfc7ef4acbbaadfa440f32500c370/starlette-1.3.1-py3-none-any.whl", hash = "sha256:c7372aae11c3c3f26a42df7bd626cec2f47d03483d261d369516a615a53714c6", size = 73632, upload-time = "2026-06-12T09:23:10.017Z" }, -] - [[package]] name = "strawberry-graphql" version = "0.312.2" @@ -1971,10 +1925,6 @@ wheels = [ ] [package.optional-dependencies] -fastapi = [ - { name = "fastapi" }, - { name = "python-multipart" }, -] litestar = [ { name = "litestar" }, ] From d7dc32a61d0fb68a4f9e538f23ac5be1c285c243 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 05:43:40 +0200 Subject: [PATCH 09/12] test(api): adapt and extend tests for the litestar swap --- api/tests/refactor/e2e/conftest.py | 36 ++- api/tests/refactor/e2e/test_authz_parity.py | 17 +- api/tests/test_auth_routers.py | 302 ++++++++++++++++++++ api/tests/test_contextfile.py | 69 +++-- api/tests/test_errors.py | 40 ++- api/tests/test_state.py | 37 ++- 6 files changed, 423 insertions(+), 78 deletions(-) create mode 100644 api/tests/test_auth_routers.py diff --git a/api/tests/refactor/e2e/conftest.py b/api/tests/refactor/e2e/conftest.py index cc4cc64a..7e2c89e3 100644 --- a/api/tests/refactor/e2e/conftest.py +++ b/api/tests/refactor/e2e/conftest.py @@ -5,9 +5,9 @@ This conftest is the only place allowed to touch the app, and only via: - the app factory (`damnit_api.main.create_app`), imported lazily inside a fixture -- `mint_session_cookie`, which forges the signed session cookie OAuth callback would set +- `mint_session_cookie`, which forges the session cookie OAuth callback would set - This helper is deliberately coupled to the session implementation - (Starlette `SessionMiddleware` over `itsdangerous`) + (Litestar's client-side `CookieBackendConfig`) - The framework-swap branch must update this one helper and nothing else Outbound HTTP is recorded and replayed with pytest-recording. Cassettes stored in @@ -18,15 +18,12 @@ headers and OAuth client credentials at record time, but you should still check. """ -import base64 -import json from pathlib import Path import httpx import pytest import pytest_asyncio from asgi_lifespan import LifespanManager -from itsdangerous import TimestampSigner SESSION_COOKIE = "session" @@ -142,16 +139,33 @@ async def logged_in_client(e2e_client): # noqa: RUF029 - must be async to recei def mint_session_cookie(user: dict | None = None) -> dict[str, str]: - """Forge the signed session cookie a completed OAuth login would set. + """Forge the session cookie a completed OAuth login would set. !!! warning This bypasses awkward OAuth internals (redirects, callbacks, etc...) while - keeping the key real session path: value is signed with the app's own session - secret, same as how as Starlette's `SessionMiddleware` does. + keeping the real session path: the value is encrypted with the app's own + session secret, exactly as Litestar's client-side `CookieBackendConfig` does + (see `main.py`). """ + import hashlib + + from litestar.middleware.session.client_side import ( + ClientSideSessionBackend, + CookieBackendConfig, + ) + from damnit_api.shared.settings import settings - secret = settings.session_secret.get_secret_value() # pyright: ignore[reportOptionalMemberAccess] - payload = base64.b64encode(json.dumps({"user": user or TEST_USER}).encode("utf-8")) - return {SESSION_COOKIE: TimestampSigner(secret).sign(payload).decode("utf-8")} + secret = settings.session_secret.get_secret_value() # ty: ignore[unresolved-attribute] # pyright: ignore[reportOptionalMemberAccess] + config = CookieBackendConfig( + secret=hashlib.sha256(secret.encode()).digest(), key=SESSION_COOKIE + ) + backend = ClientSideSessionBackend(config=config) + chunks = backend.dump_data({"user": user or TEST_USER}) + if len(chunks) == 1: + return {SESSION_COOKIE: chunks[0].decode("utf-8")} + return { + f"{SESSION_COOKIE}-{i}": chunk.decode("utf-8") + for i, chunk in enumerate(chunks) + } diff --git a/api/tests/refactor/e2e/test_authz_parity.py b/api/tests/refactor/e2e/test_authz_parity.py index 800fe262..432ca271 100644 --- a/api/tests/refactor/e2e/test_authz_parity.py +++ b/api/tests/refactor/e2e/test_authz_parity.py @@ -77,14 +77,17 @@ async def test_member_passes_authorization_unchanged(logged_in_client): assert payload["data"]["runs"] -async def test_graphql_query_without_session_unchanged(e2e_client): - """Today an unauthenticated GraphQL request is NOT turned into an HTTP error. +async def test_graphql_query_without_session_rejected(e2e_client): + """An unauthenticated GraphQL request is rejected, not served. + + The session lookup raises `ValueError` ("No user info in session") while + building the GraphQL context. Litestar's exception handling turns that into + a 500 response; under FastAPI the same error propagated unhandled through the + raw ASGI transport instead. !!! todo - Currently session-lookup `ValueError` is unhandled (surfaced here by raw ASGI - transport). This should become a proper 401 error, test and must be updated - this test when it does. + This should become a proper 401 error; update this test when it does. """ - with pytest.raises(ValueError, match="No user info in session"): - await e2e_client.post("/graphql", json=runs_query(MEMBER_PROPOSAL)) + response = await e2e_client.post("/graphql", json=runs_query(MEMBER_PROPOSAL)) + assert response.status_code == 500 diff --git a/api/tests/test_auth_routers.py b/api/tests/test_auth_routers.py new file mode 100644 index 00000000..2fd0a9f8 --- /dev/null +++ b/api/tests/test_auth_routers.py @@ -0,0 +1,302 @@ +"""Tests for the OAuth2 flow in `auth/routers.py`.""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from litestar.di import Provide +from litestar.middleware.session.client_side import CookieBackendConfig +from litestar.testing import create_test_client + +from damnit_api.auth.routers import OAuthController +from damnit_api.auth.token_store import InMemoryTokenStore +from damnit_api.state import SESSION_COOKIE_KEY, OAuthClient + +SERVER_METADATA = { + "authorization_endpoint": "https://idp.example/authorize", + "token_endpoint": "https://idp.example/token", + "userinfo_endpoint": "https://idp.example/userinfo", + "revocation_endpoint": "https://idp.example/revoke", + "end_session_endpoint": "https://idp.example/logout", +} + +USERINFO = { + "sub": "user-1", + "email": "user@example.com", + "family_name": "User", + "given_name": "Test", + "groups": [], + "name": "Test User", + "preferred_username": "tuser", +} + + +def _oauth_config() -> OAuthClient: + return OAuthClient( + client_id="test-client", + client_secret="test-secret", # noqa: S106 + scope="openid email groups", + server_metadata_url="https://idp.example/.well-known/openid-configuration", + server_metadata=dict(SERVER_METADATA), + ) + + +@pytest.fixture +def token_store(): + return InMemoryTokenStore() + + +@pytest.fixture +def session_config(): + return CookieBackendConfig(secret=b"0" * 32, key=SESSION_COOKIE_KEY) + + +@pytest.fixture +def client(token_store, session_config): + # oauth_config/token_store are app-level dependencies in the composition + # root (main.py); provide fakes at the same layer here. + with create_test_client( + route_handlers=[OAuthController], + dependencies={ + "oauth_config": Provide(_oauth_config, sync_to_thread=False), + "token_store": Provide(lambda: token_store, sync_to_thread=False), + }, + session_config=session_config, + middleware=[session_config.middleware], + ) as c: + yield c + + +@pytest.fixture +def mock_oauth_client(): + """Patch authlib's AsyncOAuth2Client so no real HTTP is exchanged.""" + with patch("authlib.integrations.httpx_client.AsyncOAuth2Client") as cls: + instance = MagicMock() + instance.create_authorization_url = MagicMock( + return_value=("https://idp.example/authorize?...", "csrf-state") + ) + instance.fetch_token = AsyncMock( + return_value={ + "access_token": "tok123", + "refresh_token": "ref456", + "id_token": "idtok789", + } + ) + userinfo_resp = MagicMock() + userinfo_resp.raise_for_status = MagicMock() + userinfo_resp.json = MagicMock(return_value=dict(USERINFO)) + instance.get = AsyncMock(return_value=userinfo_resp) + instance.post = AsyncMock() + instance.aclose = AsyncMock() + cls.return_value = instance + yield instance + + +# ── /oauth/callback ────────────────────────────────────────────────────────── + + +def test_callback_state_mismatch_returns_401(client): + client.set_session_data({"_oauth_state": "correct-state"}) + + resp = client.get("/oauth/callback", params={"state": "wrong-state", "code": "abc"}) + + assert resp.status_code == 401 + + +def test_callback_success_stores_session_user_and_token( + client, token_store, mock_oauth_client +): + client.set_session_data({"_oauth_state": "csrf-state"}) + + resp = client.get( + "/oauth/callback", + params={"state": "csrf-state", "code": "abc"}, + follow_redirects=False, + ) + + assert resp.status_code == 302 + assert resp.headers["location"] == "/app/home" + + session = client.get_session_data() + assert session["user"]["sub"] == "user-1" + + assert token_store.pop_token_field("user-1", "access_token") == "tok123" + assert token_store.pop_token_field("user-1", "refresh_token") == "ref456" + + +# ── /oauth/logout ──────────────────────────────────────────────────────────── + + +def test_logout_revokes_tokens_and_clears_session_cookie( + client, token_store, mock_oauth_client +): + token_store.store( + "user-1", + {"access_token": "tok123", "refresh_token": "ref456", "id_token": "idtok789"}, + ) + client.set_session_data({"user": {"sub": "user-1"}}) + + resp = client.post("/oauth/logout") + + assert resp.status_code == 201 + assert resp.json()["logout_url"].startswith("https://idp.example/logout") + + # Both refresh and access tokens are revoked via the TokenStore. + assert mock_oauth_client.post.await_count == 2 + assert token_store.pop_token_field("user-1", "access_token") is None + assert token_store.pop_token_field("user-1", "refresh_token") is None + + # The session cookie is cleared, keyed by the shared session cookie name. + set_cookie = resp.headers.get("set-cookie", "") + assert f"{SESSION_COOKIE_KEY}=" in set_cookie + assert "Max-Age=0" in set_cookie + + assert "user" not in client.get_session_data() + + +# ── x-forwarded-host trust gate ────────────────────────────────────────────── + + +def test_forwarded_host_ignored_when_untrusted(client, mock_oauth_client): + resp = client.get( + "/oauth/login", + params={"redirect_uri": "/app/home"}, + headers={"x-forwarded-host": "evil.example"}, + follow_redirects=False, + ) + + assert resp.status_code == 302 + callback_uri = mock_oauth_client.create_authorization_url.call_args.kwargs[ + "redirect_uri" + ] + assert "evil.example" not in callback_uri + + +def test_forwarded_host_honoured_when_trusted(client, mock_oauth_client, monkeypatch): + from damnit_api.shared import settings as settings_module + + monkeypatch.setattr(settings_module.settings, "trust_forwarded_host", True) + + resp = client.get( + "/oauth/login", + params={"redirect_uri": "/app/home"}, + headers={"x-forwarded-host": "proxy.example"}, + follow_redirects=False, + ) + + assert resp.status_code == 302 + callback_uri = mock_oauth_client.create_authorization_url.call_args.kwargs[ + "redirect_uri" + ] + assert "proxy.example" in callback_uri + + +# ── post-login redirect allow-list ─────────────────────────────────────────── + + +def test_login_rejects_absolute_redirect_target(client, mock_oauth_client): + client.set_session_data({"user": dict(USERINFO)}) + + resp = client.get( + "/oauth/login", + params={"redirect_uri": "https://evil.example/phish"}, + follow_redirects=False, + ) + + assert resp.status_code == 302 + assert resp.headers["location"] == "/app/home" + + +def test_login_rejects_protocol_relative_redirect_target(client, mock_oauth_client): + client.set_session_data({"user": dict(USERINFO)}) + + resp = client.get( + "/oauth/login", + params={"redirect_uri": "//evil.example/phish"}, + follow_redirects=False, + ) + + assert resp.status_code == 302 + assert resp.headers["location"] == "/app/home" + + +def test_relative_redirect_carried_through_login_and_callback( + client, token_store, mock_oauth_client +): + resp = client.get( + "/oauth/login", + params={"redirect_uri": "/proposal/1234"}, + follow_redirects=False, + ) + + assert resp.status_code == 302 # off to the IdP + assert client.get_session_data()["_login_redirect"] == "/proposal/1234" + + resp = client.get( + "/oauth/callback", + params={"state": "csrf-state", "code": "abc"}, + follow_redirects=False, + ) + + assert resp.status_code == 302 + assert resp.headers["location"] == "/proposal/1234" + + +def test_callback_sanitizes_redirect_carried_in_session( + client, token_store, mock_oauth_client +): + client.set_session_data( + { + "_oauth_state": "csrf-state", + "_login_redirect": "https://evil.example/phish", + } + ) + + resp = client.get( + "/oauth/callback", + params={"state": "csrf-state", "code": "abc"}, + follow_redirects=False, + ) + + assert resp.status_code == 302 + assert resp.headers["location"] == "/app/home" + + +# ── websocket session auth (websockets authenticate identically to HTTP) ───── + + +def _ws_echo_user_handler(): + from litestar import websocket + from litestar.connection import WebSocket + + from damnit_api.auth.models import OAuthUserInfo + + @websocket("/ws") + async def ws_handler(socket: WebSocket) -> None: + await socket.accept() + user = OAuthUserInfo.from_connection(socket) + await socket.send_json({"email": user.email}) + await socket.close() + + return ws_handler + + +def test_websocket_resolves_user_from_session(session_config): + with create_test_client( + route_handlers=[_ws_echo_user_handler()], + session_config=session_config, + middleware=[session_config.middleware], + ) as c: + c.set_session_data({"user": dict(USERINFO)}) + with c.websocket_connect("/ws") as ws: + assert ws.receive_json() == {"email": "user@example.com"} + + +def test_websocket_without_session_user_is_rejected(session_config): + from litestar.exceptions import WebSocketDisconnect + + with create_test_client( + route_handlers=[_ws_echo_user_handler()], + session_config=session_config, + middleware=[session_config.middleware], + ) as c, pytest.raises(WebSocketDisconnect), c.websocket_connect("/ws") as ws: + ws.receive_json() diff --git a/api/tests/test_contextfile.py b/api/tests/test_contextfile.py index fe5fa51d..e1fd0cd5 100644 --- a/api/tests/test_contextfile.py +++ b/api/tests/test_contextfile.py @@ -1,36 +1,45 @@ import asyncio import time -from types import SimpleNamespace import pytest -from fastapi.testclient import TestClient +from litestar import Router +from litestar.di import Provide +from litestar.testing import create_test_client from damnit_api.contextfile import models -from damnit_api.main import create_app -from damnit_api.metadata.routers import get_proposal_meta - - -@pytest.fixture -def app(monkeypatch): - monkeypatch.setenv("DW_API_AUTH__CLIENT_ID", "test") - monkeypatch.setenv("DW_API_AUTH__CLIENT_SECRET", "test") - monkeypatch.setenv( - "DW_API_AUTH__SERVER_METADATA_URL", "https://example.com/.well-known" +from damnit_api.contextfile.routers import get_content, get_modified +from damnit_api.metadata.models import ProposalMeta + + +def _stub_proposal(damnit_path: str) -> ProposalMeta: + """Create a minimal ProposalMeta for tests (transient, no DB required).""" + return ProposalMeta( + number=1, + cycle="202401", + instrument="TEST", + path="/fake", + title="Test", + principal_investigator="Test PI", + start_date=None, + end_date=None, + updated_at=None, + damnit_path=damnit_path, ) - monkeypatch.setenv("DW_API_SESSION_SECRET", "test") - - # 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 - app.dependency_overrides.clear() @pytest.fixture -def client(app): - with TestClient(app) as c: +def client(temp_dir): + test_router = Router( + path="/contextfile", + route_handlers=[get_content, get_modified], + dependencies={ + "proposal": Provide( + lambda: _stub_proposal(str(temp_dir)), + sync_to_thread=False, + ) + }, + ) + with create_test_client(route_handlers=[test_router]) as c: yield c @@ -49,11 +58,8 @@ def clear_cache(): @pytest.mark.asyncio -async def test_watcher_detects_change(app, client, temp_dir): +async def test_watcher_detects_change(client, temp_dir): temp_path = temp_dir / "context.py" - app.dependency_overrides[get_proposal_meta] = lambda: SimpleNamespace( - damnit_path=str(temp_dir) - ) resp = client.get("/contextfile/last_modified") assert resp.status_code == 200 @@ -64,15 +70,8 @@ async def test_watcher_detects_change(app, client, temp_dir): assert await wait_for_change(client, "/contextfile/last_modified", initial_modified) -# FastAPI's TestClient creates a fresh event loop per request, so any -# alru_cached helper called by the handler binds to that loop. The -# loop-reset warning is intrinsic to this testing pattern. @pytest.mark.filterwarnings("ignore::async_lru.AlruCacheLoopResetWarning") -def test_file_fetching(app, client, temp_dir): - app.dependency_overrides[get_proposal_meta] = lambda: SimpleNamespace( - damnit_path=str(temp_dir) - ) - +def test_file_fetching(client): resp = client.get("/contextfile/content") assert resp.status_code == 200 assert resp.json()["fileContent"] == "initial content" diff --git a/api/tests/test_errors.py b/api/tests/test_errors.py index 05b0e908..7b906cd5 100644 --- a/api/tests/test_errors.py +++ b/api/tests/test_errors.py @@ -2,7 +2,8 @@ import pytest import structlog -from fastapi.testclient import TestClient +from litestar import get +from litestar.testing import TestClient from damnit_api.main import create_app from damnit_api.shared.errors import ( @@ -75,31 +76,19 @@ def test_request_id_none_when_not_bound(): # ----------------------------------------------------------------------------- -# FastAPI exception handler +# Litestar exception handler @pytest.fixture def app(monkeypatch): - monkeypatch.setenv("DW_API_AUTH__CLIENT_ID", "test") - monkeypatch.setenv("DW_API_AUTH__CLIENT_SECRET", "test") - monkeypatch.setenv( - "DW_API_AUTH__SERVER_METADATA_URL", "https://example.com/.well-known" - ) - monkeypatch.setenv("DW_API_SESSION_SECRET", "test") + # AppState is built in the lifespan; stub the only startup step that + # does network I/O (OIDC discovery). + async def noop_load(self): + pass - # 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) + monkeypatch.setattr("damnit_api.state.OAuthClient.load_server_metadata", noop_load) - app = create_app() - yield app - app.dependency_overrides.clear() - - -@pytest.fixture -def client(app): - with TestClient(app, raise_server_exceptions=False) as c: - yield c + return create_app() @pytest.mark.parametrize( @@ -110,13 +99,16 @@ def client(app): (UpstreamServiceError, 502), ], ) -def test_handler_maps_dwerror_to_status_code(app, client, exc_class, expected_status): - @app.get("/__test_raise__") - def _raise(): +def test_handler_maps_dwerror_to_status_code(app, exc_class, expected_status): + @get("/__test_raise__", sync_to_thread=False) + def _raise() -> None: msg = "boom" raise exc_class(msg, details="extra") - resp = client.get("/__test_raise__") + app.register(_raise) + + with TestClient(app) as client: + resp = client.get("/__test_raise__") assert resp.status_code == expected_status body = resp.json() diff --git a/api/tests/test_state.py b/api/tests/test_state.py index 59931daf..0f7dbd02 100644 --- a/api/tests/test_state.py +++ b/api/tests/test_state.py @@ -2,12 +2,15 @@ import ast from pathlib import Path +from unittest.mock import MagicMock + +import pytest from damnit_api.auth.token_store import InMemoryTokenStore from damnit_api.runs.repository import DamnitRepositoryRegistry from damnit_api.shared.models import ProposalNumber from damnit_api.shared.settings import Settings -from damnit_api.state import create_oauth_client +from damnit_api.state import create_mymdc_client, create_oauth_client def test_appstate_only_imported_by_composition_root(): @@ -39,6 +42,29 @@ def test_create_oauth_client_returns_none_when_auth_disabled(tmp_path): assert create_oauth_client(settings) is None +def test_create_oauth_client_returns_populated_client_when_auth_set(): + # The repo's `.env` provides valid auth settings; non-local mode requires them. + settings = Settings() + assert settings.auth is not None + client = create_oauth_client(settings) + assert client is not None + assert client.client_id == settings.auth.client_id + assert client.scope == "openid email groups" + + +def test_create_mymdc_client_raises_on_unsupported_config(tmp_path): + settings = Settings(damnit_path=tmp_path) + # Neither MyMdCHTTPSettings nor MyMdCMockSettings. + object.__setattr__(settings, "mymdc", MagicMock()) + with pytest.raises(ValueError, match="Invalid MyMdC configuration"): + create_mymdc_client(settings) + + +def test_create_mymdc_client_builds_mock_client(tmp_path): + settings = Settings(damnit_path=tmp_path) # default mymdc is the mock backend + assert create_mymdc_client(settings) is not None + + def test_repository_registry_memoizes_per_proposal(): created = [] @@ -61,3 +87,12 @@ def test_token_store_stores_and_pops_fields(): 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 + + +def test_token_store_drops_entry_once_empty(): + store = InMemoryTokenStore() + store.store("sub-1", {"access_token": "a"}) + + assert store.pop_token_field("sub-1", "access_token") == "a" + # The now-empty entry is removed, not left as a dangling empty dict. + assert "sub-1" not in store._tokens From b33dd578365b3f0a88fba3bae32e0a1ba2d2efdf Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 05:48:06 +0200 Subject: [PATCH 10/12] docs(api/adr): add ADR-006 litestar and sweep back-references --- .../adr/000-vertical-slice-architecture.md | 2 +- api/docs/adr/006-litestar.md | 37 +++++++++++++++++++ 2 files changed, 38 insertions(+), 1 deletion(-) create mode 100644 api/docs/adr/006-litestar.md diff --git a/api/docs/adr/000-vertical-slice-architecture.md b/api/docs/adr/000-vertical-slice-architecture.md index 42b8a520..b27e7b62 100644 --- a/api/docs/adr/000-vertical-slice-architecture.md +++ b/api/docs/adr/000-vertical-slice-architecture.md @@ -104,7 +104,7 @@ damnit_api/ - Feature packages (`runs`, `proposals`, `auth`, `contextfile`) may import `core` and infrastructure (`mymdc`, `appdb`), never each other's internals. Allowed cross-feature edges are explicit and narrow: `auth → proposals` (membership needs proposal metadata) - never the reverse. - Infrastructure (`mymdc`, `appdb`) imports only `core` and `settings`. - `graphql/schema.py` and `app.py` may import everything (composition). -- Domain and service modules never import Litestar or Strawberry; framework types appear only in `routers.py`, `gql.py`, `dependencies.py`, and permission classes. +- Domain and service modules never import Litestar (see [ADR-006](006-litestar.md)) or Strawberry; framework types appear only in `routers.py`, `gql.py`, `dependencies.py`, and permission classes. - Private (`_`-prefixed) functions are module-internal. Anything imported across module boundaries is public API and named accordingly. - Function-body imports are allowed only in the composition root and for documented, cycle-free lazy loading. diff --git a/api/docs/adr/006-litestar.md b/api/docs/adr/006-litestar.md new file mode 100644 index 00000000..1ca57b49 --- /dev/null +++ b/api/docs/adr/006-litestar.md @@ -0,0 +1,37 @@ +--- +date: 2026-07-08 +--- + +# ADR-006 - Web framework: Litestar + +## Context and Problem Statement + +The API needs an async Python web framework. It has to provide REST routing, ASGI websockets for GraphQL subscriptions, session middleware, dependency injection, and OpenAPI generation. + +Two requirements discriminate between candidates. Runtime dependencies are built once at startup and injected into handlers, so the framework must support typed application state with lifespan management and dependency declaration that is not welded to route signatures ([ADR-002](002-no-global-mutable-state.md)). The web framework must also stay at the edges of the codebase, so that domain and service code never imports it ([ADR-000](000-vertical-slice-architecture.md)). + +## Considered Options + +- FastAPI (the status quo). +- Litestar. + +## Decision Outcome + +Chosen option: "Litestar", because it provides typed `State`, layered `Provide`-based dependency injection, lifespan context managers, native session middleware, and an official Strawberry integration. + +FastAPI satisfies the basics. Its dependency injection is expressed per route through `Depends` in signatures, its application state is an untyped `app.state` namespace, and its authlib OAuth integration is Starlette-specific. Litestar avoids each of these. + +### Consequences + +- Good: application state is typed and lifespan-managed, and dependencies are declared off the route signatures. +- Good: FastAPI and Starlette are no longer dependencies. +- Bad: the OAuth flow is implemented natively rather than through a framework integration, so it needs its own tests and security review. +- Bad: Litestar has a smaller ecosystem than FastAPI. + +## Details + +Framework types (`Request`, `ASGIConnection`, `Provide`) stay at the edges: route handlers, dependency providers, and permission classes. The narrow per-slice providers read Litestar's injected `State` and return a single `AppState` attribute, rather than declaring an `AppState` parameter, because `AppState`'s fields are `TYPE_CHECKING`-only forward references and only the composition root may import it ([ADR-002](002-no-global-mutable-state.md)). + +Dependency injection resolves by parameter name against the `Provide` map, not by a `Depends` default, so the old `Annotated[T, Depends(...)]` aliases collapse to plain type aliases. Injected collaborators whose type is a union of concrete classes are annotated with `SkipValidation`, because Litestar's msgspec-based signature validation cannot build a decoder for a union of custom types. + +FastAPI's `@app.exception_handler` decorators become a Litestar `exception_handlers` mapping. The 401-to-login redirect for known paths is preserved inside the `HTTPException` handler. The uvicorn proxy-headers middleware is dropped in favour of a `trust_forwarded_host` setting that gates use of the `x-forwarded-host` header in the OAuth callback URL only. From 7a465848418580754ff38e058f0edf98d628e17f Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 05:50:47 +0200 Subject: [PATCH 11/12] docs(api/adr): add ADR-007 graphql-transport-only and sweep refs --- .../adr/000-vertical-slice-architecture.md | 2 +- api/docs/adr/007-graphql-transport-only.md | 38 +++++++++++++++++++ api/docs/architecture.md | 2 +- 3 files changed, 40 insertions(+), 2 deletions(-) create mode 100644 api/docs/adr/007-graphql-transport-only.md diff --git a/api/docs/adr/000-vertical-slice-architecture.md b/api/docs/adr/000-vertical-slice-architecture.md index b27e7b62..e4bde4d6 100644 --- a/api/docs/adr/000-vertical-slice-architecture.md +++ b/api/docs/adr/000-vertical-slice-architecture.md @@ -88,7 +88,7 @@ damnit_api/ ├── mymdc/ # MyMdC port ├── appdb/ # application-DB engine/session/models │ -└── graphql/ # transport composition: +└── graphql/ # transport composition only (ADR-007): ├── schema.py # assemble Query/Subscription from feature gql modules └── directives.py ``` diff --git a/api/docs/adr/007-graphql-transport-only.md b/api/docs/adr/007-graphql-transport-only.md new file mode 100644 index 00000000..461110ec --- /dev/null +++ b/api/docs/adr/007-graphql-transport-only.md @@ -0,0 +1,38 @@ +--- +date: 2026-07-08 +--- + +# ADR-007 - GraphQL as a transport layer; per-feature schema contributions + +## Context and Problem Statement + +GraphQL is the API's primary query surface. Left unmanaged, a GraphQL layer attracts logic that belongs elsewhere: domain serialisation ends up inside type definitions, resolvers accumulate data-access code, and one schema module becomes a central coupling point that imports from the whole codebase. + +Two structural questions need stable answers. Who owns the types and resolvers? With vertical slices each feature owns its domain, so its GraphQL surface belongs to that slice, not to a central package ([ADR-000](000-vertical-slice-architecture.md)). What does the composition layer do? Something must assemble the feature contributions into one schema, configure scalars and naming, and bind the schema to the web framework. + +A hard external constraint applies. The frontend depends on the public schema: snake_case field names, scalar names, and subscription payload shapes. Internal restructuring must not change that schema. + +## Considered Options + +- A central `graphql` package that owns all types and resolvers. +- Feature-owned GraphQL surfaces, with the composition layer assembling them. + +## Decision Outcome + +Chosen option: "feature-owned GraphQL surfaces", because it keeps each slice's GraphQL surface inside the slice and reduces the shared layer to composition. + +- Each feature exposes a `gql.py` with its Strawberry types and its `Query`/`Subscription` contributions. +- The composition layer merges those contributions, registers scalars, sets `StrawberryConfig(auto_camel_case=False)`, builds the framework controller, and defines the request `Context`. +- The `Context` is a typed object built from injected dependencies; resolvers reach collaborators only through `info.context`, never module imports. +- Serialisation is a domain concern, not a type concern: it lives in framework-free feature modules, and Strawberry types are thin. +- Resolvers are orchestration only: apply permission classes, fetch through the repository ([ADR-005](005-repository-pattern.md)), convert through serialisation, and raise `DamnitWebError` subclasses ([ADR-001](001-error-classes.md)). + +### Consequences + +- Good: the composition layer imports features; features never import it back, the narrow exception being type-only `Context` annotations under `TYPE_CHECKING`. +- Good: serialisation is unit-testable without Strawberry, and the GraphQL layer is testable against the CSV repository ([ADR-005](005-repository-pattern.md)). +- Bad: the public schema is frozen, so any change to it is deliberate and coordinated with the frontend. + +## Details + +Pushing a sub-selection down to the data layer is legitimate resolver logic rather than leaked domain code: shaping a `variables(names:)` selection via `info.selected_fields` is genuinely about the transport. The frozen-schema guard regenerates its snapshot only through an explicit environment flag, so an accidental schema change fails the parity test rather than silently updating the golden file. diff --git a/api/docs/architecture.md b/api/docs/architecture.md index c290efba..ff00ee13 100644 --- a/api/docs/architecture.md +++ b/api/docs/architecture.md @@ -22,7 +22,7 @@ For more information, see [ADR-000](adr/000-vertical-slice-architecture.md). | `proposals/` | Proposal metadata and lookup | Proposal models, MyMdC-backed metadata services, path locator (see [ADR-004](adr/004-proposal-path-locator.md)) | Planned | `metadata/` | | `auth/` | Authentication and authorisation | OAuth flow, sessions, token store, `User`, permission classes, the membership policy | Partial | Policy still in `metadata/services.py` | | `contextfile/` | Context-file viewing | File reading, watching, its routes | Done | As-is | -| `graphql/` | GraphQL transport only | Schema assembly, context, directives, controller binding - no resolvers, no domain logic | Partial | Assembly still in `shared/gql.py`; resolvers still here | +| `graphql/` | GraphQL transport only | Schema assembly, context, directives, controller binding - no resolvers, no domain logic (see [ADR-007](adr/007-graphql-transport-only.md)) | Partial | Assembly still in `shared/gql.py`; resolvers still here | | `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 (see [ADR-004](adr/004-proposal-path-locator.md)), converters | Planned | `shared/` + `utils.py` | From 500321df018a635991a5143e8574bf7c0d48c792 Mon Sep 17 00:00:00 2001 From: Robert Rosca <32569096+RobertRosca@users.noreply.github.com> Date: Wed, 8 Jul 2026 05:58:30 +0200 Subject: [PATCH 12/12] build(api): drop unused itsdangerous dependency --- api/pyproject.toml | 1 - uv.lock | 11 ----------- 2 files changed, 12 deletions(-) diff --git a/api/pyproject.toml b/api/pyproject.toml index b6db2bb3..813a1bf1 100644 --- a/api/pyproject.toml +++ b/api/pyproject.toml @@ -21,7 +21,6 @@ dependencies = [ "aiosqlite~=0.19", "scipy~=1.11", "authlib~=1.3", - "itsdangerous~=2.1", "httpx~=0.27", "pydantic~=2.12", "pydantic-settings~=2.2", diff --git a/uv.lock b/uv.lock index fa2a4fcb..2476189f 100644 --- a/uv.lock +++ b/uv.lock @@ -388,7 +388,6 @@ dependencies = [ { name = "damnit" }, { name = "h5py" }, { name = "httpx" }, - { name = "itsdangerous" }, { name = "ldap3" }, { name = "litestar" }, { name = "matplotlib" }, @@ -465,7 +464,6 @@ requires-dist = [ { name = "damnit", specifier = "~=0.2.1" }, { name = "h5py", specifier = "~=3.9" }, { name = "httpx", specifier = "~=0.27" }, - { name = "itsdangerous", specifier = "~=2.1" }, { name = "ldap3", specifier = "~=2.9" }, { name = "litestar", specifier = "~=2.24.0" }, { name = "matplotlib", specifier = "~=3.7" }, @@ -855,15 +853,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/3e/95/c7c34aa53c16353c56d0b802fba48d5f5caa2cdee7958acbcb795c830416/isort-8.0.1-py3-none-any.whl", hash = "sha256:28b89bc70f751b559aeca209e6120393d43fbe2490de0559662be7a9787e3d75", size = 89733, upload-time = "2026-02-28T10:08:19.466Z" }, ] -[[package]] -name = "itsdangerous" -version = "2.2.0" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/9c/cb/8ac0172223afbccb63986cc25049b154ecfb5e85932587206f42317be31d/itsdangerous-2.2.0.tar.gz", hash = "sha256:e0050c0b7da1eea53ffaf149c0cfbb5c6e2e2b69c4bef22c81fa6eb73e5f6173", size = 54410, upload-time = "2024-04-16T21:28:15.614Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/04/96/92447566d16df59b2a776c0fb82dbc4d9e07cd95062562af01e408583fc4/itsdangerous-2.2.0-py3-none-any.whl", hash = "sha256:c6242fc49e35958c8b15141343aa660db5fc54d4f13a1db01a3f5891b98700ef", size = 16234, upload-time = "2024-04-16T21:28:14.499Z" }, -] - [[package]] name = "jinja2" version = "3.1.6"