diff --git a/.github/workflows/pr.yml b/.github/workflows/pr.yml index 1a3704e5..356003d8 100644 --- a/.github/workflows/pr.yml +++ b/.github/workflows/pr.yml @@ -147,6 +147,9 @@ jobs: SM_SECRET_KEY: ci-test-key SM_USERS_BOOTSTRAP_EMAIL: admin@example.com SM_USERS_BOOTSTRAP_PASSWORD: admin + # Exclude Keycloak module — SM020 prevents both users and keycloak + # from running simultaneously. E2E tests use the users module. + SM_MODULES_ENABLED: '["Auth","Users","Dashboard","Permissions","Settings","BackgroundTasks","FileStorage","FeatureFlags"]' E2E_BASE_URL: http://localhost:8000 steps: - uses: actions/checkout@v6 diff --git a/.verify/01-login-page.png b/.verify/01-login-page.png new file mode 100644 index 00000000..4d1a1efc Binary files /dev/null and b/.verify/01-login-page.png differ diff --git a/.verify/02-dashboard.png b/.verify/02-dashboard.png new file mode 100644 index 00000000..2b0e9085 Binary files /dev/null and b/.verify/02-dashboard.png differ diff --git a/.verify/03-keycloak-login.png b/.verify/03-keycloak-login.png new file mode 100644 index 00000000..01425ee7 Binary files /dev/null and b/.verify/03-keycloak-login.png differ diff --git a/.verify/04-keycloak-authenticated.png b/.verify/04-keycloak-authenticated.png new file mode 100644 index 00000000..7386ae53 Binary files /dev/null and b/.verify/04-keycloak-authenticated.png differ diff --git a/.verify/05-session-survives-expiry.png b/.verify/05-session-survives-expiry.png new file mode 100644 index 00000000..025c45b9 Binary files /dev/null and b/.verify/05-session-survives-expiry.png differ diff --git a/.verify/06-after-logout.png b/.verify/06-after-logout.png new file mode 100644 index 00000000..94ba9af6 Binary files /dev/null and b/.verify/06-after-logout.png differ diff --git a/CLAUDE.md b/CLAUDE.md index 5342a295..e0b3524f 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -78,7 +78,7 @@ Standard mixins in `simple_module_db.mixins`: `AuditMixin`, `SoftDeleteMixin` (b ## Diagnostic codes -Meaningful codes when reading `make doctor` output: `SM001` missing meta (error), `SM003` orphan page / `SM004` phantom render (warn), `SM007` module overrides no hooks (info), `SM008` duplicate name (error), `SM009` framework→plugin import (error), `SM010` DB revision behind head (error), `SM011` module table not in migration history (warn), `SM012` `register_settings` overridden but nothing on `app.state.` (warn, fires at dev boot only), `SM013`–`SM016` locale issues, `SM017` module ships `.tsx` pages but is missing `package.json`/`tsconfig.json` (warn), `SM018` Inertia `router.{post,patch,put,delete}()` in a page targets a JSON `/api/*` endpoint (warn — Inertia rejects non-Inertia responses), `SM019` module registers view routes (non-empty `view_prefix` + overrides `register_routes`) but overrides neither `register_menu_items` nor `register_permissions` (warn — pages exist with no sidebar entry and no role-editor visibility; admins can't reach them through the UI). Modules whose views are sub-pages of another module typically register permissions to stay discoverable in the role editor without needing their own sidebar entry. In production, errors fail boot. +Meaningful codes when reading `make doctor` output: `SM001` missing meta (error), `SM003` orphan page / `SM004` phantom render (warn), `SM007` module overrides no hooks (info), `SM008` duplicate name (error), `SM009` framework→plugin import (error), `SM010` DB revision behind head (error), `SM011` module table not in migration history (warn), `SM012` `register_settings` overridden but nothing on `app.state.` (warn, fires at dev boot only), `SM013`–`SM016` locale issues, `SM017` module ships `.tsx` pages but is missing `package.json`/`tsconfig.json` (warn), `SM018` Inertia `router.{post,patch,put,delete}()` in a page targets a JSON `/api/*` endpoint (warn — Inertia rejects non-Inertia responses), `SM019` module registers view routes (non-empty `view_prefix` + overrides `register_routes`) but overrides neither `register_menu_items` nor `register_permissions` (warn — pages exist with no sidebar entry and no role-editor visibility; admins can't reach them through the UI). Modules whose views are sub-pages of another module typically register permissions to stay discoverable in the role editor without needing their own sidebar entry. `SM020` multiple auth provider modules installed (error), `SM021` no auth provider module installed (warn). In production, errors fail boot. ## Tests & fixtures diff --git a/docs/framework-conventions.md b/docs/framework-conventions.md index 99f24cb8..42a517f8 100644 --- a/docs/framework-conventions.md +++ b/docs/framework-conventions.md @@ -241,12 +241,31 @@ async def create_order(...): ... The `auth` module exposes a principal-resolver chain on `app.state.auth.principal_resolvers` — a list of async callables that -`users.AuthMiddleware` consults after the session-cookie path. Use it to add +`AuthMiddleware` consults after the session-cookie path. Use it to add non-cookie credential sources (Personal Access Tokens, API keys, JWTs) without forking the middleware. See [`docs/framework/principal-resolvers.md`](framework/principal-resolvers.md) for the contract, ordering rules, and a worked Bearer-token example. +### Auth Provider Contract + +The framework supports swappable authentication backends. Exactly one auth provider +module must be installed — either `simple-module-users` (local credentials + OAuth) +or `simple-module-keycloak` (Keycloak OIDC). Both implement the `AuthProvider` +protocol from `auth.contracts.provider`. + +**Module authors never import from `users` or `keycloak` directly.** Use only: +- `from auth.deps import CurrentUser, require_permission` +- `from auth.contracts.schemas import UserContext` + +The `AuthMiddleware` (in `auth/middleware.py`) delegates to the active provider's +`resolve_user()` method, then falls through to the principal-resolver chain. +API paths (`/api/*`) receive 401 JSON when unauthenticated; view paths receive +a 302 redirect to the provider's login URL. + +Boot-time diagnostic `SM020` fails if multiple auth providers are installed. +`SM021` warns if none is installed. + ## Events Base class: `Event` from `simple_module_core.events`. Subclass per domain event: @@ -290,6 +309,8 @@ Dispatch walks the event's MRO, so subscribing to a base class delivers subclass | SM014 | WARNING | Non-default locale missing keys present in the default | | SM015 | WARNING | Non-default locale has keys not in the default | | SM016 | ERROR | Locale JSON invalid or contains non-string leaves | +| SM020 | ERROR | Multiple auth provider modules installed | +| SM021 | WARNING | No auth provider module installed | Diagnostics run automatically at boot — warnings print to stderr in dev, errors abort startup in production. There's no separate "run diagnostics" step in a scaffolded app: just start the server. diff --git a/docs/superpowers/plans/2026-05-27-pluggable-auth-keycloak.md b/docs/superpowers/plans/2026-05-27-pluggable-auth-keycloak.md new file mode 100644 index 00000000..b064121c --- /dev/null +++ b/docs/superpowers/plans/2026-05-27-pluggable-auth-keycloak.md @@ -0,0 +1,3046 @@ +# Pluggable Auth + Keycloak Module Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Make the auth layer swappable — framework users install either `users` (local credentials) or a new `keycloak` module (Keycloak OIDC), both implementing the same `AuthProvider` contract. Mobile clients get bearer-token support regardless of provider. + +**Architecture:** The `auth` module (contract layer) gains an `AuthProvider` protocol and a provider-agnostic `AuthMiddleware`. The `users` module implements `AuthProvider` and adds bearer-token + refresh-token endpoints. A new `keycloak` module implements `AuthProvider` via OIDC + JWKS JWT validation. A boot-time diagnostic (SM020) ensures only one provider is installed. + +**Tech Stack:** Python 3.12, FastAPI, SQLModel, PyJWT, httpx, Starlette SessionMiddleware, Alembic + +**Prerequisite:** The principal-resolver chain from `2026-05-21-auth-principal-resolver-design.md` is already partially implemented (`AuthState` with `principal_resolvers` exists on `app.state.auth`). + +--- + +## File Map + +### Auth module (contract layer) — modify existing + +| File | Action | Responsibility | +|------|--------|---------------| +| `modules/auth/auth/contracts/provider.py` | Create | `AuthProvider` protocol definition | +| `modules/auth/auth/contracts/__init__.py` | Modify | Re-export `AuthProvider` | +| `modules/auth/auth/__init__.py` | Modify | Re-export `AuthProvider` | +| `modules/auth/auth/state.py` | Modify | Add `auth_provider` field to `AuthState` | +| `modules/auth/auth/middleware.py` | Create | Provider-agnostic `AuthMiddleware` | +| `modules/auth/auth/module.py` | Modify | Register `AuthMiddleware` + `principal_serializer` | +| `modules/auth/tests/test_auth_provider_protocol.py` | Create | Protocol conformance tests | +| `modules/auth/tests/test_auth_middleware.py` | Create | Provider-agnostic middleware tests | + +### Users module — modify existing + +| File | Action | Responsibility | +|------|--------|---------------| +| `modules/users/users/provider.py` | Create | `UsersAuthProvider` implementing `AuthProvider` | +| `modules/users/users/models/refresh_token.py` | Create | `RefreshToken` SQLModel table | +| `modules/users/users/auth_local/token_api.py` | Create | `POST/DELETE /api/users/auth/token`, `POST .../token/refresh` | +| `modules/users/users/module.py` | Modify | Register as `auth_provider`, remove `AuthMiddleware` + `principal_serializer` registration | +| `modules/users/users/middleware.py` | Delete (or keep as thin import wrapper) | Logic moved to `auth/middleware.py` + `users/provider.py` | +| `modules/users/users/settings.py` | Modify | Add `bearer_token_lifetime_seconds` setting | +| `modules/users/tests/test_users_provider.py` | Create | `UsersAuthProvider` tests | +| `modules/users/tests/test_token_api.py` | Create | Token endpoint tests | + +### Framework diagnostics — modify existing + +| File | Action | Responsibility | +|------|--------|---------------| +| `framework/core/simple_module_core/diagnostics/_module.py` | Modify | Add `SM020`/`SM021` checks | +| `framework/core/tests/test_diagnostics.py` | Modify | Tests for new diagnostics | + +### Keycloak module — new package + +| File | Action | Responsibility | +|------|--------|---------------| +| `modules/keycloak/pyproject.toml` | Create | Package metadata + entry point | +| `modules/keycloak/keycloak/__init__.py` | Create | Package marker | +| `modules/keycloak/keycloak/module.py` | Create | `KeycloakModule(ModuleBase)` | +| `modules/keycloak/keycloak/settings.py` | Create | `KeycloakSettings` | +| `modules/keycloak/keycloak/state.py` | Create | `KeycloakState` dataclass | +| `modules/keycloak/keycloak/provider.py` | Create | `KeycloakAuthProvider(AuthProvider)` | +| `modules/keycloak/keycloak/jwks.py` | Create | JWKS key cache + JWT validation | +| `modules/keycloak/keycloak/oidc.py` | Create | OIDC discovery, token exchange | +| `modules/keycloak/keycloak/models.py` | Create | `KeycloakUserCache` table | +| `modules/keycloak/keycloak/endpoints/api.py` | Create | Login redirect, callback, userinfo | +| `modules/keycloak/keycloak/endpoints/views.py` | Create | Inertia login/logout pages | +| `modules/keycloak/keycloak/contracts/__init__.py` | Create | Empty | +| `modules/keycloak/keycloak/locales/en.json` | Create | English translations | +| `modules/keycloak/keycloak/pages/Login.tsx` | Create | Auto-redirect login page | +| `modules/keycloak/keycloak/pages/LoggedOut.tsx` | Create | Post-logout landing | +| `modules/keycloak/package.json` | Create | JS workspace member | +| `modules/keycloak/tsconfig.json` | Create | TypeScript config | +| `modules/keycloak/tests/test_jwks.py` | Create | JWKS cache + JWT validation tests | +| `modules/keycloak/tests/test_oidc.py` | Create | OIDC helper tests | +| `modules/keycloak/tests/test_keycloak_provider.py` | Create | Provider implementation tests | +| `modules/keycloak/tests/test_keycloak_module.py` | Create | Module lifecycle tests | +| `modules/keycloak/tests/conftest.py` | Create | Keycloak test fixtures | + +### Alembic migration + +| File | Action | Responsibility | +|------|--------|---------------| +| `host/migrations/versions/_add_users_refresh_token.py` | Create | `users_refresh_token` table | +| `host/migrations/versions/_keycloak_user_cache.py` | Create | `keycloak_user_cache` table | + +### Workspace config + +| File | Action | Responsibility | +|------|--------|---------------| +| `pyproject.toml` (root) | Modify | Add `modules/keycloak` to `tool.ty.environment.extra-paths` and `tool.pytest.ini_options.testpaths` | + +--- + +## Task 1: AuthProvider Protocol + +**Files:** +- Create: `modules/auth/auth/contracts/provider.py` +- Modify: `modules/auth/auth/contracts/__init__.py` +- Modify: `modules/auth/auth/__init__.py` +- Create: `modules/auth/tests/test_auth_provider_protocol.py` + +- [ ] **Step 1: Write the test file** + +```python +# modules/auth/tests/test_auth_provider_protocol.py +"""Tests for the AuthProvider protocol.""" + +from __future__ import annotations + +from auth.contracts.provider import AuthProvider +from auth.contracts.schemas import UserContext +from starlette.requests import Request +from starlette.testclient import TestClient + + +class _FakeProvider: + """Minimal implementation to verify protocol conformance.""" + + name = "fake" + + async def resolve_user(self, request: Request) -> UserContext | None: + return None + + def get_login_url(self, request: Request, next_url: str | None = None) -> str: + return "/fake/login" + + def get_logout_url(self, request: Request) -> str: + return "/fake/logout" + + def get_public_paths(self) -> tuple[tuple[str, ...], tuple[str, ...]]: + return (("/fake/login",), ()) + + def is_bearer_request(self, request: Request) -> bool: + return False + + +def test_fake_provider_satisfies_protocol(): + provider = _FakeProvider() + assert isinstance(provider, AuthProvider) + + +def test_protocol_rejects_incomplete_implementation(): + class _Incomplete: + name = "broken" + + assert not isinstance(_Incomplete(), AuthProvider) + + +def test_auth_package_reexports_auth_provider(): + import auth + + assert hasattr(auth, "AuthProvider") + assert "AuthProvider" in auth.__all__ + from auth.contracts.provider import AuthProvider as Canonical + + assert auth.AuthProvider is Canonical + + +def test_contracts_package_reexports_auth_provider(): + from auth.contracts import AuthProvider + + assert AuthProvider is not None +``` + +- [ ] **Step 2: Run tests — they should fail (AuthProvider not defined yet)** + +Run: `uv run pytest modules/auth/tests/test_auth_provider_protocol.py -v` +Expected: `ModuleNotFoundError` or `ImportError` — `auth.contracts.provider` doesn't exist. + +- [ ] **Step 3: Create the AuthProvider protocol** + +```python +# modules/auth/auth/contracts/provider.py +"""AuthProvider protocol — the contract both users and keycloak modules implement.""" + +from __future__ import annotations + +from typing import Protocol, runtime_checkable + +from starlette.requests import Request + +from auth.contracts.schemas import UserContext + + +@runtime_checkable +class AuthProvider(Protocol): + """Extension point for swappable authentication backends. + + Exactly one module (``users`` or ``keycloak``) registers an implementation + on ``app.state.auth.auth_provider`` during ``register_settings``. + The ``AuthMiddleware`` delegates to it on every request. + """ + + name: str + + async def resolve_user(self, request: Request) -> UserContext | None: ... + + def get_login_url(self, request: Request, next_url: str | None = None) -> str: ... + + def get_logout_url(self, request: Request) -> str: ... + + def get_public_paths(self) -> tuple[tuple[str, ...], tuple[str, ...]]: ... + + def is_bearer_request(self, request: Request) -> bool: ... + + +__all__ = ["AuthProvider"] +``` + +- [ ] **Step 4: Update contracts `__init__.py`** + +Change `modules/auth/auth/contracts/__init__.py` from: +```python +"""Auth contracts — public types for other modules.""" + +from auth.contracts.schemas import UserContext + +__all__ = ["UserContext"] +``` +to: +```python +"""Auth contracts — public types for other modules.""" + +from auth.contracts.provider import AuthProvider +from auth.contracts.schemas import UserContext + +__all__ = ["AuthProvider", "UserContext"] +``` + +- [ ] **Step 5: Update auth package `__init__.py`** + +Change `modules/auth/auth/__init__.py` from: +```python +"""Auth module — shared contracts (UserContext, PrincipalResolver, deps).""" + +from auth.contracts.resolver import PrincipalResolver +from auth.contracts.schemas import UserContext + +__all__ = ["PrincipalResolver", "UserContext"] +``` +to: +```python +"""Auth module — shared contracts (UserContext, AuthProvider, PrincipalResolver, deps).""" + +from auth.contracts.provider import AuthProvider +from auth.contracts.resolver import PrincipalResolver +from auth.contracts.schemas import UserContext + +__all__ = ["AuthProvider", "PrincipalResolver", "UserContext"] +``` + +- [ ] **Step 6: Run tests — they should pass** + +Run: `uv run pytest modules/auth/tests/test_auth_provider_protocol.py -v` +Expected: All 4 tests pass. + +- [ ] **Step 7: Commit** + +```bash +git add modules/auth/auth/contracts/provider.py modules/auth/auth/contracts/__init__.py modules/auth/auth/__init__.py modules/auth/tests/test_auth_provider_protocol.py +git commit -m "feat(auth): add AuthProvider protocol for swappable auth backends" +``` + +--- + +## Task 2: Add auth_provider to AuthState + +**Files:** +- Modify: `modules/auth/auth/state.py` +- Modify: `modules/auth/tests/test_resolver_registry.py` + +- [ ] **Step 1: Write the test** + +Add to `modules/auth/tests/test_resolver_registry.py`: + +```python +def test_auth_state_has_auth_provider_field(): + from auth.state import AuthState + + state = AuthState() + assert state.auth_provider is None + + +def test_auth_state_accepts_auth_provider(): + from auth.contracts.provider import AuthProvider + from auth.state import AuthState + + class FakeProvider: + name = "fake" + + async def resolve_user(self, request): + return None + + def get_login_url(self, request, next_url=None): + return "/login" + + def get_logout_url(self, request): + return "/logout" + + def get_public_paths(self): + return ((), ()) + + def is_bearer_request(self, request): + return False + + provider = FakeProvider() + state = AuthState(auth_provider=provider) + assert state.auth_provider is provider + assert isinstance(state.auth_provider, AuthProvider) +``` + +- [ ] **Step 2: Run test — should fail** + +Run: `uv run pytest modules/auth/tests/test_resolver_registry.py::test_auth_state_has_auth_provider_field -v` +Expected: `TypeError` — `AuthState` doesn't accept `auth_provider`. + +- [ ] **Step 3: Add auth_provider field to AuthState** + +Change `modules/auth/auth/state.py` from: +```python +"""Module-owned state attached to ``app.state.auth`` by ``AuthModule.register_settings``. + +Holds the principal-resolver registry (see +``auth.contracts.resolver.PrincipalResolver``). Apps register additional +resolvers from their ``on_startup`` hook:: + + app.state.auth.principal_resolvers.append(my_pat_resolver) +""" + +from __future__ import annotations + +from dataclasses import dataclass, field + +from auth.contracts.resolver import PrincipalResolver + + +@dataclass +class AuthState: + """Per-app auth registry. Initialized empty; modules append resolvers.""" + + principal_resolvers: list[PrincipalResolver] = field(default_factory=list) + + +__all__ = ["AuthState"] +``` +to: +```python +"""Module-owned state attached to ``app.state.auth`` by ``AuthModule.register_settings``. + +Holds the auth provider (set by one of ``users`` or ``keycloak``) and the +principal-resolver registry. Apps register additional resolvers from their +``on_startup`` hook:: + + app.state.auth.principal_resolvers.append(my_pat_resolver) +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import TYPE_CHECKING + +from auth.contracts.resolver import PrincipalResolver + +if TYPE_CHECKING: + from auth.contracts.provider import AuthProvider + + +@dataclass +class AuthState: + """Per-app auth registry. Initialized empty; provider modules populate at boot.""" + + auth_provider: AuthProvider | None = None + principal_resolvers: list[PrincipalResolver] = field(default_factory=list) + + +__all__ = ["AuthState"] +``` + +- [ ] **Step 4: Run all auth tests — should pass** + +Run: `uv run pytest modules/auth/tests/ -v` +Expected: All pass, including the two new tests. + +- [ ] **Step 5: Commit** + +```bash +git add modules/auth/auth/state.py modules/auth/tests/test_resolver_registry.py +git commit -m "feat(auth): add auth_provider slot to AuthState" +``` + +--- + +## Task 3: Provider-Agnostic AuthMiddleware in auth/ + +**Files:** +- Create: `modules/auth/auth/middleware.py` +- Create: `modules/auth/tests/test_auth_middleware.py` + +- [ ] **Step 1: Write middleware tests** + +```python +# modules/auth/tests/test_auth_middleware.py +"""Tests for the provider-agnostic AuthMiddleware.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import httpx +import pytest +from auth.contracts.schemas import UserContext +from auth.middleware import AuthMiddleware +from auth.state import AuthState +from fastapi import FastAPI, Request +from starlette.middleware.sessions import SessionMiddleware +from starlette.responses import JSONResponse + +SECRET = "test-middleware-secret" + +_TEST_USER = UserContext( + id="aaaaaaaa-bbbb-cccc-dddd-eeeeeeeeeeee", + email="test@example.com", + name="Test User", + roles=["admin"], +) + + +class _StubProvider: + name = "stub" + + def __init__(self, *, user: UserContext | None = None): + self._user = user + + async def resolve_user(self, request): + return self._user + + def get_login_url(self, request, next_url=None): + return "/stub/login" + + def get_logout_url(self, request): + return "/stub/logout" + + def get_public_paths(self): + return (("/stub/login", "/stub/public/"), ()) + + def is_bearer_request(self, request): + auth = request.headers.get("authorization", "") + return auth.startswith("Bearer ") + + +def _build_app(provider, *, principal_resolvers=None): + app = FastAPI() + app.state.auth = AuthState( + auth_provider=provider, + principal_resolvers=list(principal_resolvers or []), + ) + + @app.get("/{path:path}") + async def catch_all(request: Request, path: str = ""): + user = getattr(request.state, "user", None) + return JSONResponse({ + "user": user.to_session_dict() if user else None, + }) + + app.add_middleware(AuthMiddleware) + app.add_middleware(SessionMiddleware, secret_key=SECRET) + return app + + +@pytest.fixture +def authenticated_app(): + return _build_app(_StubProvider(user=_TEST_USER)) + + +@pytest.fixture +def unauthenticated_app(): + return _build_app(_StubProvider(user=None)) + + +async def test_authenticated_request_sets_user(authenticated_app): + async with httpx.AsyncClient(app=authenticated_app, base_url="http://test") as c: + resp = await c.get("/some/page") + assert resp.status_code == 200 + assert resp.json()["user"]["email"] == "test@example.com" + + +async def test_unauthenticated_browser_redirects_to_login(unauthenticated_app): + async with httpx.AsyncClient( + app=unauthenticated_app, base_url="http://test", follow_redirects=False + ) as c: + resp = await c.get("/protected/page") + assert resp.status_code == 302 + assert resp.headers["location"] == "/stub/login" + + +async def test_unauthenticated_api_returns_401(unauthenticated_app): + async with httpx.AsyncClient( + app=unauthenticated_app, base_url="http://test" + ) as c: + resp = await c.get("/api/protected") + assert resp.status_code == 401 + assert resp.json()["detail"] == "Not authenticated" + + +async def test_unauthenticated_bearer_returns_401(unauthenticated_app): + async with httpx.AsyncClient( + app=unauthenticated_app, base_url="http://test" + ) as c: + resp = await c.get("/some/page", headers={"Authorization": "Bearer bad"}) + assert resp.status_code == 401 + + +async def test_public_paths_skip_auth(unauthenticated_app): + async with httpx.AsyncClient( + app=unauthenticated_app, base_url="http://test" + ) as c: + resp = await c.get("/stub/login") + assert resp.status_code == 200 + + +async def test_framework_public_paths_skip_auth(unauthenticated_app): + async with httpx.AsyncClient( + app=unauthenticated_app, base_url="http://test" + ) as c: + resp = await c.get("/health") + assert resp.status_code == 200 + + +async def test_root_is_public(unauthenticated_app): + async with httpx.AsyncClient( + app=unauthenticated_app, base_url="http://test" + ) as c: + resp = await c.get("/") + assert resp.status_code == 200 + + +async def test_resolver_chain_fallback(): + """When provider returns None, fall through to principal resolvers.""" + + async def fake_resolver(request): + auth = request.headers.get("authorization", "") + if auth == "Bearer good-token": + return _TEST_USER + return None + + app = _build_app(_StubProvider(user=None), principal_resolvers=[fake_resolver]) + async with httpx.AsyncClient(app=app, base_url="http://test") as c: + resp = await c.get("/protected", headers={"Authorization": "Bearer good-token"}) + assert resp.status_code == 200 + assert resp.json()["user"]["email"] == "test@example.com" + + +async def test_resolver_exception_is_logged_and_skipped(): + """A resolver that raises should be caught; middleware continues.""" + + async def bad_resolver(request): + raise RuntimeError("boom") + + app = _build_app(_StubProvider(user=None), principal_resolvers=[bad_resolver]) + async with httpx.AsyncClient( + app=app, base_url="http://test", follow_redirects=False + ) as c: + resp = await c.get("/protected/page") + assert resp.status_code == 302 +``` + +- [ ] **Step 2: Run tests — they should fail** + +Run: `uv run pytest modules/auth/tests/test_auth_middleware.py -v` +Expected: `ImportError` — `auth.middleware` doesn't exist. + +- [ ] **Step 3: Create the provider-agnostic AuthMiddleware** + +```python +# modules/auth/auth/middleware.py +"""Provider-agnostic authentication middleware. + +Delegates user resolution to the ``AuthProvider`` registered on +``app.state.auth.auth_provider``, then falls through to the +principal-resolver chain. Sets ``request.state.user`` and the +``current_user_id`` ContextVar for audit listeners. +""" + +from __future__ import annotations + +import logging + +from simple_module_db.listeners import current_user_id +from starlette.requests import Request +from starlette.responses import JSONResponse, RedirectResponse +from starlette.types import ASGIApp, Receive, Scope, Send + +logger = logging.getLogger(__name__) + +_FRAMEWORK_PUBLIC_PREFIXES = ( + "/health", + "/static/", + "/api/docs", + "/api/redoc", + "/openapi.json", + "/i18n/", +) +_FRAMEWORK_PUBLIC_EXACT = ("/",) +_SESSION_NEXT_KEY = "next" + + +class AuthMiddleware: + """Authenticate requests via the registered AuthProvider. + + On cache miss (provider returns None), falls through to the + principal-resolver chain. Unauthenticated API requests get 401 JSON; + unauthenticated browser requests get a redirect to the provider's + login URL. + """ + + def __init__(self, app: ASGIApp) -> None: + self.app = app + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] != "http": + await self.app(scope, receive, send) + return + + path: str = scope["path"] + auth_state = scope["app"].state.auth + provider = auth_state.auth_provider + + if provider is None: + await self.app(scope, receive, send) + return + + is_public = ( + any(path.startswith(p) for p in _FRAMEWORK_PUBLIC_PREFIXES) + or path in _FRAMEWORK_PUBLIC_EXACT + ) + if not is_public: + prefix_paths, exact_paths = provider.get_public_paths() + is_public = ( + any(path.startswith(p) for p in prefix_paths) or path in exact_paths + ) + + request = Request(scope) + user_ctx = await provider.resolve_user(request) + + if user_ctx is None: + for resolver in auth_state.principal_resolvers: + try: + user_ctx = await resolver(request) + except Exception: + logger.exception( + "Principal resolver %r raised; treating as no-match", + resolver, + ) + continue + if user_ctx is not None: + break + + if user_ctx is None and not is_public: + if path.startswith("/api/") or provider.is_bearer_request(request): + response = JSONResponse( + {"detail": "Not authenticated"}, status_code=401 + ) + else: + session = scope.get("session", {}) + session[_SESSION_NEXT_KEY] = str(request.url) + response = RedirectResponse( + provider.get_login_url(request), status_code=302 + ) + await response(scope, receive, send) + return + + if user_ctx is not None: + request.state.user = user_ctx + token = current_user_id.set(user_ctx.id) + try: + await self.app(scope, receive, send) + finally: + current_user_id.reset(token) + return + + await self.app(scope, receive, send) +``` + +- [ ] **Step 4: Run tests — should pass** + +Run: `uv run pytest modules/auth/tests/test_auth_middleware.py -v` +Expected: All 9 tests pass. + +- [ ] **Step 5: Commit** + +```bash +git add modules/auth/auth/middleware.py modules/auth/tests/test_auth_middleware.py +git commit -m "feat(auth): add provider-agnostic AuthMiddleware" +``` + +--- + +## Task 4: Move principal_serializer + AuthMiddleware Registration to AuthModule + +**Files:** +- Modify: `modules/auth/auth/module.py` +- Modify: `modules/users/users/module.py` +- Modify: `modules/auth/tests/test_resolver_registry.py` + +- [ ] **Step 1: Write tests for the new AuthModule behavior** + +Add to `modules/auth/tests/test_resolver_registry.py`: + +```python +def test_auth_module_registers_middleware(): + """AuthModule.register_middleware should add AuthMiddleware.""" + from auth.module import AuthModule + from fastapi import FastAPI + + app = FastAPI() + AuthModule().register_middleware(app) + middleware_classes = [m.cls.__name__ for m in app.user_middleware] + assert "AuthMiddleware" in middleware_classes + + +def test_auth_module_registers_principal_serializer(): + """AuthModule.register_settings should set principal_serializer on app.state.""" + from auth.module import AuthModule + from fastapi import FastAPI + + app = FastAPI() + AuthModule().register_settings(app) + serializer = getattr(app.state, "principal_serializer", None) + assert serializer is not None + + from auth.contracts.schemas import UserContext + + ctx = UserContext(id="123", email="a@b.com", name="Test", roles=["admin"]) + result = serializer(ctx) + assert result == {"id": "123", "name": "Test", "email": "a@b.com", "roles": ["admin"]} +``` + +- [ ] **Step 2: Run tests — should fail** + +Run: `uv run pytest modules/auth/tests/test_resolver_registry.py::test_auth_module_registers_middleware -v` +Expected: FAIL — `AuthModule` doesn't override `register_middleware`. + +- [ ] **Step 3: Update AuthModule to register middleware + principal_serializer** + +Change `modules/auth/auth/module.py` to: + +```python +"""Auth module — shared contracts (UserContext, AuthProvider, deps). + +Intentionally minimal: this module owns the PUBLIC interface (UserContext, +AuthProvider, PrincipalResolver, get_current_user, CurrentUser, require_permission) +that every other module imports. Keeping it stable prevents churn when auth +internals change. + +The ``auth_provider`` slot on ``app.state.auth`` is the extension point +auth-provider modules (``users``, ``keycloak``) use to register themselves. +The ``principal_resolvers`` registry lets downstream modules add extra +credential sources (PAT bearer tokens, API keys, etc.). +""" + +from __future__ import annotations + +import importlib.resources +from pathlib import Path +from typing import TYPE_CHECKING + +from simple_module_core.module import ModuleBase, ModuleMeta + +if TYPE_CHECKING: + from fastapi import FastAPI + + from auth.contracts.schemas import UserContext + + +def _serialize_principal(user: UserContext) -> dict: + return { + "id": user.id, + "name": user.name, + "email": user.email, + "roles": user.roles, + } + + +class AuthModule(ModuleBase): + meta = ModuleMeta( + name="Auth", + route_prefix="/auth", + ) + + def register_settings(self, app: FastAPI) -> None: + from auth.state import AuthState + + app.state.auth = AuthState() + app.state.principal_serializer = _serialize_principal + + def register_middleware(self, app: FastAPI) -> None: + from auth.middleware import AuthMiddleware + + app.add_middleware(AuthMiddleware) + + def locale_dirs(self) -> dict[str, Path]: + return {"auth": Path(str(importlib.resources.files(__package__) / "locales"))} +``` + +- [ ] **Step 4: Remove middleware + serializer registration from UsersModule** + +In `modules/users/users/module.py`: + +Remove the `register_middleware` method entirely (lines 159-162): +```python + # DELETE THIS METHOD: + def register_middleware(self, app: FastAPI) -> None: + from users.middleware import AuthMiddleware + + app.add_middleware(AuthMiddleware) +``` + +In `register_settings`, remove the `serialize_principal` function and `app.state.principal_serializer` line (lines 61-69): +```python + # DELETE THESE LINES from register_settings: + def serialize_principal(user: UserContext) -> dict: + return { + "id": user.id, + "name": user.name, + "email": user.email, + "roles": user.roles, + } + + app.state.principal_serializer = serialize_principal +``` + +Also remove the unused `from auth.contracts.schemas import UserContext` import from `register_settings` (it was only used by the deleted serializer). + +- [ ] **Step 5: Run auth + users tests** + +Run: `uv run pytest modules/auth/tests/ modules/users/tests/ -v` +Expected: All pass. The middleware behavior is unchanged — it still delegates to the provider; users module just no longer registers it. + +- [ ] **Step 6: Commit** + +```bash +git add modules/auth/auth/module.py modules/users/users/module.py modules/auth/tests/test_resolver_registry.py +git commit -m "refactor(auth,users): move AuthMiddleware + principal_serializer to auth module" +``` + +--- + +## Task 5: UsersAuthProvider Implementation + +**Files:** +- Create: `modules/users/users/provider.py` +- Create: `modules/users/tests/test_users_provider.py` +- Modify: `modules/users/users/module.py` + +- [ ] **Step 1: Write provider tests** + +```python +# modules/users/tests/test_users_provider.py +"""Tests for UsersAuthProvider.""" + +from __future__ import annotations + +import uuid + +import httpx +import pytest +from auth.contracts.provider import AuthProvider +from auth.contracts.schemas import UserContext +from auth.state import AuthState +from fastapi import FastAPI, Request +from starlette.middleware.sessions import SessionMiddleware +from starlette.responses import JSONResponse +from users.provider import UsersAuthProvider + +SECRET = "test-provider-secret" + + +def test_users_provider_satisfies_protocol(): + provider = UsersAuthProvider() + assert isinstance(provider, AuthProvider) + + +def test_login_url(): + provider = UsersAuthProvider() + assert provider.get_login_url(None) == "/users/login" + + +def test_logout_url(): + provider = UsersAuthProvider() + assert provider.get_logout_url(None) == "/users/logout" + + +def test_public_paths(): + provider = UsersAuthProvider() + prefixes, exact = provider.get_public_paths() + assert "/users/login" in prefixes + assert "/api/users/auth/" in prefixes + + +def test_is_bearer_request_true(): + from unittest.mock import MagicMock + + request = MagicMock() + request.headers = {"authorization": "Bearer abc123"} + provider = UsersAuthProvider() + assert provider.is_bearer_request(request) is True + + +def test_is_bearer_request_false(): + from unittest.mock import MagicMock + + request = MagicMock() + request.headers = {} + provider = UsersAuthProvider() + assert provider.is_bearer_request(request) is False +``` + +- [ ] **Step 2: Run tests — should fail** + +Run: `uv run pytest modules/users/tests/test_users_provider.py -v` +Expected: `ImportError` — `users.provider` doesn't exist. + +- [ ] **Step 3: Create UsersAuthProvider** + +```python +# modules/users/users/provider.py +"""UsersAuthProvider — AuthProvider implementation for the users module. + +Resolves users from session cookies (browser) or the principal-resolver chain +(bearer tokens, PATs). Session handling mirrors the original AuthMiddleware +logic: fast path from ``session["user_ctx"]``, slow path via DB lookup. +""" + +from __future__ import annotations + +import logging +import uuid as uuid_mod + +from auth.contracts.provider import AuthProvider +from auth.contracts.schemas import UserContext +from starlette.requests import Request + +logger = logging.getLogger(__name__) + +_SESSION_USER_ID_KEY = "user_id" +_SESSION_USER_CTX_KEY = "user_ctx" + + +class UsersAuthProvider: + """Cookie-based auth provider using fastapi-users' DatabaseStrategy.""" + + name = "users" + _is_auth_provider = True + + async def resolve_user(self, request: Request) -> UserContext | None: + session = request.scope.get("session", {}) + raw_user_id = session.get(_SESSION_USER_ID_KEY) + if not raw_user_id: + return None + + user_id_str = str(raw_user_id) + + cached = UserContext.from_session_dict(session.get(_SESSION_USER_CTX_KEY)) + if cached is not None and cached.id == user_id_str: + return cached + + try: + user_uuid = uuid_mod.UUID(user_id_str) + except (ValueError, TypeError): + logger.warning("Invalid user_id in session: %r", raw_user_id) + session.pop(_SESSION_USER_ID_KEY, None) + session.pop(_SESSION_USER_CTX_KEY, None) + return None + + user_ctx = await self._load_user(request.scope, user_uuid) + if user_ctx is None: + session.pop(_SESSION_USER_ID_KEY, None) + session.pop(_SESSION_USER_CTX_KEY, None) + else: + session[_SESSION_USER_CTX_KEY] = user_ctx.to_session_dict() + return user_ctx + + def get_login_url(self, request: Request | None, next_url: str | None = None) -> str: + return "/users/login" + + def get_logout_url(self, request: Request | None) -> str: + return "/users/logout" + + def get_public_paths(self) -> tuple[tuple[str, ...], tuple[str, ...]]: + return ( + ( + "/users/login", + "/users/register", + "/users/forgot-password", + "/users/reset-password", + "/users/verify", + "/users/invite/accept", + "/api/users/auth/", + "/api/users/register", + ), + (), + ) + + def is_bearer_request(self, request: Request | None) -> bool: + if request is None: + return False + return request.headers.get("authorization", "").startswith("Bearer ") + + async def _load_user(self, scope, user_id: uuid_mod.UUID) -> UserContext | None: + try: + from sqlalchemy import select + from sqlalchemy.orm import selectinload + + from users.models import User + + session_factory = scope["app"].state.sm.db.session_factory + async with session_factory() as db_session: + stmt = ( + select(User) + .where(User.id == user_id) + .options(selectinload(User.roles)) + ) + user = (await db_session.execute(stmt)).scalar_one_or_none() + if user is None or not user.is_active or user.disabled_at is not None: + return None + return UserContext.from_user(user) + except Exception: + logger.exception( + "Failed to load user %s from DB; treating as unauthenticated", + user_id, + ) + return None +``` + +- [ ] **Step 4: Run tests — should pass** + +Run: `uv run pytest modules/users/tests/test_users_provider.py -v` +Expected: All 6 tests pass. + +- [ ] **Step 5: Register UsersAuthProvider in UsersModule** + +In `modules/users/users/module.py`, update `register_settings` to add (after the `register_module_settings` call): + +```python + from users.provider import UsersAuthProvider + + app.state.auth.auth_provider = UsersAuthProvider() +``` + +- [ ] **Step 6: Run full test suite** + +Run: `uv run pytest modules/auth/tests/ modules/users/tests/ -v` +Expected: All pass. + +- [ ] **Step 7: Commit** + +```bash +git add modules/users/users/provider.py modules/users/tests/test_users_provider.py modules/users/users/module.py +git commit -m "feat(users): implement UsersAuthProvider with session-cookie resolution" +``` + +--- + +## Task 6: Remove Old AuthMiddleware from Users Module + +**Files:** +- Modify: `modules/users/users/middleware.py` (remove or convert to thin re-export) +- Modify: `modules/users/tests/_middleware_support.py` +- Modify: `modules/users/tests/test_users_middleware.py` (update to use new middleware) + +- [ ] **Step 1: Update middleware test support to use auth.middleware** + +In `modules/users/tests/_middleware_support.py`, change the import: + +From: +```python +from users.middleware import AuthMiddleware +``` +To: +```python +from auth.middleware import AuthMiddleware +``` + +And update `_build_app` to set `auth_provider` on the `AuthState`: + +From: +```python + app.state.auth = AuthState( + principal_resolvers=list(principal_resolvers or []), + ) +``` +To: +```python + from users.provider import UsersAuthProvider + + app.state.auth = AuthState( + auth_provider=UsersAuthProvider(), + principal_resolvers=list(principal_resolvers or []), + ) +``` + +- [ ] **Step 2: Run existing middleware tests** + +Run: `uv run pytest modules/users/tests/test_users_middleware.py -v` +Expected: All pass — the behavior is identical, just routed through `auth.middleware` → `UsersAuthProvider` instead of `users.middleware.AuthMiddleware` directly. + +- [ ] **Step 3: Replace users/middleware.py with a deprecation re-export** + +Replace `modules/users/users/middleware.py` contents with: + +```python +"""Backwards-compatibility re-export. + +The canonical AuthMiddleware now lives in ``auth.middleware``. This shim +exists only to avoid breaking imports in downstream apps that referenced +``users.middleware.AuthMiddleware`` directly. +""" + +from auth.middleware import AuthMiddleware + +__all__ = ["AuthMiddleware"] +``` + +- [ ] **Step 4: Run full test suite** + +Run: `uv run pytest modules/auth/tests/ modules/users/tests/ -v` +Expected: All pass. + +- [ ] **Step 5: Commit** + +```bash +git add modules/users/users/middleware.py modules/users/tests/_middleware_support.py +git commit -m "refactor(users): delegate to auth.middleware, keep thin re-export for compat" +``` + +--- + +## Task 7: SM020/SM021 Diagnostics + +**Files:** +- Modify: `framework/core/simple_module_core/diagnostics/_module.py` +- Modify or create test in: `framework/core/tests/test_diagnostics.py` (or equivalent) + +- [ ] **Step 1: Write diagnostic tests** + +Find the existing diagnostics test file and add: + +```python +def test_sm020_multiple_auth_providers(): + """SM020 fires when two modules both set _is_auth_provider.""" + from simple_module_core.diagnostics._module import ModuleDiagnostics + from simple_module_core.module import ModuleBase, ModuleMeta + + class FakeUsersModule(ModuleBase): + meta = ModuleMeta(name="Users") + _is_auth_provider = True + + class FakeKeycloakModule(ModuleBase): + meta = ModuleMeta(name="Keycloak") + _is_auth_provider = True + + diags = ModuleDiagnostics() + results = diags._check_auth_provider_conflict([FakeUsersModule(), FakeKeycloakModule()]) + assert len(results) == 1 + assert results[0].code == "SM020" + assert results[0].level.name == "ERROR" + + +def test_sm021_no_auth_provider(): + """SM021 fires when no module sets _is_auth_provider.""" + from simple_module_core.diagnostics._module import ModuleDiagnostics + from simple_module_core.module import ModuleBase, ModuleMeta + + class FakeDashboard(ModuleBase): + meta = ModuleMeta(name="Dashboard") + + diags = ModuleDiagnostics() + results = diags._check_auth_provider_conflict([FakeDashboard()]) + assert len(results) == 1 + assert results[0].code == "SM021" + assert results[0].level.name == "WARNING" + + +def test_sm020_single_provider_passes(): + """No diagnostic when exactly one auth provider is installed.""" + from simple_module_core.diagnostics._module import ModuleDiagnostics + from simple_module_core.module import ModuleBase, ModuleMeta + + class FakeUsersModule(ModuleBase): + meta = ModuleMeta(name="Users") + _is_auth_provider = True + + class FakeDashboard(ModuleBase): + meta = ModuleMeta(name="Dashboard") + + diags = ModuleDiagnostics() + results = diags._check_auth_provider_conflict([FakeUsersModule(), FakeDashboard()]) + assert results == [] +``` + +- [ ] **Step 2: Run tests — should fail** + +Run: `uv run pytest framework/core/tests/ -k "sm020 or sm021 or auth_provider_conflict" -v` +Expected: `AttributeError` — `_check_auth_provider_conflict` doesn't exist. + +- [ ] **Step 3: Add the check method and wire it into `run()`** + +In `framework/core/simple_module_core/diagnostics/_module.py`, add to the `run()` method: +```python + diagnostics.extend(self._check_auth_provider_conflict(modules)) +``` + +Add the new method to `ModuleDiagnostics`: + +```python + def _check_auth_provider_conflict(self, modules: list[ModuleBase]) -> list[Diagnostic]: + """SM020/SM021: exactly one auth provider module must be installed.""" + providers = [m for m in modules if getattr(m, "_is_auth_provider", False)] + diags: list[Diagnostic] = [] + if len(providers) > 1: + names = ", ".join(m.meta.name for m in providers) + diags.append( + Diagnostic( + level=DiagnosticLevel.ERROR, + code="SM020", + message=f"Multiple auth provider modules installed: {names}", + module_name=providers[0].meta.name, + suggestion=( + "Install only one auth provider " + "(e.g. 'users' OR 'keycloak', not both)" + ), + ) + ) + elif len(providers) == 0: + diags.append( + Diagnostic( + level=DiagnosticLevel.WARNING, + code="SM021", + message="No auth provider module installed", + module_name="(none)", + suggestion=( + "Install an auth provider module " + "(e.g. 'simple-module-users' or 'simple-module-keycloak')" + ), + ) + ) + return diags +``` + +- [ ] **Step 4: Add `_is_auth_provider = True` to UsersModule** + +In `modules/users/users/module.py`, add the class attribute: +```python +class UsersModule(ModuleBase): + meta = ModuleMeta(...) + _is_auth_provider = True +``` + +- [ ] **Step 5: Run tests** + +Run: `uv run pytest framework/core/tests/ -k "sm020 or sm021 or auth_provider" -v` +Expected: All 3 new tests pass. + +- [ ] **Step 6: Commit** + +```bash +git add framework/core/simple_module_core/diagnostics/_module.py modules/users/users/module.py +git commit -m "feat(diagnostics): SM020/SM021 — exactly one auth provider required" +``` + +--- + +## Task 8: Keycloak Module Scaffold — Package + Module + Settings + +**Files:** +- Create: `modules/keycloak/pyproject.toml` +- Create: `modules/keycloak/keycloak/__init__.py` +- Create: `modules/keycloak/keycloak/module.py` +- Create: `modules/keycloak/keycloak/settings.py` +- Create: `modules/keycloak/keycloak/state.py` +- Create: `modules/keycloak/keycloak/contracts/__init__.py` +- Create: `modules/keycloak/keycloak/locales/en.json` +- Create: `modules/keycloak/package.json` +- Create: `modules/keycloak/tsconfig.json` +- Create: `modules/keycloak/tests/__init__.py` +- Create: `modules/keycloak/tests/conftest.py` +- Create: `modules/keycloak/tests/test_keycloak_module.py` +- Modify: `pyproject.toml` (root — add to `extra-paths` and `testpaths`) + +- [ ] **Step 1: Write module registration test** + +```python +# modules/keycloak/tests/test_keycloak_module.py +"""Tests for KeycloakModule lifecycle.""" + +from __future__ import annotations + +from auth.contracts.provider import AuthProvider + + +def test_keycloak_module_meta(): + from keycloak.module import KeycloakModule + + mod = KeycloakModule() + assert mod.meta.name == "Keycloak" + assert mod.meta.depends_on == ["Auth"] + assert mod._is_auth_provider is True + + +def test_keycloak_module_registers_provider(): + from auth.state import AuthState + from fastapi import FastAPI + from keycloak.module import KeycloakModule + + app = FastAPI() + app.state.auth = AuthState() + KeycloakModule().register_settings(app) + + assert app.state.auth.auth_provider is not None + assert isinstance(app.state.auth.auth_provider, AuthProvider) + assert app.state.auth.auth_provider.name == "keycloak" +``` + +- [ ] **Step 2: Create pyproject.toml** + +```toml +# modules/keycloak/pyproject.toml +[project] +name = "simple_module_keycloak" +version = "0.0.15" +description = "Keycloak OIDC authentication provider for simple_module — swap with simple_module_users" +readme = "README.md" +license = "MIT" +requires-python = ">=3.12" +authors = [{ name = "Anto Subash", email = "antosubash@live.com" }] +keywords = ["simple-module", "keycloak", "oidc", "authentication", "fastapi"] +classifiers = [ + "Development Status :: 3 - Alpha", + "Framework :: FastAPI", + "Intended Audience :: Developers", + "License :: OSI Approved :: MIT License", + "Operating System :: OS Independent", + "Programming Language :: Python :: 3", + "Programming Language :: Python :: 3.12", + "Topic :: Internet :: WWW/HTTP", + "Topic :: Software Development :: Libraries :: Application Frameworks", + "Typing :: Typed", +] +dependencies = [ + "simple_module_core==0.0.15", + "simple_module_db==0.0.15", + "simple_module_hosting==0.0.15", + "simple_module_settings==0.0.15", + "simple_module_auth==0.0.15", + "PyJWT[crypto]>=2.8", + "httpx>=0.27", +] + +[project.entry-points.simple_module] +keycloak = "keycloak.module:KeycloakModule" + +[project.urls] +Homepage = "https://github.com/antosubash/simple_module_python" +Repository = "https://github.com/antosubash/simple_module_python" + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[tool.hatch.build.targets.wheel] +packages = ["keycloak"] + +[tool.hatch.build.targets.wheel.force-include] +"package.json" = "keycloak/package.json" + +[tool.uv.sources] +simple_module_core = { workspace = true } +simple_module_db = { workspace = true } +simple_module_hosting = { workspace = true } +simple_module_settings = { workspace = true } +simple_module_auth = { workspace = true } +``` + +- [ ] **Step 3: Create package files** + +`modules/keycloak/keycloak/__init__.py`: +```python +"""Keycloak OIDC authentication provider for simple_module.""" +``` + +`modules/keycloak/keycloak/contracts/__init__.py`: +```python +"""Keycloak module contracts.""" +``` + +`modules/keycloak/keycloak/locales/en.json`: +```json +{ + "login": { + "redirecting": "Redirecting to identity provider…", + "title": "Sign In" + }, + "logout": { + "title": "Signed Out", + "message": "You have been signed out successfully." + }, + "errors": { + "callback_failed": "Authentication failed. Please try again.", + "invalid_state": "Invalid authentication state. Please try again.", + "token_validation_failed": "Token validation failed." + } +} +``` + +`modules/keycloak/package.json`: +```json +{ + "name": "@simple-module/keycloak", + "private": true, + "version": "0.0.0", + "type": "module", + "dependencies": {} +} +``` + +`modules/keycloak/tsconfig.json`: +```json +{ + "extends": "../../host/client_app/tsconfig.json", + "include": ["keycloak/**/*.ts", "keycloak/**/*.tsx"] +} +``` + +`modules/keycloak/tests/__init__.py`: empty file. + +`modules/keycloak/tests/conftest.py`: +```python +"""Keycloak module test fixtures.""" +``` + +- [ ] **Step 4: Create KeycloakSettings** + +```python +# modules/keycloak/keycloak/settings.py +"""Keycloak module settings — DB-backed via ``register_module_settings``.""" + +from __future__ import annotations + +from pydantic import Field, model_validator +from pydantic_settings import BaseSettings, SettingsConfigDict +from simple_module_core.dotenv import env_str +from simple_module_core.environments import NON_PROD_ENVIRONMENTS + + +class KeycloakSettings(BaseSettings): + """Keycloak OIDC configuration.""" + + model_config = SettingsConfigDict(extra="ignore") + + server_url: str = env_str("SM_KEYCLOAK_SERVER_URL", "") + realm: str = env_str("SM_KEYCLOAK_REALM", "") + client_id: str = env_str("SM_KEYCLOAK_CLIENT_ID", "") + client_secret: str = env_str("SM_KEYCLOAK_CLIENT_SECRET", "") + + roles_claim_path: str = "realm_access.roles" + admin_role: str = "admin" + login_redirect_url: str = "/dashboard/" + jwks_cache_ttl_seconds: int = 3600 + + role_mapping: dict[str, str] = Field( + default_factory=lambda: {"admin": "admin", "user": "user"}, + ) + + @model_validator(mode="after") + def _check_required_in_production(self) -> KeycloakSettings: + import os + + env = os.environ.get("SM_ENVIRONMENT", "development") + if env in NON_PROD_ENVIRONMENTS: + return self + missing = [] + if not self.server_url: + missing.append("SM_KEYCLOAK_SERVER_URL") + if not self.realm: + missing.append("SM_KEYCLOAK_REALM") + if not self.client_id: + missing.append("SM_KEYCLOAK_CLIENT_ID") + if not self.client_secret: + missing.append("SM_KEYCLOAK_CLIENT_SECRET") + if missing: + msg = f"Keycloak settings required in production: {', '.join(missing)}" + raise ValueError(msg) + return self +``` + +- [ ] **Step 5: Create KeycloakState** + +```python +# modules/keycloak/keycloak/state.py +"""Module-scoped state container for the keycloak module.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from keycloak.jwks import JWKSCache + from keycloak.settings import KeycloakSettings + + +@dataclass +class KeycloakState: + """Keycloak-module singletons. Single slot at ``app.state.keycloak``.""" + + settings: KeycloakSettings + jwks_cache: JWKSCache | None = None +``` + +- [ ] **Step 6: Create KeycloakModule (minimal — provider wired in next tasks)** + +```python +# modules/keycloak/keycloak/module.py +"""Keycloak OIDC authentication module.""" + +from __future__ import annotations + +import importlib.resources +from pathlib import Path +from typing import TYPE_CHECKING + +from simple_module_core.menu import MenuItem, MenuRegistry, MenuSection +from simple_module_core.module import ModuleBase, ModuleMeta +from simple_module_core.permissions import PermissionRegistry + +if TYPE_CHECKING: + from fastapi import APIRouter, FastAPI + + +class KeycloakModule(ModuleBase): + meta = ModuleMeta( + name="Keycloak", + route_prefix="/api/keycloak", + view_prefix="/keycloak", + depends_on=["Auth"], + ) + _is_auth_provider = True + + def register_settings(self, app: FastAPI) -> None: + import importlib + + from keycloak.provider import KeycloakAuthProvider + from keycloak.settings import KeycloakSettings + from keycloak.state import KeycloakState + + register_module_settings = importlib.import_module( + "settings.registration" + ).register_module_settings + + register_module_settings( + app, "keycloak", KeycloakSettings, lambda s: KeycloakState(settings=s) + ) + + app.state.auth.auth_provider = KeycloakAuthProvider(app.state.keycloak.settings) + + def register_menu_items(self, registry: MenuRegistry) -> None: + registry.add( + MenuItem( + label="Logout", + url="/keycloak/logout", + icon="log-out", + order=999, + section=MenuSection.USER_DROPDOWN, + method="post", + ) + ) + + def register_routes(self, api_router: APIRouter, view_router: APIRouter) -> None: + from keycloak.endpoints.api import router as api + from keycloak.endpoints.views import router as views + + api_router.include_router(api) + view_router.include_router(views) + + async def on_startup(self, app: FastAPI) -> None: + from keycloak.jwks import JWKSCache + + state = app.state.keycloak + s = state.settings + if s.server_url and s.realm: + state.jwks_cache = JWKSCache( + jwks_url=f"{s.server_url}/realms/{s.realm}/protocol/openid-connect/certs", + ttl_seconds=s.jwks_cache_ttl_seconds, + ) + provider = app.state.auth.auth_provider + provider.jwks_cache = state.jwks_cache + + def locale_dirs(self) -> dict[str, Path]: + return { + "keycloak": Path( + str(importlib.resources.files(__package__) / "locales") + ) + } +``` + +- [ ] **Step 7: Update root pyproject.toml** + +Add `"modules/keycloak"` to `tool.ty.environment.extra-paths` and `"modules/keycloak/tests"` to `tool.pytest.ini_options.testpaths`. + +- [ ] **Step 8: Run `uv sync --all-packages` then tests** + +Run: `uv sync --all-packages && uv run pytest modules/keycloak/tests/test_keycloak_module.py -v` +Expected: All pass. + +- [ ] **Step 9: Commit** + +```bash +git add modules/keycloak/ pyproject.toml +git commit -m "feat(keycloak): scaffold keycloak module with settings + provider registration" +``` + +--- + +## Task 9: JWKS Cache + JWT Validation + +**Files:** +- Create: `modules/keycloak/keycloak/jwks.py` +- Create: `modules/keycloak/tests/test_jwks.py` + +- [ ] **Step 1: Write JWKS/JWT tests** + +```python +# modules/keycloak/tests/test_jwks.py +"""Tests for JWKS key cache and JWT validation.""" + +from __future__ import annotations + +import json +import time + +import jwt +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa + +from keycloak.jwks import JWKSCache + + +def _generate_rsa_keypair(): + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_key = private_key.public_key() + return private_key, public_key + + +def _make_jwks_response(public_key, kid="test-key-1"): + from jwt.algorithms import RSAAlgorithm + + jwk = json.loads(RSAAlgorithm.to_jwk(public_key)) + jwk["kid"] = kid + jwk["use"] = "sig" + jwk["alg"] = "RS256" + return {"keys": [jwk]} + + +def _sign_token(private_key, payload, kid="test-key-1"): + return jwt.encode(payload, private_key, algorithm="RS256", headers={"kid": kid}) + + +@pytest.fixture +def rsa_keys(): + return _generate_rsa_keypair() + + +@pytest.fixture +def valid_payload(): + now = int(time.time()) + return { + "sub": "user-123", + "email": "test@example.com", + "preferred_username": "testuser", + "iss": "https://auth.example.com/realms/test", + "aud": "my-client", + "exp": now + 3600, + "iat": now, + "realm_access": {"roles": ["admin", "user"]}, + } + + +async def test_validate_jwt_valid_token(rsa_keys, valid_payload, httpx_mock): + private_key, public_key = rsa_keys + jwks_data = _make_jwks_response(public_key) + httpx_mock.add_response(url="https://auth.example.com/jwks", json=jwks_data) + + cache = JWKSCache( + jwks_url="https://auth.example.com/jwks", + ttl_seconds=3600, + issuer="https://auth.example.com/realms/test", + audience="my-client", + ) + + token = _sign_token(private_key, valid_payload) + claims = await cache.validate_jwt(token) + assert claims is not None + assert claims["sub"] == "user-123" + assert claims["email"] == "test@example.com" + + +async def test_validate_jwt_expired_token(rsa_keys, valid_payload, httpx_mock): + private_key, public_key = rsa_keys + valid_payload["exp"] = int(time.time()) - 100 + jwks_data = _make_jwks_response(public_key) + httpx_mock.add_response(url="https://auth.example.com/jwks", json=jwks_data) + + cache = JWKSCache( + jwks_url="https://auth.example.com/jwks", + ttl_seconds=3600, + issuer="https://auth.example.com/realms/test", + audience="my-client", + ) + + token = _sign_token(private_key, valid_payload) + claims = await cache.validate_jwt(token) + assert claims is None + + +async def test_validate_jwt_wrong_issuer(rsa_keys, valid_payload, httpx_mock): + private_key, public_key = rsa_keys + jwks_data = _make_jwks_response(public_key) + httpx_mock.add_response(url="https://auth.example.com/jwks", json=jwks_data) + + cache = JWKSCache( + jwks_url="https://auth.example.com/jwks", + ttl_seconds=3600, + issuer="https://wrong-issuer.com/realms/test", + audience="my-client", + ) + + token = _sign_token(private_key, valid_payload) + claims = await cache.validate_jwt(token) + assert claims is None + + +async def test_validate_jwt_wrong_audience(rsa_keys, valid_payload, httpx_mock): + private_key, public_key = rsa_keys + jwks_data = _make_jwks_response(public_key) + httpx_mock.add_response(url="https://auth.example.com/jwks", json=jwks_data) + + cache = JWKSCache( + jwks_url="https://auth.example.com/jwks", + ttl_seconds=3600, + issuer="https://auth.example.com/realms/test", + audience="wrong-client", + ) + + token = _sign_token(private_key, valid_payload) + claims = await cache.validate_jwt(token) + assert claims is None + + +async def test_jwks_cache_refetches_on_unknown_kid(rsa_keys, valid_payload, httpx_mock): + """When a token has a kid not in cache, refetch JWKS once before rejecting.""" + private_key, public_key = rsa_keys + jwks_data = _make_jwks_response(public_key, kid="rotated-key") + httpx_mock.add_response(url="https://auth.example.com/jwks", json={"keys": []}) + httpx_mock.add_response(url="https://auth.example.com/jwks", json=jwks_data) + + cache = JWKSCache( + jwks_url="https://auth.example.com/jwks", + ttl_seconds=3600, + issuer="https://auth.example.com/realms/test", + audience="my-client", + ) + + token = _sign_token(private_key, valid_payload, kid="rotated-key") + claims = await cache.validate_jwt(token) + assert claims is not None + assert claims["sub"] == "user-123" +``` + +Note: These tests require `pytest-httpx` for mocking. Add to dev dependencies if not already present, or use `httpx_mock` fixture from `pytest-httpx`. Alternatively, mock `httpx.AsyncClient.get` directly if `pytest-httpx` is not available. + +- [ ] **Step 2: Run tests — should fail** + +Run: `uv run pytest modules/keycloak/tests/test_jwks.py -v` +Expected: `ImportError` — `keycloak.jwks` doesn't exist. + +- [ ] **Step 3: Implement JWKSCache** + +```python +# modules/keycloak/keycloak/jwks.py +"""JWKS key cache and JWT validation for Keycloak tokens.""" + +from __future__ import annotations + +import logging +import time +from typing import Any + +import httpx +import jwt +from jwt.algorithms import RSAAlgorithm + +logger = logging.getLogger(__name__) + + +class JWKSCache: + """Caches Keycloak's public signing keys and validates JWTs. + + On validation failure with cached keys, refetches JWKS once before + rejecting — this handles Keycloak key rotation gracefully. + """ + + def __init__( + self, + jwks_url: str, + ttl_seconds: int = 3600, + issuer: str = "", + audience: str = "", + ) -> None: + self._jwks_url = jwks_url + self._ttl = ttl_seconds + self._issuer = issuer + self._audience = audience + self._keys: dict[str, Any] = {} + self._fetched_at: float = 0 + + async def validate_jwt(self, token: str) -> dict[str, Any] | None: + """Decode and validate a JWT. Returns claims dict or None.""" + try: + unverified = jwt.get_unverified_header(token) + except jwt.exceptions.DecodeError: + return None + + kid = unverified.get("kid") + if kid is None: + return None + + key = await self._get_key(kid) + if key is None: + return None + + return self._decode(token, key) + + def _decode(self, token: str, key: Any) -> dict[str, Any] | None: + try: + return jwt.decode( + token, + key, + algorithms=["RS256"], + issuer=self._issuer if self._issuer else None, + audience=self._audience if self._audience else None, + options={ + "verify_iss": bool(self._issuer), + "verify_aud": bool(self._audience), + }, + ) + except (jwt.ExpiredSignatureError, jwt.InvalidIssuerError, jwt.InvalidAudienceError): + return None + except jwt.PyJWTError: + logger.exception("JWT validation failed") + return None + + async def _get_key(self, kid: str) -> Any | None: + if self._is_stale() or kid not in self._keys: + await self._fetch_keys() + + if kid in self._keys: + return self._keys[kid] + + # Key rotation: refetch once more if kid still missing + await self._fetch_keys(force=True) + return self._keys.get(kid) + + def _is_stale(self) -> bool: + return time.monotonic() - self._fetched_at > self._ttl + + async def _fetch_keys(self, *, force: bool = False) -> None: + if not force and not self._is_stale(): + return + try: + async with httpx.AsyncClient() as client: + resp = await client.get(self._jwks_url, timeout=10) + resp.raise_for_status() + jwks_data = resp.json() + except Exception: + logger.exception("Failed to fetch JWKS from %s", self._jwks_url) + return + + new_keys: dict[str, Any] = {} + for key_data in jwks_data.get("keys", []): + kid = key_data.get("kid") + if kid and key_data.get("alg") == "RS256": + try: + public_key = RSAAlgorithm.from_jwk(key_data) + new_keys[kid] = public_key + except Exception: + logger.warning("Failed to parse JWK kid=%s", kid) + self._keys = new_keys + self._fetched_at = time.monotonic() +``` + +- [ ] **Step 4: Run tests — should pass** + +Run: `uv run pytest modules/keycloak/tests/test_jwks.py -v` +Expected: All 5 tests pass. (If `pytest-httpx` is not installed, install it: `uv add --dev pytest-httpx` or use inline mocking.) + +- [ ] **Step 5: Commit** + +```bash +git add modules/keycloak/keycloak/jwks.py modules/keycloak/tests/test_jwks.py +git commit -m "feat(keycloak): JWKS key cache with JWT validation and key-rotation retry" +``` + +--- + +## Task 10: OIDC Discovery + Token Exchange + +**Files:** +- Create: `modules/keycloak/keycloak/oidc.py` +- Create: `modules/keycloak/tests/test_oidc.py` + +- [ ] **Step 1: Write OIDC helper tests** + +```python +# modules/keycloak/tests/test_oidc.py +"""Tests for OIDC discovery and token exchange helpers.""" + +from __future__ import annotations + +import secrets + +import pytest +from keycloak.oidc import OIDCClient + + +@pytest.fixture +def oidc_client(): + return OIDCClient( + server_url="https://auth.example.com", + realm="test", + client_id="my-app", + client_secret="secret123", + ) + + +def test_authorization_url(oidc_client): + url, state = oidc_client.build_authorization_url( + redirect_uri="https://app.example.com/callback", + nonce="test-nonce", + ) + assert "auth.example.com/realms/test/protocol/openid-connect/auth" in url + assert "client_id=my-app" in url + assert "redirect_uri=" in url + assert "response_type=code" in url + assert "scope=openid" in url + assert "nonce=test-nonce" in url + assert state is not None + assert len(state) > 0 + + +def test_token_endpoint_url(oidc_client): + assert oidc_client.token_endpoint == ( + "https://auth.example.com/realms/test/protocol/openid-connect/token" + ) + + +def test_logout_url(oidc_client): + url = oidc_client.build_logout_url( + post_logout_redirect_uri="https://app.example.com/login", + id_token_hint="token123", + ) + assert "auth.example.com/realms/test/protocol/openid-connect/logout" in url + assert "post_logout_redirect_uri=" in url + assert "id_token_hint=token123" in url + + +def test_issuer(oidc_client): + assert oidc_client.issuer == "https://auth.example.com/realms/test" + + +async def test_exchange_code(oidc_client, httpx_mock): + httpx_mock.add_response( + url=oidc_client.token_endpoint, + json={ + "access_token": "at-123", + "id_token": "id-123", + "refresh_token": "rt-123", + "token_type": "Bearer", + "expires_in": 300, + }, + ) + tokens = await oidc_client.exchange_code( + code="auth-code-xyz", + redirect_uri="https://app.example.com/callback", + ) + assert tokens["access_token"] == "at-123" + assert tokens["id_token"] == "id-123" +``` + +- [ ] **Step 2: Run tests — should fail** + +Run: `uv run pytest modules/keycloak/tests/test_oidc.py -v` +Expected: `ImportError`. + +- [ ] **Step 3: Implement OIDCClient** + +```python +# modules/keycloak/keycloak/oidc.py +"""OIDC helpers for Keycloak — authorization URL, token exchange, logout.""" + +from __future__ import annotations + +import secrets +from typing import Any +from urllib.parse import urlencode + +import httpx + + +class OIDCClient: + """Thin wrapper around Keycloak's OIDC endpoints.""" + + def __init__( + self, + server_url: str, + realm: str, + client_id: str, + client_secret: str, + ) -> None: + self._base = f"{server_url.rstrip('/')}/realms/{realm}/protocol/openid-connect" + self._client_id = client_id + self._client_secret = client_secret + self._server_url = server_url.rstrip("/") + self._realm = realm + + @property + def issuer(self) -> str: + return f"{self._server_url}/realms/{self._realm}" + + @property + def token_endpoint(self) -> str: + return f"{self._base}/token" + + @property + def jwks_url(self) -> str: + return f"{self._base}/certs" + + def build_authorization_url( + self, + redirect_uri: str, + nonce: str, + scope: str = "openid email profile", + ) -> tuple[str, str]: + """Build the OIDC authorization URL. Returns (url, state).""" + state = secrets.token_urlsafe(32) + params = { + "client_id": self._client_id, + "redirect_uri": redirect_uri, + "response_type": "code", + "scope": scope, + "state": state, + "nonce": nonce, + } + url = f"{self._base}/auth?{urlencode(params)}" + return url, state + + async def exchange_code( + self, + code: str, + redirect_uri: str, + ) -> dict[str, Any]: + """Exchange an authorization code for tokens.""" + data = { + "grant_type": "authorization_code", + "code": code, + "redirect_uri": redirect_uri, + "client_id": self._client_id, + "client_secret": self._client_secret, + } + async with httpx.AsyncClient() as client: + resp = await client.post(self.token_endpoint, data=data, timeout=10) + resp.raise_for_status() + return resp.json() + + def build_logout_url( + self, + post_logout_redirect_uri: str, + id_token_hint: str | None = None, + ) -> str: + params: dict[str, str] = { + "post_logout_redirect_uri": post_logout_redirect_uri, + } + if id_token_hint: + params["id_token_hint"] = id_token_hint + return f"{self._base}/logout?{urlencode(params)}" +``` + +- [ ] **Step 4: Run tests — should pass** + +Run: `uv run pytest modules/keycloak/tests/test_oidc.py -v` +Expected: All 5 tests pass. + +- [ ] **Step 5: Commit** + +```bash +git add modules/keycloak/keycloak/oidc.py modules/keycloak/tests/test_oidc.py +git commit -m "feat(keycloak): OIDC client with authorization URL, token exchange, logout" +``` + +--- + +## Task 11: KeycloakUserCache Model + Migration + +**Files:** +- Create: `modules/keycloak/keycloak/models.py` +- Create migration via Alembic + +- [ ] **Step 1: Create the model** + +```python +# modules/keycloak/keycloak/models.py +"""Keycloak user cache — maps Keycloak sub to a stable framework UUID.""" + +from __future__ import annotations + +import uuid as uuid_mod +from datetime import datetime + +from simple_module_db.base import create_module_base +from sqlmodel import Field + +Base = create_module_base("keycloak") + + +class KeycloakUserCache(Base, table=True): + __tablename__ = "keycloak_user_cache" + + id: uuid_mod.UUID = Field(default_factory=uuid_mod.uuid4, primary_key=True) + keycloak_sub: str = Field(unique=True, index=True) + email: str = "" + full_name: str | None = None + last_login_at: datetime | None = None +``` + +- [ ] **Step 2: Generate migration** + +Run: `make migration msg="add keycloak_user_cache table"` +Expected: New migration file created in `host/migrations/versions/`. + +- [ ] **Step 3: Apply migration** + +Run: `make migrate` +Expected: Migration applies cleanly. + +- [ ] **Step 4: Commit** + +```bash +git add modules/keycloak/keycloak/models.py host/migrations/versions/*keycloak* +git commit -m "feat(keycloak): KeycloakUserCache model + migration" +``` + +--- + +## Task 12: KeycloakAuthProvider Implementation + +**Files:** +- Create: `modules/keycloak/keycloak/provider.py` +- Create: `modules/keycloak/tests/test_keycloak_provider.py` + +- [ ] **Step 1: Write provider tests** + +```python +# modules/keycloak/tests/test_keycloak_provider.py +"""Tests for KeycloakAuthProvider.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest +from auth.contracts.provider import AuthProvider +from auth.contracts.schemas import UserContext +from keycloak.provider import KeycloakAuthProvider +from keycloak.settings import KeycloakSettings + + +@pytest.fixture +def settings(): + return KeycloakSettings( + server_url="https://auth.example.com", + realm="test", + client_id="my-app", + client_secret="secret", + role_mapping={"admin": "admin", "user": "user", "editor": "editor"}, + ) + + +@pytest.fixture +def provider(settings): + return KeycloakAuthProvider(settings) + + +def test_satisfies_protocol(provider): + assert isinstance(provider, AuthProvider) + + +def test_name(provider): + assert provider.name == "keycloak" + + +def test_login_url(provider): + assert provider.get_login_url(None) == "/keycloak/login" + + +def test_logout_url(provider): + assert provider.get_logout_url(None) == "/keycloak/logout" + + +def test_public_paths(provider): + prefixes, exact = provider.get_public_paths() + assert "/keycloak/login" in prefixes + assert "/api/keycloak/auth/" in prefixes + + +def test_is_bearer_request(provider): + req = MagicMock() + req.headers = {"authorization": "Bearer abc"} + assert provider.is_bearer_request(req) is True + + req.headers = {} + assert provider.is_bearer_request(req) is False + + +def test_claims_to_user_context(provider): + claims = { + "sub": "kc-user-123", + "email": "test@example.com", + "preferred_username": "testuser", + "realm_access": {"roles": ["admin", "unknown_role", "user"]}, + } + ctx = provider._claims_to_user_context(claims, cache_id="aaaaaaaa-0000-0000-0000-000000000001") + assert isinstance(ctx, UserContext) + assert ctx.id == "aaaaaaaa-0000-0000-0000-000000000001" + assert ctx.email == "test@example.com" + assert ctx.name == "testuser" + assert sorted(ctx.roles) == ["admin", "user"] + # "unknown_role" not in mapping, so excluded + + +def test_claims_to_user_context_no_roles(provider): + claims = {"sub": "kc-user-456", "email": "noroles@example.com"} + ctx = provider._claims_to_user_context(claims, cache_id="bbbb") + assert ctx.roles == [] + + +def test_extract_roles_custom_claim_path(settings): + settings.roles_claim_path = "resource_access.my-app.roles" + provider = KeycloakAuthProvider(settings) + claims = {"sub": "x", "resource_access": {"my-app": {"roles": ["admin"]}}} + ctx = provider._claims_to_user_context(claims, cache_id="cccc") + assert ctx.roles == ["admin"] +``` + +- [ ] **Step 2: Run tests — should fail** + +Run: `uv run pytest modules/keycloak/tests/test_keycloak_provider.py -v` +Expected: `ImportError`. + +- [ ] **Step 3: Implement KeycloakAuthProvider** + +```python +# modules/keycloak/keycloak/provider.py +"""KeycloakAuthProvider — resolves users from Keycloak JWTs or session.""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Any + +from auth.contracts.schemas import UserContext +from starlette.requests import Request + +if TYPE_CHECKING: + from keycloak.jwks import JWKSCache + from keycloak.settings import KeycloakSettings + +logger = logging.getLogger(__name__) + +_SESSION_USER_CTX_KEY = "user_ctx" + + +class KeycloakAuthProvider: + """OIDC auth provider backed by Keycloak.""" + + name = "keycloak" + _is_auth_provider = True + + def __init__(self, settings: KeycloakSettings) -> None: + self._settings = settings + self.jwks_cache: JWKSCache | None = None + + async def resolve_user(self, request: Request) -> UserContext | None: + auth_header = request.headers.get("authorization", "") + if auth_header.startswith("Bearer "): + return await self._resolve_bearer(request, auth_header[7:]) + + session = request.scope.get("session", {}) + return UserContext.from_session_dict(session.get(_SESSION_USER_CTX_KEY)) + + def get_login_url(self, request: Request | None, next_url: str | None = None) -> str: + return "/keycloak/login" + + def get_logout_url(self, request: Request | None) -> str: + return "/keycloak/logout" + + def get_public_paths(self) -> tuple[tuple[str, ...], tuple[str, ...]]: + return ( + ("/keycloak/login", "/keycloak/logout", "/api/keycloak/auth/"), + (), + ) + + def is_bearer_request(self, request: Request | None) -> bool: + if request is None: + return False + return request.headers.get("authorization", "").startswith("Bearer ") + + async def _resolve_bearer(self, request: Request, token: str) -> UserContext | None: + if self.jwks_cache is None: + logger.warning("JWKS cache not initialized; rejecting bearer token") + return None + claims = await self.jwks_cache.validate_jwt(token) + if claims is None: + return None + + cache_id = await self._upsert_user_cache(request, claims) + return self._claims_to_user_context(claims, cache_id=cache_id) + + def _claims_to_user_context( + self, + claims: dict[str, Any], + *, + cache_id: str, + ) -> UserContext: + roles_raw = _extract_nested(claims, self._settings.roles_claim_path) + mapped = [ + self._settings.role_mapping[r] + for r in (roles_raw or []) + if r in self._settings.role_mapping + ] + return UserContext( + id=cache_id, + email=claims.get("email", ""), + name=claims.get("preferred_username") or claims.get("name", ""), + roles=mapped, + tenant_id=claims.get("tenant_id"), + ) + + async def _upsert_user_cache(self, request: Request, claims: dict) -> str: + """Upsert KeycloakUserCache and return its UUID as string.""" + try: + from keycloak.models import KeycloakUserCache + from sqlalchemy import select + + session_factory = request.app.state.sm.db.session_factory + sub = claims["sub"] + async with session_factory() as db: + stmt = select(KeycloakUserCache).where( + KeycloakUserCache.keycloak_sub == sub + ) + row = (await db.execute(stmt)).scalar_one_or_none() + if row is None: + import uuid as uuid_mod + from datetime import datetime, timezone + + row = KeycloakUserCache( + id=uuid_mod.uuid4(), + keycloak_sub=sub, + email=claims.get("email", ""), + full_name=claims.get("preferred_username"), + last_login_at=datetime.now(timezone.utc), + ) + db.add(row) + await db.flush() + else: + from datetime import datetime, timezone + + row.email = claims.get("email", row.email) + row.full_name = claims.get("preferred_username", row.full_name) + row.last_login_at = datetime.now(timezone.utc) + await db.flush() + return str(row.id) + except Exception: + logger.exception("Failed to upsert KeycloakUserCache for sub=%s", claims.get("sub")) + return claims.get("sub", "unknown") + + +def _extract_nested(data: dict, path: str) -> list[str] | None: + """Extract a value from a nested dict using a dot-separated path.""" + parts = path.split(".") + current: Any = data + for part in parts: + if not isinstance(current, dict): + return None + current = current.get(part) + if current is None: + return None + return current if isinstance(current, list) else None +``` + +- [ ] **Step 4: Run tests — should pass** + +Run: `uv run pytest modules/keycloak/tests/test_keycloak_provider.py -v` +Expected: All tests pass. + +- [ ] **Step 5: Commit** + +```bash +git add modules/keycloak/keycloak/provider.py modules/keycloak/tests/test_keycloak_provider.py +git commit -m "feat(keycloak): KeycloakAuthProvider with JWT resolution and role mapping" +``` + +--- + +## Task 13: Keycloak Endpoints (API + Views) + +**Files:** +- Create: `modules/keycloak/keycloak/endpoints/api.py` +- Create: `modules/keycloak/keycloak/endpoints/__init__.py` +- Create: `modules/keycloak/keycloak/endpoints/views.py` +- Create: `modules/keycloak/keycloak/pages/Login.tsx` +- Create: `modules/keycloak/keycloak/pages/LoggedOut.tsx` + +- [ ] **Step 1: Create endpoints `__init__.py`** + +```python +# modules/keycloak/keycloak/endpoints/__init__.py +"""Keycloak endpoint routers.""" +``` + +- [ ] **Step 2: Create API endpoints** + +```python +# modules/keycloak/keycloak/endpoints/api.py +"""Keycloak OIDC API endpoints — login redirect, callback.""" + +from __future__ import annotations + +import logging +import secrets +from typing import TYPE_CHECKING + +from fastapi import APIRouter, HTTPException, Request +from starlette.responses import RedirectResponse + +if TYPE_CHECKING: + from keycloak.settings import KeycloakSettings + +logger = logging.getLogger(__name__) +router = APIRouter(prefix="/auth", tags=["keycloak-auth"]) + +_SESSION_OIDC_STATE = "keycloak_oidc_state" +_SESSION_OIDC_NONCE = "keycloak_oidc_nonce" +_SESSION_USER_CTX = "user_ctx" +_SESSION_ID_TOKEN = "keycloak_id_token" +_SESSION_NEXT = "next" + + +def _get_settings(request: Request) -> KeycloakSettings: + return request.app.state.keycloak.settings + + +def _get_oidc_client(request: Request): + from keycloak.oidc import OIDCClient + + s = _get_settings(request) + return OIDCClient( + server_url=s.server_url, + realm=s.realm, + client_id=s.client_id, + client_secret=s.client_secret, + ) + + +@router.get("/login") +async def oidc_login(request: Request): + """Redirect to Keycloak's authorization endpoint.""" + client = _get_oidc_client(request) + callback_url = str(request.url_for("oidc_callback")) + nonce = secrets.token_urlsafe(32) + url, state = client.build_authorization_url( + redirect_uri=callback_url, + nonce=nonce, + ) + request.session[_SESSION_OIDC_STATE] = state + request.session[_SESSION_OIDC_NONCE] = nonce + return RedirectResponse(url, status_code=302) + + +@router.get("/callback") +async def oidc_callback(request: Request): + """Handle Keycloak's authorization code callback.""" + code = request.query_params.get("code") + state = request.query_params.get("state") + + expected_state = request.session.pop(_SESSION_OIDC_STATE, None) + nonce = request.session.pop(_SESSION_OIDC_NONCE, None) + + if not code or not state or state != expected_state: + raise HTTPException(status_code=400, detail="Invalid OIDC state") + + client = _get_oidc_client(request) + callback_url = str(request.url_for("oidc_callback")) + + try: + tokens = await client.exchange_code(code=code, redirect_uri=callback_url) + except Exception: + logger.exception("Token exchange failed") + raise HTTPException(status_code=502, detail="Token exchange failed") + + id_token = tokens.get("id_token", "") + access_token = tokens.get("access_token", "") + + jwks_cache = request.app.state.keycloak.jwks_cache + claims = await jwks_cache.validate_jwt(access_token) if jwks_cache else None + if claims is None: + raise HTTPException(status_code=401, detail="Token validation failed") + + provider = request.app.state.auth.auth_provider + cache_id = await provider._upsert_user_cache(request, claims) + user_ctx = provider._claims_to_user_context(claims, cache_id=cache_id) + + request.session[_SESSION_USER_CTX] = user_ctx.to_session_dict() + request.session[_SESSION_ID_TOKEN] = id_token + + s = _get_settings(request) + next_url = request.session.pop(_SESSION_NEXT, None) or s.login_redirect_url + return RedirectResponse(next_url, status_code=303) +``` + +- [ ] **Step 3: Create view endpoints** + +```python +# modules/keycloak/keycloak/endpoints/views.py +"""Keycloak Inertia view routes — login page, logout.""" + +from __future__ import annotations + +from fastapi import APIRouter, Request +from simple_module_hosting.inertia_deps import InertiaDep +from starlette.responses import RedirectResponse + +router = APIRouter(tags=["keycloak-views"]) + +_SESSION_USER_CTX = "user_ctx" +_SESSION_ID_TOKEN = "keycloak_id_token" + + +@router.get("/login") +async def login_page(request: Request, inertia: InertiaDep): + """Render a minimal login page that can auto-redirect to Keycloak.""" + return inertia.render("Keycloak/Login") + + +@router.post("/logout") +async def logout(request: Request): + """Clear framework session and redirect to Keycloak's logout endpoint.""" + from keycloak.oidc import OIDCClient + + s = request.app.state.keycloak.settings + id_token = request.session.get(_SESSION_ID_TOKEN) + + request.session.clear() + + client = OIDCClient( + server_url=s.server_url, + realm=s.realm, + client_id=s.client_id, + client_secret=s.client_secret, + ) + base_url = str(request.base_url).rstrip("/") + logout_url = client.build_logout_url( + post_logout_redirect_uri=f"{base_url}/keycloak/login", + id_token_hint=id_token, + ) + return RedirectResponse(logout_url, status_code=303) +``` + +- [ ] **Step 4: Create frontend pages** + +`modules/keycloak/keycloak/pages/Login.tsx`: +```tsx +import { router } from "@inertiajs/react"; +import { useEffect } from "react"; + +export default function Login() { + useEffect(() => { + router.get("/api/keycloak/auth/login"); + }, []); + + return ( +
+

Redirecting to identity provider…

+
+ ); +} +``` + +`modules/keycloak/keycloak/pages/LoggedOut.tsx`: +```tsx +import { Link } from "@inertiajs/react"; + +export default function LoggedOut() { + return ( +
+

Signed Out

+

You have been signed out successfully.

+ + Sign in again + +
+ ); +} +``` + +- [ ] **Step 5: Commit** + +```bash +git add modules/keycloak/keycloak/endpoints/ modules/keycloak/keycloak/pages/ +git commit -m "feat(keycloak): OIDC login/callback endpoints + Inertia login/logout pages" +``` + +--- + +## Task 14: Users Module — Bearer Token + Refresh Token Endpoints + +**Files:** +- Create: `modules/users/users/models/refresh_token.py` +- Create: `modules/users/users/auth_local/token_api.py` +- Modify: `modules/users/users/settings.py` +- Modify: `modules/users/users/module.py` (include token_api router) +- Create: `modules/users/tests/test_token_api.py` + +- [ ] **Step 1: Write token endpoint tests** + +```python +# modules/users/tests/test_token_api.py +"""Tests for bearer token endpoints (mobile auth).""" + +from __future__ import annotations + +import pytest + + +async def test_token_login_returns_tokens(client, authenticated_client): + """POST /api/users/auth/token with valid credentials returns token pair.""" + # First create a user via the authenticated_client (admin) + resp = await client.post( + "/api/users/auth/token", + json={"email": "admin@example.com", "password": "Admin1234!"}, + ) + # May need to seed user first — adjust based on test fixtures + assert resp.status_code in (200, 400) # 200 if user exists, 400 if not + if resp.status_code == 200: + data = resp.json() + assert "access_token" in data + assert "refresh_token" in data + assert data["token_type"] == "bearer" + + +async def test_token_login_invalid_credentials(client): + resp = await client.post( + "/api/users/auth/token", + json={"email": "wrong@example.com", "password": "wrong"}, + ) + assert resp.status_code == 401 + + +async def test_token_refresh(client): + """POST /api/users/auth/token/refresh swaps refresh token for new pair.""" + # This test requires a valid refresh token — integration-level + resp = await client.post( + "/api/users/auth/token/refresh", + json={"refresh_token": "invalid-token"}, + ) + assert resp.status_code == 401 +``` + +- [ ] **Step 2: Add `bearer_token_lifetime_seconds` to UsersSettings** + +In `modules/users/users/settings.py`, add after the existing cookie settings: + +```python + bearer_token_lifetime_seconds: int = 60 * 15 # 15 minutes + refresh_token_lifetime_seconds: int = 60 * 60 * 24 * 30 # 30 days +``` + +- [ ] **Step 3: Create RefreshToken model** + +```python +# modules/users/users/models/refresh_token.py +"""Refresh token for mobile/API bearer auth.""" + +from __future__ import annotations + +import uuid as uuid_mod +from datetime import datetime, timezone + +from simple_module_db.base import create_module_base +from sqlmodel import Field + +from users.models.user import Base + + +class RefreshToken(Base, table=True): + __tablename__ = "users_refresh_token" + + token: uuid_mod.UUID = Field(default_factory=uuid_mod.uuid4, primary_key=True) + user_id: uuid_mod.UUID = Field(foreign_key="users_user.id", index=True) + created_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc)) + expires_at: datetime + revoked_at: datetime | None = None +``` + +Update `modules/users/users/models/__init__.py` to export `RefreshToken`. + +- [ ] **Step 4: Create token_api endpoints** + +```python +# modules/users/users/auth_local/token_api.py +"""Bearer token endpoints for mobile/API clients.""" + +from __future__ import annotations + +import uuid as uuid_mod +from datetime import datetime, timedelta, timezone + +from fastapi import APIRouter, Depends, HTTPException, Request +from simple_module_db.deps import get_db +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession +from sqlmodel import SQLModel + +from users.models import User +from users.models.refresh_token import RefreshToken + +router = APIRouter(prefix="/auth", tags=["users-token"]) + + +class TokenRequest(SQLModel): + email: str + password: str + + +class TokenResponse(SQLModel): + access_token: str + refresh_token: str + token_type: str = "bearer" + expires_in: int + + +class RefreshRequest(SQLModel): + refresh_token: str + + +@router.post("/token", response_model=TokenResponse) +async def token_login(body: TokenRequest, request: Request, db: AsyncSession = Depends(get_db)): + """Exchange email+password for access + refresh token pair (mobile auth).""" + from fastapi_users.password import PasswordHelper + + stmt = select(User).where(User.email == body.email) + user = (await db.execute(stmt)).scalar_one_or_none() + if user is None or not user.is_active or user.disabled_at is not None: + raise HTTPException(status_code=401, detail="Invalid credentials") + + helper = PasswordHelper() + verified, _ = helper.verify_and_update(body.password, user.hashed_password) + if not verified: + raise HTTPException(status_code=401, detail="Invalid credentials") + + settings = request.app.state.users.settings + return await _create_token_pair(db, user.id, settings) + + +@router.post("/token/refresh", response_model=TokenResponse) +async def token_refresh( + body: RefreshRequest, request: Request, db: AsyncSession = Depends(get_db) +): + """Exchange a refresh token for a new token pair (rotation).""" + try: + token_uuid = uuid_mod.UUID(body.refresh_token) + except (ValueError, TypeError): + raise HTTPException(status_code=401, detail="Invalid refresh token") + + now = datetime.now(timezone.utc) + stmt = select(RefreshToken).where( + RefreshToken.token == token_uuid, + RefreshToken.revoked_at.is_(None), # type: ignore[union-attr] + RefreshToken.expires_at > now, + ) + rt = (await db.execute(stmt)).scalar_one_or_none() + if rt is None: + raise HTTPException(status_code=401, detail="Invalid or expired refresh token") + + rt.revoked_at = now + await db.flush() + + settings = request.app.state.users.settings + return await _create_token_pair(db, rt.user_id, settings) + + +@router.delete("/token") +async def token_revoke( + body: RefreshRequest, db: AsyncSession = Depends(get_db) +): + """Revoke a refresh token (mobile logout).""" + try: + token_uuid = uuid_mod.UUID(body.refresh_token) + except (ValueError, TypeError): + raise HTTPException(status_code=400, detail="Invalid token format") + + stmt = select(RefreshToken).where(RefreshToken.token == token_uuid) + rt = (await db.execute(stmt)).scalar_one_or_none() + if rt and rt.revoked_at is None: + rt.revoked_at = datetime.now(timezone.utc) + await db.flush() + return {"status": "ok"} + + +async def _create_token_pair( + db: AsyncSession, + user_id: uuid_mod.UUID, + settings, +) -> TokenResponse: + """Create access + refresh token pair and persist the refresh token.""" + from users.models import UserAccessToken + + now = datetime.now(timezone.utc) + + access_token = UserAccessToken( + token=str(uuid_mod.uuid4()), + user_id=user_id, + created_at=now, + ) + db.add(access_token) + + refresh = RefreshToken( + token=uuid_mod.uuid4(), + user_id=user_id, + created_at=now, + expires_at=now + timedelta(seconds=settings.refresh_token_lifetime_seconds), + ) + db.add(refresh) + await db.flush() + + return TokenResponse( + access_token=access_token.token, + refresh_token=str(refresh.token), + token_type="bearer", + expires_in=settings.bearer_token_lifetime_seconds, + ) +``` + +- [ ] **Step 5: Wire token_api router into UsersModule.register_routes** + +In `modules/users/users/module.py`, inside `register_routes`, add: + +```python + from users.auth_local.token_api import router as token_router + api_router.include_router(token_router) +``` + +- [ ] **Step 6: Generate migration for refresh_token table** + +Run: `make migration msg="add users_refresh_token table"` + +- [ ] **Step 7: Run tests** + +Run: `uv run pytest modules/users/tests/test_token_api.py -v` +Expected: Tests pass. + +- [ ] **Step 8: Commit** + +```bash +git add modules/users/users/models/refresh_token.py modules/users/users/auth_local/token_api.py modules/users/users/settings.py modules/users/users/module.py modules/users/users/models/__init__.py host/migrations/versions/*refresh_token* +git commit -m "feat(users): bearer token + refresh token endpoints for mobile auth" +``` + +--- + +## Task 15: Integration Tests + Full Suite Verification + +**Files:** +- Create: `tests/integration/test_pluggable_auth.py` + +- [ ] **Step 1: Write integration tests** + +```python +# tests/integration/test_pluggable_auth.py +"""Integration tests for pluggable auth — verifying both providers work.""" + +from __future__ import annotations + +from auth.contracts.provider import AuthProvider +from auth.state import AuthState + + +def test_users_module_is_auth_provider(): + from users.module import UsersModule + + assert UsersModule._is_auth_provider is True + + +def test_keycloak_module_is_auth_provider(): + from keycloak.module import KeycloakModule + + assert KeycloakModule._is_auth_provider is True + + +def test_sm020_fires_with_both_modules(): + from simple_module_core.diagnostics._module import ModuleDiagnostics + from keycloak.module import KeycloakModule + from users.module import UsersModule + + diags = ModuleDiagnostics() + results = diags._check_auth_provider_conflict([UsersModule(), KeycloakModule()]) + assert any(d.code == "SM020" for d in results) + + +def test_sm021_fires_with_neither(): + from simple_module_core.diagnostics._module import ModuleDiagnostics + from simple_module_core.module import ModuleBase, ModuleMeta + + class StubModule(ModuleBase): + meta = ModuleMeta(name="Stub") + + diags = ModuleDiagnostics() + results = diags._check_auth_provider_conflict([StubModule()]) + assert any(d.code == "SM021" for d in results) + + +def test_auth_provider_protocol_satisfied_by_users(): + from users.provider import UsersAuthProvider + + assert isinstance(UsersAuthProvider(), AuthProvider) + + +def test_auth_provider_protocol_satisfied_by_keycloak(): + from keycloak.provider import KeycloakAuthProvider + from keycloak.settings import KeycloakSettings + + settings = KeycloakSettings( + server_url="https://example.com", + realm="test", + client_id="app", + client_secret="secret", + ) + assert isinstance(KeycloakAuthProvider(settings), AuthProvider) +``` + +- [ ] **Step 2: Run integration tests** + +Run: `uv run pytest tests/integration/test_pluggable_auth.py -v` +Expected: All pass. + +- [ ] **Step 3: Run full test suite** + +Run: `make test` +Expected: All tests pass. No regressions in existing users, auth, or framework tests. + +- [ ] **Step 4: Run linter** + +Run: `make lint` +Expected: Clean. + +- [ ] **Step 5: Commit** + +```bash +git add tests/integration/test_pluggable_auth.py +git commit -m "test: integration tests for pluggable auth provider system" +``` + +--- + +## Task 16: Documentation Update + +**Files:** +- Modify: `docs/framework-conventions.md` (add auth provider section) +- The spec document is already committed + +- [ ] **Step 1: Add auth provider section to framework conventions** + +Add a section to `docs/framework-conventions.md` under the existing auth documentation: + +```markdown +### Auth Provider Contract + +The framework supports swappable authentication backends. Exactly one auth provider +module must be installed — either `simple-module-users` (local credentials + OAuth) +or `simple-module-keycloak` (Keycloak OIDC). Both implement the `AuthProvider` +protocol from `auth.contracts.provider`. + +**Module authors never import from `users` or `keycloak` directly.** Use only: +- `from auth.deps import CurrentUser, require_permission` +- `from auth.contracts.schemas import UserContext` + +The `AuthMiddleware` (in `auth/middleware.py`) delegates to the active provider's +`resolve_user()` method, then falls through to the principal-resolver chain. +API paths (`/api/*`) receive 401 JSON when unauthenticated; view paths receive +a 302 redirect to the provider's login URL. + +Boot-time diagnostic `SM020` fails if multiple auth providers are installed. +`SM021` warns if none is installed. +``` + +- [ ] **Step 2: Update CLAUDE.md diagnostic codes table** + +Add to the diagnostic codes section: +``` +`SM020` multiple auth provider modules installed (error), `SM021` no auth provider module installed (warn) +``` + +- [ ] **Step 3: Commit** + +```bash +git add docs/framework-conventions.md CLAUDE.md +git commit -m "docs: document pluggable auth provider contract and SM020/SM021 diagnostics" +``` + +--- + +## Summary of Commit Sequence + +1. `feat(auth): add AuthProvider protocol for swappable auth backends` +2. `feat(auth): add auth_provider slot to AuthState` +3. `feat(auth): add provider-agnostic AuthMiddleware` +4. `refactor(auth,users): move AuthMiddleware + principal_serializer to auth module` +5. `feat(users): implement UsersAuthProvider with session-cookie resolution` +6. `refactor(users): delegate to auth.middleware, keep thin re-export for compat` +7. `feat(diagnostics): SM020/SM021 — exactly one auth provider required` +8. `feat(keycloak): scaffold keycloak module with settings + provider registration` +9. `feat(keycloak): JWKS key cache with JWT validation and key-rotation retry` +10. `feat(keycloak): OIDC client with authorization URL, token exchange, logout` +11. `feat(keycloak): KeycloakUserCache model + migration` +12. `feat(keycloak): KeycloakAuthProvider with JWT resolution and role mapping` +13. `feat(keycloak): OIDC login/callback endpoints + Inertia login/logout pages` +14. `feat(users): bearer token + refresh token endpoints for mobile auth` +15. `test: integration tests for pluggable auth provider system` +16. `docs: document pluggable auth provider contract and SM020/SM021 diagnostics` diff --git a/docs/superpowers/specs/2026-05-27-pluggable-auth-keycloak-design.md b/docs/superpowers/specs/2026-05-27-pluggable-auth-keycloak-design.md new file mode 100644 index 00000000..1cb25e03 --- /dev/null +++ b/docs/superpowers/specs/2026-05-27-pluggable-auth-keycloak-design.md @@ -0,0 +1,538 @@ +# Pluggable Auth: AuthProvider Contract + Keycloak Module + +**Date:** 2026-05-27 +**Status:** Design — approved for implementation planning +**Builds on:** [2026-05-21 Auth principal-resolver chain](2026-05-21-auth-principal-resolver-design.md) (must land first) + +## Goal + +Make the framework's authentication layer swappable. Framework users install either the existing `users` module (local credentials + OAuth) or a new `keycloak` module (Keycloak OIDC) — both implement the same `AuthProvider` contract so every other module is unaffected. Both web and mobile clients are supported regardless of which provider is active. + +## Relationship to the Principal-Resolver Spec + +The [principal-resolver design](2026-05-21-auth-principal-resolver-design.md) (issue #163) adds the extension point for bearer-token resolution and the 401-JSON-for-API behavior. That spec is a **prerequisite** for this one — it lands the resolver chain and the API-vs-browser response split. This design builds on top: + +- The principal-resolver chain becomes one of the tools an `AuthProvider` uses internally (the `users` provider registers its cookie resolver as the primary path and lets downstream modules add PAT resolvers via the chain). +- The `keycloak` provider uses its own `resolve_user` that validates Keycloak JWTs for bearer tokens and reads session for browser flows — it doesn't use the resolver chain for its core flow, but the chain remains available for downstream modules that want to add extra credential types on top of Keycloak (e.g., service-account tokens). +- The 401-JSON-for-`/api/*` behavior from the resolver spec carries forward unchanged. + +## Scope + +**In scope.** +- `AuthProvider` protocol in `auth/contracts/` — the interface both providers implement. +- Provider-agnostic `AuthMiddleware` extracted from `users/middleware.py` into `auth/`. +- Bearer token transport for mobile clients (framework-issued tokens with `users`, Keycloak-issued JWTs with `keycloak`). +- New `keycloak` module package: OIDC login (web redirect + mobile PKCE), JWT validation against JWKS, role mapping, lightweight user cache table. +- Conflict diagnostic (`SM020`) — boot fails if both `users` and `keycloak` are installed. +- `principal_serializer` registration moved from `users/module.py` into `auth/` so it works with either provider. + +**Out of scope.** +- Full OAuth2 authorization server (client registration, scopes, token introspection for third-party apps). +- Keycloak admin API integration (realm/client provisioning from the framework). +- Dual-provider mode (authenticating some users via `users` and others via `keycloak` simultaneously). +- Migration tooling to move existing `users` data into Keycloak. + +## Non-goals & invariants + +- **`auth` module stays the stable contract layer.** Other modules import only from `auth.contracts` and `auth.deps` — never from `users` or `keycloak` directly. +- **`UserContext` is the single identity type.** Every downstream consumer (menus, permissions, audit listeners, Inertia shared props) operates on `UserContext` regardless of provider. +- **`PermissionRegistry` is framework-owned.** Both providers map their roles into the same registry. Permission definitions live in application modules, not in the identity provider. +- **Session cookie (`session`) stays Starlette-managed.** Both providers use the framework's `SessionMiddleware` for server-side state (redirect targets, OIDC nonces). The auth *credential* transport differs (cookie vs. bearer), but the session layer is shared. + +## Architecture + +Five pieces, built in dependency order: + +### 1. AuthProvider Protocol (in `auth/contracts/provider.py`) + +```python +from __future__ import annotations +from typing import Protocol, runtime_checkable +from starlette.requests import Request +from auth.contracts.schemas import UserContext + +@runtime_checkable +class AuthProvider(Protocol): + name: str + + async def resolve_user(self, request: Request) -> UserContext | None: + """Extract authenticated user from the request. + + For cookie-based providers: read session/cookie, validate, return UserContext. + For token-based providers: decode Authorization header, validate, return UserContext. + Returns None if no valid credential is present. + """ + ... + + def get_login_url(self, request: Request, next_url: str | None = None) -> str: + """URL to redirect unauthenticated browser requests to.""" + ... + + def get_logout_url(self, request: Request) -> str: + """URL/endpoint to POST for logout. + + Keycloak needs RP-initiated logout (redirect to Keycloak's /logout). + Users module clears session + cookie locally. + """ + ... + + def get_public_paths(self) -> tuple[tuple[str, ...], tuple[str, ...]]: + """Return (prefix_paths, exact_paths) that skip authentication. + + Framework-level paths (/health, /static/, /api/docs, /openapi.json, /i18n/) + are always public (handled by the middleware itself). Providers return only + their own paths (e.g., /users/login or /keycloak/login). + """ + ... + + def is_bearer_request(self, request: Request) -> bool: + """True if the request carries a bearer token (mobile/API client).""" + ... +``` + +Both `users` and `keycloak` modules implement this protocol. The active provider is stored on `app.state.auth_provider` during module registration. + +### 2. Provider-Agnostic AuthMiddleware (moved to `auth/`) + +The current `users/middleware.py` hardcodes session-key reading and DB user loading. The new middleware delegates to the provider: + +```python +_FRAMEWORK_PUBLIC_PREFIXES = ("/health", "/static/", "/api/docs", "/api/redoc", "/openapi.json", "/i18n/") +_FRAMEWORK_PUBLIC_EXACT = ("/",) + +class AuthMiddleware: + def __init__(self, app: ASGIApp) -> None: + self.app = app + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] != "http": + await self.app(scope, receive, send) + return + + path = scope["path"] + provider: AuthProvider = scope["app"].state.auth_provider + + # Framework-level public paths + is_public = ( + any(path.startswith(p) for p in _FRAMEWORK_PUBLIC_PREFIXES) + or path in _FRAMEWORK_PUBLIC_EXACT + ) + # Provider-specific public paths + if not is_public: + prefix_paths, exact_paths = provider.get_public_paths() + is_public = any(path.startswith(p) for p in prefix_paths) or path in exact_paths + + request = Request(scope) + user_ctx = await provider.resolve_user(request) + + # Fall through to principal-resolver chain (from #163 spec) + if user_ctx is None: + resolvers = getattr(scope["app"].state.auth, "principal_resolvers", ()) + for resolver in resolvers: + try: + user_ctx = await resolver(request) + except Exception: + logger.exception("Principal resolver %r raised; treating as no-match", resolver) + continue + if user_ctx is not None: + break + + if user_ctx is None and not is_public: + if path.startswith("/api/"): + response = JSONResponse({"detail": "Not authenticated"}, status_code=401) + else: + session = scope["session"] + session["next"] = str(request.url) + response = RedirectResponse(provider.get_login_url(request), status_code=302) + await response(scope, receive, send) + return + + if user_ctx is not None: + request.state.user = user_ctx + token = current_user_id.set(user_ctx.id) + try: + await self.app(scope, receive, send) + finally: + current_user_id.reset(token) + return + + await self.app(scope, receive, send) +``` + +Key behaviors: +- Bearer requests get `401 JSON` instead of `302 redirect` (from the principal-resolver spec). +- The resolver chain from #163 runs after the provider's own `resolve_user`, so downstream PAT modules work with either provider. +- Framework-level public paths are hardcoded; provider-specific paths come from `get_public_paths()`. + +### 3. Users Module Changes + +The `users` module adapts to implement `AuthProvider` and gains bearer token support for mobile/API clients. + +**`AuthProvider` implementation.** A new `UsersAuthProvider` class in `users/provider.py`: + +```python +class UsersAuthProvider: + name = "users" + + async def resolve_user(self, request: Request) -> UserContext | None: + # Path 1: Bearer token (mobile/API) + auth_header = request.headers.get("authorization", "") + if auth_header.startswith("Bearer "): + token = auth_header[7:] + return await self._resolve_bearer(request, token) + + # Path 2: Session cookie (browser — existing logic moved here) + session = request.scope.get("session", {}) + raw_user_id = session.get("user_id") + if not raw_user_id: + return None + # Fast path: cached UserContext in session + user_ctx = UserContext.from_session_dict(session.get("user_ctx")) + if user_ctx and user_ctx.id == str(raw_user_id): + return user_ctx + # Slow path: DB lookup (existing _load_user logic) + return await self._load_user_from_db(request, raw_user_id) + + def get_login_url(self, request, next_url=None) -> str: + return "/users/login" + + def get_logout_url(self, request) -> str: + return "/users/logout" + + def get_public_paths(self) -> tuple[tuple[str, ...], tuple[str, ...]]: + return ( + ("/users/login", "/users/register", "/users/forgot-password", + "/users/reset-password", "/users/verify", "/users/invite/accept", + "/api/users/auth/", "/api/users/register"), + (), + ) + + def is_bearer_request(self, request) -> bool: + return request.headers.get("authorization", "").startswith("Bearer ") +``` + +**Registration.** `UsersModule.register_settings()` adds: +```python +app.state.auth_provider = UsersAuthProvider() +``` + +**New: Refresh token support for mobile.** Mobile clients need token refresh without re-authenticating: + +- **`RefreshToken` model** — new table `users_refresh_token`: `token` (UUID, PK), `user_id` (FK), `created_at`, `expires_at` (30 days), `revoked_at`. +- **`POST /api/users/auth/token`** — email+password login returning `{access_token, token_type, expires_in, refresh_token}` as JSON (no cookie set). +- **`POST /api/users/auth/token/refresh`** — exchange refresh token for new access + refresh token pair. Old refresh token is revoked (rotation). +- **`DELETE /api/users/auth/token`** — revoke refresh token (mobile logout). +- Access tokens for bearer transport use a shorter lifetime (15 minutes) than cookie tokens (14 days). This is configurable via `UsersSettings.bearer_token_lifetime_seconds`. + +**OAuth for mobile.** When `/api/users/oauth/{provider}/callback` receives a request with `Accept: application/json` (or `?response_type=token` query param), it returns tokens as JSON instead of setting cookies and redirecting. + +### 4. Keycloak Module (new package) + +**Package:** `modules/keycloak/` with entry point `simple_module.keycloak = keycloak.module:KeycloakModule`. + +**Depends on:** `Auth` (same dependency as `users`). + +**Settings** (`SM_KEYCLOAK_*`, DB-backed after bootstrap): + +| Setting | Default | Description | +|---------|---------|-------------| +| `server_url` | (required) | Keycloak base URL (e.g. `https://auth.example.com`) | +| `realm` | (required) | Realm name | +| `client_id` | (required) | OIDC client ID (confidential client) | +| `client_secret` | (required) | OIDC client secret | +| `roles_claim_path` | `realm_access.roles` | Dot-path to extract roles from ID/access token | +| `admin_role` | `admin` | Keycloak role name that maps to framework admin | +| `login_redirect_url` | `/dashboard/` | Where to redirect after successful login | +| `jwks_cache_ttl_seconds` | `3600` | How long to cache Keycloak's public keys | +| `role_mapping` | `{"admin": "admin", "user": "user"}` | Keycloak role → framework role mapping | + +**Module layout:** + +``` +modules/keycloak/keycloak/ +├── module.py # KeycloakModule(ModuleBase) — registers as auth_provider +├── models.py # KeycloakUserCache table +├── settings.py # KeycloakSettings +├── provider.py # KeycloakAuthProvider(AuthProvider) implementation +├── oidc.py # OIDC discovery, token exchange helpers +├── jwks.py # JWKS key cache + JWT signature validation +├── endpoints/ +│ ├── api.py # /api/keycloak/auth/login, /callback, /userinfo +│ └── views.py # /keycloak/login (redirect), /keycloak/logout +├── pages/ +│ ├── Login.tsx # Minimal page that auto-redirects to Keycloak +│ └── LoggedOut.tsx # Post-logout landing page +└── locales/ + └── en.json +``` + +**Web login flow:** + +1. User hits any protected page → `AuthMiddleware` redirects to `/keycloak/login`. +2. `/keycloak/login` view generates OIDC state + nonce, stores in session, redirects to Keycloak's authorization endpoint: `{server_url}/realms/{realm}/protocol/openid-connect/auth?client_id=...&redirect_uri=/api/keycloak/auth/callback&response_type=code&scope=openid+email+profile&state=...&nonce=...`. +3. User authenticates at Keycloak's hosted login page. +4. Keycloak redirects to `/api/keycloak/auth/callback?code=...&state=...`. +5. Callback validates state against session, exchanges code for tokens via Keycloak's token endpoint. +6. Validates ID token signature (JWKS), checks `iss`, `aud`, `exp`, `nonce`. +7. Extracts claims: `sub`, `email`, `preferred_username` or `name`, `realm_access.roles`. +8. Upserts `KeycloakUserCache` row (maps `keycloak_sub` → stable framework UUID). +9. Maps Keycloak roles to framework roles via `role_mapping` setting. +10. Builds `UserContext`, stores in `session["user_ctx"]`. Stores `id_token` in session for logout. +11. Redirects to `session["next"]` or `login_redirect_url`. + +**Mobile login flow (Authorization Code + PKCE):** + +Mobile clients authenticate directly with Keycloak — the framework is not in the login path: + +1. Mobile app opens system browser to Keycloak's auth endpoint with `code_challenge` + `code_challenge_method=S256` (PKCE). +2. User authenticates at Keycloak. +3. Keycloak redirects to mobile app's registered redirect URI (deep link / custom scheme) with auth code. +4. Mobile app exchanges code + `code_verifier` for tokens directly with Keycloak's token endpoint. +5. Mobile app sends requests to the framework API with `Authorization: Bearer `. +6. Framework's `KeycloakAuthProvider.resolve_user()` validates the JWT against Keycloak's JWKS keys. +7. On token expiry, mobile app refreshes directly with Keycloak (`grant_type=refresh_token`). The framework is not involved in token refresh. + +**`KeycloakAuthProvider` implementation:** + +```python +class KeycloakAuthProvider: + name = "keycloak" + + def __init__(self, settings: KeycloakSettings, jwks_cache: JWKSCache): + self.settings = settings + self.jwks_cache = jwks_cache + + async def resolve_user(self, request: Request) -> UserContext | None: + # Path 1: Bearer token (mobile/API) + auth_header = request.headers.get("authorization", "") + if auth_header.startswith("Bearer "): + token = auth_header[7:] + claims = await self.jwks_cache.validate_jwt(token) + if claims is None: + return None + return self._claims_to_user_context(claims) + + # Path 2: Session (browser) + session = request.scope.get("session", {}) + return UserContext.from_session_dict(session.get("user_ctx")) + + def get_login_url(self, request, next_url=None) -> str: + return "/keycloak/login" + + def get_logout_url(self, request) -> str: + return "/keycloak/logout" + + def get_public_paths(self): + return ( + ("/keycloak/login", "/keycloak/logout", "/api/keycloak/auth/"), + (), + ) + + def is_bearer_request(self, request) -> bool: + return request.headers.get("authorization", "").startswith("Bearer ") + + def _claims_to_user_context(self, claims: dict) -> UserContext: + roles_raw = _extract_nested(claims, self.settings.roles_claim_path) + mapped_roles = [ + self.settings.role_mapping[r] + for r in (roles_raw or []) + if r in self.settings.role_mapping + ] + return UserContext( + id=str(self._get_or_create_cache_id(claims["sub"])), + email=claims.get("email", ""), + name=claims.get("preferred_username") or claims.get("name", ""), + roles=mapped_roles, + tenant_id=claims.get("tenant_id"), + ) +``` + +**JWT validation (`jwks.py`):** + +- Fetches Keycloak's JWKS endpoint (`{server_url}/realms/{realm}/protocol/openid-connect/certs`) on first request. +- Caches keys in memory with configurable TTL (default 1 hour). +- On validation failure with cached keys: refetch JWKS once before rejecting (handles key rotation). +- Validates: signature (RS256), `iss` (must match `{server_url}/realms/{realm}`), `aud` (must contain `client_id`), `exp` (not expired), `iat` (not in future). + +**`KeycloakUserCache` model:** + +```python +class KeycloakUserCache(Base, table=True): + __tablename__ = "keycloak_user_cache" + + id: uuid.UUID = Field(default_factory=uuid4, primary_key=True) + keycloak_sub: str = Field(unique=True, index=True) + email: str + full_name: str | None = None + last_login_at: datetime | None = None +``` + +Purpose: +- Provides a stable UUID for audit trails and foreign keys from other modules (they reference `keycloak_user_cache.id`, not Keycloak's `sub` string). +- Caches user metadata so the Inertia shared props and menu system have a name/email without calling Keycloak's userinfo endpoint. +- Upserted on each web login callback and on first bearer-token request from a new user. + +**Logout (RP-initiated):** + +`POST /keycloak/logout` clears the framework session, then redirects to Keycloak's logout endpoint: +`{server_url}/realms/{realm}/protocol/openid-connect/logout?post_logout_redirect_uri={base_url}/keycloak/login&id_token_hint={id_token_from_session}` + +### 5. Conflict Detection + +**SM020 — Multiple auth providers.** New diagnostic in `ModuleDiagnostics`: + +```python +def _check_auth_provider_conflict(self, modules: list[ModuleBase]) -> list[Diagnostic]: + providers = [m for m in modules if getattr(m, '_is_auth_provider', False)] + if len(providers) > 1: + names = ", ".join(m.meta.name for m in providers) + return [Diagnostic( + level=DiagnosticLevel.ERROR, + code="SM020", + message=f"Multiple auth provider modules installed: {names}", + suggestion="Install only one auth provider (e.g. 'users' OR 'keycloak', not both)", + )] + return [] +``` + +**SM021 — No auth provider.** Warns (not errors) if no auth provider is installed — allows headless/API-only deployments that handle auth externally. + +```python + if len(providers) == 0: + return [Diagnostic( + level=DiagnosticLevel.WARNING, + code="SM021", + message="No auth provider module installed", + suggestion="Install an auth provider module (e.g. 'simple-module-users' or 'simple-module-keycloak')", + )] +``` + +The marker `_is_auth_provider = True` is a class attribute on both `UsersModule` and `KeycloakModule`. In production (`strict=True`), SM020 (error) fails boot; SM021 (warning) logs only. + +### 6. Shared Infrastructure Moved to `auth/` + +These pieces currently live in `users/` but are provider-agnostic: + +| What | From | To | Reason | +|------|------|----|--------| +| `AuthMiddleware` | `users/middleware.py` | `auth/middleware.py` | Provider-agnostic; delegates to `AuthProvider` | +| `principal_serializer` | `users/module.py` | `auth/module.py` (`AuthModule.register_settings`) | Serializes `UserContext` for Inertia shared props; provider-agnostic since it only reads `UserContext` fields. Registered on `app.state.principal_serializer` where `InertiaLayoutDataMiddleware` reads it. | +| Framework public paths | `users/middleware.py` | `auth/middleware.py` | `/health`, `/static/`, `/api/docs` are framework-level | + +Provider-specific public paths (e.g., `/users/login`, `/keycloak/login`) are returned by each provider's `get_public_paths()`. + +`auth/deps.py` (`get_current_user`, `CurrentUser`, `require_permission`) stays unchanged — it reads from `request.state.user` which the middleware sets. + +`AuthModule.register_middleware()` now registers the `AuthMiddleware` instead of `UsersModule.register_middleware()`. Since `Users` depends on `Auth`, the middleware ordering is preserved (Auth middleware wraps Users routes). + +## Menu Integration + +Both providers register appropriate menu items: + +**Users module (unchanged):** Users admin, Profile, Logout in the user dropdown. + +**Keycloak module:** +- Logout menu item (`POST /keycloak/logout`, user dropdown, order 999). +- Optional: link to Keycloak account console (`{server_url}/realms/{realm}/account`, user dropdown, order 980) — configurable, off by default. +- No "Users admin" menu — user management happens in Keycloak's admin console. +- Role mapping admin page (sidebar, admin-only, order 100) — configure which Keycloak roles map to which framework permission groups. + +## Role Mapping + +Both providers feed roles into the same `PermissionRegistry`: + +**Users module (unchanged):** `User.roles` → `[r.name for r in user.roles]` → `UserContext.roles`. + +**Keycloak module:** JWT `realm_access.roles` → filtered through `role_mapping` dict → `UserContext.roles`. + +Default mapping (DB-backed, editable via admin UI): + +```python +role_mapping: dict[str, str] = { + "admin": "admin", + "user": "user", +} +``` + +Unknown Keycloak roles are silently ignored. The admin UI page lets operators add custom mappings (e.g., Keycloak `editor` → framework `editor` if the app defines that role). + +## Module Selection + +Framework users choose at install time in their app's `pyproject.toml`: + +```toml +# Pick one auth provider: +dependencies = [ + "simple-module-auth", # always required (contracts layer) + "simple-module-users", # local auth — OR: + # "simple-module-keycloak", # Keycloak auth +] +``` + +Both declare the same entry-point group (`simple_module`), both implement `AuthProvider`, and the SM020 diagnostic ensures only one is active. + +## Migration Path + +For framework users switching from `users` to `keycloak`: + +1. Provision Keycloak realm + client (operator responsibility, out of scope). +2. Create matching users in Keycloak (manual or bulk import, out of scope). +3. Replace `simple-module-users` with `simple-module-keycloak` in `pyproject.toml`. +4. Set `SM_KEYCLOAK_*` env vars (or configure via settings UI after first boot with a superuser session cookie). +5. Run `uv sync --all-packages && make migrate` to apply keycloak module's migration. +6. Optionally clean up users module tables: `alembic downgrade users@base`. + +No automated data migration — user identity lives in Keycloak. + +## Testing Strategy + +**Unit tests (keycloak module):** +- `KeycloakAuthProvider.resolve_user()` with mocked JWKS validation. +- JWT validation: valid token, expired, wrong issuer, wrong audience, malformed, key rotation (cache refresh). +- Role mapping: known roles, unknown roles, empty roles, admin bypass. +- OIDC flow: state generation, state validation, code exchange (mocked HTTP). +- `KeycloakUserCache` upsert logic. + +**Unit tests (auth middleware):** +- `AuthMiddleware` with a mock `AuthProvider` — verify redirect vs. 401 behavior. +- Bearer token request through the full middleware stack. +- Public path matching (framework-level + provider-specific). +- Principal-resolver chain fallthrough after provider returns None. + +**Unit tests (users module changes):** +- `UsersAuthProvider.resolve_user()` — bearer path and session path. +- Refresh token creation, rotation, revocation. +- `POST /api/users/auth/token` returns JSON tokens. +- `POST /api/users/auth/token/refresh` rotates tokens. +- OAuth callback with `Accept: application/json` returns JSON. + +**Integration tests:** +- `SM020` / `SM021` diagnostics with various module combinations. +- Full request flow: login → authenticated request → logout, with each provider. +- Bearer token flow: obtain token → API request → token expiry → refresh (users provider only). + +**E2E tests (require running Keycloak via Docker):** +- Web login redirect flow with a test Keycloak instance. +- Post-login redirect to `session["next"]`. +- Logout clears both framework session and Keycloak session. +- Role mapping reflected in menu visibility and permission checks. + +## Dependencies + +**New Python packages (keycloak module only):** +- `PyJWT` — JWT decoding and validation. Already a transitive dep via fastapi-users; used directly for Keycloak JWT validation with RS256. +- `cryptography` — RSA key handling for JWKS. Already a transitive dep. +- `httpx` — async HTTP for OIDC discovery, token exchange, JWKS fetching. Already a project dep. + +No new frontend dependencies. No new deps for the `auth` or `users` module changes. + +## Risks & Mitigations + +| Risk | Mitigation | +|------|-----------| +| JWKS cache staleness (Keycloak rotates signing keys) | Cache TTL + retry: on validation failure with cached keys, refetch JWKS once before rejecting | +| Keycloak downtime blocks all logins | Health check: keycloak module registers a check that probes the OIDC discovery endpoint; alerts fire before users notice | +| `UserContext.id` format mismatch (UUID vs. Keycloak sub) | `KeycloakUserCache` provides a stable UUID; `UserContext.id` always uses the cache table's UUID | +| Other modules FK to `users.User` table | Document that modules should reference `UserContext.id` (UUID string) for audit trails; existing modules that import `users.models.User` directly won't work with keycloak (by design — SM009 already prevents framework→module imports) | +| Session size with Keycloak tokens | Only `UserContext` dict + `id_token` (for logout hint) stored in session; access/refresh tokens are not stored server-side | +| Breaking change: `AuthMiddleware` moves from `users` to `auth` | The middleware was internal to the users module; no other module imports it directly. The observable behavior (request.state.user set, redirects, 401s) is identical. | diff --git a/framework/cli/tests/test_cli_package_update.py b/framework/cli/tests/test_cli_package_update.py index 019b552f..ae944542 100644 --- a/framework/cli/tests/test_cli_package_update.py +++ b/framework/cli/tests/test_cli_package_update.py @@ -157,14 +157,23 @@ def fetcher(url: str) -> dict: def test_missing_pyproject_exits_nonzero(tmp_path: Path) -> None: - with pytest.raises(click.exceptions.Exit) as exc: + # typer >= 0.26 vendors click as ``typer._click``; the raised ``Exit`` + # no longer inherits from ``click.exceptions.Exit``. Catch both. + _exit_types: tuple[type[BaseException], ...] = (click.exceptions.Exit,) + try: + from typer._click.exceptions import Exit as _TyExit + + _exit_types = (*_exit_types, _TyExit) + except ImportError: + pass + with pytest.raises(_exit_types) as exc: pu.run_update( path=tmp_path, dry_run=False, include_pre=False, fetcher=_fake_pypi({}), ) - assert exc.value.exit_code == 1 + assert getattr(exc.value, "exit_code", getattr(exc.value, "code", None)) == 1 def test_cli_command_registered() -> None: diff --git a/framework/core/simple_module_core/diagnostics/_module.py b/framework/core/simple_module_core/diagnostics/_module.py index de3db839..b58e1a68 100644 --- a/framework/core/simple_module_core/diagnostics/_module.py +++ b/framework/core/simple_module_core/diagnostics/_module.py @@ -26,6 +26,7 @@ def run(self, modules: list[ModuleBase]) -> list[Diagnostic]: diagnostics.extend(self._check_empty_modules(modules)) diagnostics.extend(self._check_missing_meta(modules)) diagnostics.extend(self._check_views_without_menu(modules)) + diagnostics.extend(self._check_auth_provider_conflict(modules)) diagnostics.extend(check_framework_module_coupling(modules)) # File-based checks (need to find module source directories) @@ -162,6 +163,38 @@ def _check_views_without_menu(self, modules: list[ModuleBase]) -> list[Diagnosti ) return diags + def _check_auth_provider_conflict(self, modules: list[ModuleBase]) -> list[Diagnostic]: + """SM020/SM021: exactly one auth provider module must be installed.""" + providers = [m for m in modules if getattr(m, "_is_auth_provider", False)] + diags: list[Diagnostic] = [] + if len(providers) > 1: + names = ", ".join(m.meta.name for m in providers) + diags.append( + Diagnostic( + level=DiagnosticLevel.ERROR, + code="SM020", + message=f"Multiple auth provider modules installed: {names}", + module_name=providers[0].meta.name, + suggestion=( + "Install only one auth provider (e.g. 'users' OR 'keycloak', not both)" + ), + ) + ) + elif len(providers) == 0: + diags.append( + Diagnostic( + level=DiagnosticLevel.WARNING, + code="SM021", + message="No auth provider module installed", + module_name="(none)", + suggestion=( + "Install an auth provider module " + "(e.g. 'simple-module-users' or 'simple-module-keycloak')" + ), + ) + ) + return diags + def _check_missing_meta(self, modules: list[ModuleBase]) -> list[Diagnostic]: diags: list[Diagnostic] = [] for mod in modules: diff --git a/framework/core/tests/test_diagnostics.py b/framework/core/tests/test_diagnostics.py index 8333b04c..5a329a9d 100644 --- a/framework/core/tests/test_diagnostics.py +++ b/framework/core/tests/test_diagnostics.py @@ -72,3 +72,51 @@ async def test_empty_is_quiet(self, capsys): captured = capsys.readouterr() assert captured.out == "" assert captured.err == "" + + +class TestAuthProviderDiagnostics: + def test_sm020_multiple_auth_providers(self): + from simple_module_core.diagnostics._module import ModuleDiagnostics + from simple_module_core.module import ModuleBase, ModuleMeta + + class FakeUsers(ModuleBase): + meta = ModuleMeta(name="Users") + _is_auth_provider = True + + class FakeKeycloak(ModuleBase): + meta = ModuleMeta(name="Keycloak") + _is_auth_provider = True + + diags = ModuleDiagnostics() + results = diags._check_auth_provider_conflict([FakeUsers(), FakeKeycloak()]) + assert len(results) == 1 + assert results[0].code == "SM020" + assert results[0].level == DiagnosticLevel.ERROR + + def test_sm021_no_auth_provider(self): + from simple_module_core.diagnostics._module import ModuleDiagnostics + from simple_module_core.module import ModuleBase, ModuleMeta + + class FakeDashboard(ModuleBase): + meta = ModuleMeta(name="Dashboard") + + diags = ModuleDiagnostics() + results = diags._check_auth_provider_conflict([FakeDashboard()]) + assert len(results) == 1 + assert results[0].code == "SM021" + assert results[0].level == DiagnosticLevel.WARNING + + def test_single_provider_passes(self): + from simple_module_core.diagnostics._module import ModuleDiagnostics + from simple_module_core.module import ModuleBase, ModuleMeta + + class FakeUsers(ModuleBase): + meta = ModuleMeta(name="Users") + _is_auth_provider = True + + class FakeDashboard(ModuleBase): + meta = ModuleMeta(name="Dashboard") + + diags = ModuleDiagnostics() + results = diags._check_auth_provider_conflict([FakeUsers(), FakeDashboard()]) + assert results == [] diff --git a/framework/hosting/tests/test_app.py b/framework/hosting/tests/test_app.py index f6c6a50f..e434f510 100644 --- a/framework/hosting/tests/test_app.py +++ b/framework/hosting/tests/test_app.py @@ -77,7 +77,13 @@ async def test_app_state_has_sm_services( (tmp_path / "host" / "templates").mkdir(parents=True) (tmp_path / "host" / "templates" / "index.html").write_text("") - app = create_app() + # Exclude Keycloak — both ``users`` and ``keycloak`` are installed as + # entry points in the dev workspace. SM020 fires if both are present + # because the app is only meant to run with one auth provider. + from simple_module_core.discovery import discover_modules + + all_names = [m.meta.name for m in discover_modules() if m.meta.name != "Keycloak"] + app = create_app(Settings(modules_enabled=all_names)) sm = app.state.sm assert isinstance(sm, Services) diff --git a/modules/auth/auth/__init__.py b/modules/auth/auth/__init__.py index 1188109b..69f6671c 100644 --- a/modules/auth/auth/__init__.py +++ b/modules/auth/auth/__init__.py @@ -1,6 +1,7 @@ -"""Auth module — shared contracts (UserContext, PrincipalResolver, deps).""" +"""Auth module — shared contracts (UserContext, AuthProvider, PrincipalResolver, deps).""" +from auth.contracts.provider import AuthProvider from auth.contracts.resolver import PrincipalResolver from auth.contracts.schemas import UserContext -__all__ = ["PrincipalResolver", "UserContext"] +__all__ = ["AuthProvider", "PrincipalResolver", "UserContext"] diff --git a/modules/auth/auth/contracts/__init__.py b/modules/auth/auth/contracts/__init__.py index c9439776..afe7dd4d 100644 --- a/modules/auth/auth/contracts/__init__.py +++ b/modules/auth/auth/contracts/__init__.py @@ -1,5 +1,6 @@ """Auth contracts — public types for other modules.""" +from auth.contracts.provider import AuthProvider from auth.contracts.schemas import UserContext -__all__ = ["UserContext"] +__all__ = ["AuthProvider", "UserContext"] diff --git a/modules/auth/auth/contracts/provider.py b/modules/auth/auth/contracts/provider.py new file mode 100644 index 00000000..4714a145 --- /dev/null +++ b/modules/auth/auth/contracts/provider.py @@ -0,0 +1,34 @@ +"""AuthProvider protocol — the contract both users and keycloak modules implement.""" + +from __future__ import annotations + +from typing import Protocol, runtime_checkable + +from starlette.requests import Request + +from auth.contracts.schemas import UserContext + + +@runtime_checkable +class AuthProvider(Protocol): + """Extension point for swappable authentication backends. + + Exactly one module (``users`` or ``keycloak``) registers an implementation + on ``app.state.auth.auth_provider`` during ``register_settings``. + The ``AuthMiddleware`` delegates to it on every request. + """ + + name: str + + async def resolve_user(self, request: Request) -> UserContext | None: ... + + def get_login_url(self, request: Request, next_url: str | None = None) -> str: ... + + def get_logout_url(self, request: Request) -> str: ... + + def get_public_paths(self) -> tuple[tuple[str, ...], tuple[str, ...]]: ... + + def is_bearer_request(self, request: Request) -> bool: ... + + +__all__ = ["AuthProvider"] diff --git a/modules/auth/auth/middleware.py b/modules/auth/auth/middleware.py new file mode 100644 index 00000000..995a1799 --- /dev/null +++ b/modules/auth/auth/middleware.py @@ -0,0 +1,103 @@ +"""Provider-agnostic authentication middleware. + +Delegates user resolution to the ``AuthProvider`` registered on +``app.state.auth.auth_provider``, then falls through to the +principal-resolver chain. Sets ``request.state.user`` and the +``current_user_id`` ContextVar for audit listeners. +""" + +from __future__ import annotations + +import logging + +from simple_module_db.listeners import current_user_id +from starlette.requests import Request +from starlette.responses import JSONResponse, RedirectResponse +from starlette.types import ASGIApp, Receive, Scope, Send + +logger = logging.getLogger(__name__) + +_FRAMEWORK_PUBLIC_PREFIXES = ( + "/health", + "/static/", + "/api/docs", + "/api/redoc", + "/openapi.json", + "/i18n/", +) +_FRAMEWORK_PUBLIC_EXACT = ("/",) +_SESSION_NEXT_KEY = "next" + + +class AuthMiddleware: + """Authenticate requests via the registered AuthProvider. + + On cache miss (provider returns None), falls through to the + principal-resolver chain. Unauthenticated API requests get 401 JSON; + unauthenticated browser requests get a redirect to the provider's + login URL. + """ + + def __init__(self, app: ASGIApp) -> None: + self.app = app + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] != "http": + await self.app(scope, receive, send) + return + + path: str = scope["path"] + auth_state = scope["app"].state.auth + provider = auth_state.auth_provider + + if provider is None: + await self.app(scope, receive, send) + return + + is_public = ( + any(path.startswith(p) for p in _FRAMEWORK_PUBLIC_PREFIXES) + or path in _FRAMEWORK_PUBLIC_EXACT + ) + if not is_public: + prefix_paths, exact_paths = provider.get_public_paths() + is_public = any(path.startswith(p) for p in prefix_paths) or path in exact_paths + + request = Request(scope) + user_ctx = await provider.resolve_user(request) + + if user_ctx is None: + for resolver in auth_state.principal_resolvers: + try: + user_ctx = await resolver(request) + except Exception: + logger.exception( + "Principal resolver %r raised; treating as no-match", + resolver, + ) + continue + if user_ctx is not None: + break + + if user_ctx is None and not is_public: + if path.startswith("/api/") or provider.is_bearer_request(request): + response = JSONResponse({"detail": "Not authenticated"}, status_code=401) + else: + session = scope.get("session", {}) + session[_SESSION_NEXT_KEY] = str(request.url) + response = RedirectResponse(provider.get_login_url(request), status_code=302) + await response(scope, receive, send) + return + + if user_ctx is not None: + request.state.user = user_ctx + token = current_user_id.set(user_ctx.id) + try: + await self.app(scope, receive, send) + finally: + current_user_id.reset(token) + return + + await self.app(scope, receive, send) + + +__all__ = ["AuthMiddleware"] diff --git a/modules/auth/auth/module.py b/modules/auth/auth/module.py index 6a75bd4b..9f047a89 100644 --- a/modules/auth/auth/module.py +++ b/modules/auth/auth/module.py @@ -1,15 +1,14 @@ -"""Auth module — shared contracts (UserContext, deps). +"""Auth module — shared contracts (UserContext, AuthProvider, deps). Intentionally minimal: this module owns the PUBLIC interface (UserContext, -PrincipalResolver, get_current_user, CurrentUser, require_permission) that -every other module imports. Keeping it stable prevents churn when auth +AuthProvider, PrincipalResolver, get_current_user, CurrentUser, require_permission) +that every other module imports. Keeping it stable prevents churn when auth internals change. -All authentication logic (middleware, login, signup, OAuth) lives in the -users module. The ``principal_resolvers`` registry on ``app.state.auth`` is -the extension point downstream modules use to plug in additional credential -sources (PAT bearer tokens, API keys, etc.) — see -``docs/framework/principal-resolvers.md`` for the worked example. +The ``auth_provider`` slot on ``app.state.auth`` is the extension point +auth-provider modules (``users``, ``keycloak``) use to register themselves. +The ``principal_resolvers`` registry lets downstream modules add extra +credential sources (PAT bearer tokens, API keys, etc.). """ from __future__ import annotations @@ -23,6 +22,17 @@ if TYPE_CHECKING: from fastapi import FastAPI + from auth.contracts.schemas import UserContext + + +def _serialize_principal(user: UserContext) -> dict: + return { + "id": user.id, + "name": user.name, + "email": user.email, + "roles": user.roles, + } + class AuthModule(ModuleBase): meta = ModuleMeta( @@ -34,6 +44,12 @@ def register_settings(self, app: FastAPI) -> None: from auth.state import AuthState app.state.auth = AuthState() + app.state.principal_serializer = _serialize_principal + + def register_middleware(self, app: FastAPI) -> None: + from auth.middleware import AuthMiddleware + + app.add_middleware(AuthMiddleware) def locale_dirs(self) -> dict[str, Path]: return {"auth": Path(str(importlib.resources.files(__package__) / "locales"))} diff --git a/modules/auth/auth/state.py b/modules/auth/auth/state.py index cc457da7..9e44eb2b 100644 --- a/modules/auth/auth/state.py +++ b/modules/auth/auth/state.py @@ -1,8 +1,8 @@ """Module-owned state attached to ``app.state.auth`` by ``AuthModule.register_settings``. -Holds the principal-resolver registry (see -``auth.contracts.resolver.PrincipalResolver``). Apps register additional -resolvers from their ``on_startup`` hook:: +Holds the auth provider (set by one of ``users`` or ``keycloak``) and the +principal-resolver registry. Apps register additional resolvers from their +``on_startup`` hook:: app.state.auth.principal_resolvers.append(my_pat_resolver) """ @@ -10,14 +10,19 @@ from __future__ import annotations from dataclasses import dataclass, field +from typing import TYPE_CHECKING from auth.contracts.resolver import PrincipalResolver +if TYPE_CHECKING: + from auth.contracts.provider import AuthProvider + @dataclass class AuthState: - """Per-app auth registry. Initialized empty; modules append resolvers.""" + """Per-app auth registry. Initialized empty; provider modules populate at boot.""" + auth_provider: AuthProvider | None = None principal_resolvers: list[PrincipalResolver] = field(default_factory=list) diff --git a/modules/auth/tests/test_auth_middleware.py b/modules/auth/tests/test_auth_middleware.py new file mode 100644 index 00000000..e990b3a8 --- /dev/null +++ b/modules/auth/tests/test_auth_middleware.py @@ -0,0 +1,161 @@ +"""Tests for the provider-agnostic AuthMiddleware.""" + +from __future__ import annotations + +import httpx +import pytest +from auth.contracts.schemas import UserContext +from auth.middleware import AuthMiddleware +from auth.state import AuthState +from fastapi import FastAPI, Request +from starlette.middleware.sessions import SessionMiddleware +from starlette.responses import JSONResponse + +SECRET = "test-middleware-secret" + +_TEST_USER = UserContext( + id="aaaaaaaa-bbbb-cccc-dddd-eeeeeeeeeeee", + email="test@example.com", + name="Test User", + roles=["admin"], +) + + +class _StubProvider: + name = "stub" + + def __init__(self, *, user: UserContext | None = None): + self._user = user + + async def resolve_user(self, request): + return self._user + + def get_login_url(self, request, next_url=None): + return "/stub/login" + + def get_logout_url(self, request): + return "/stub/logout" + + def get_public_paths(self): + return (("/stub/login", "/stub/public/"), ()) + + def is_bearer_request(self, request): + auth = request.headers.get("authorization", "") + return auth.startswith("Bearer ") + + +def _build_app(provider, *, principal_resolvers=None): + app = FastAPI() + app.state.auth = AuthState( + auth_provider=provider, + principal_resolvers=list(principal_resolvers or []), + ) + + @app.get("/{path:path}") + async def catch_all(request: Request, path: str = ""): + user = getattr(request.state, "user", None) + return JSONResponse( + { + "user": user.to_session_dict() if user else None, + } + ) + + app.add_middleware(AuthMiddleware) + app.add_middleware(SessionMiddleware, secret_key=SECRET) + return app + + +@pytest.fixture +def authenticated_app(): + return _build_app(_StubProvider(user=_TEST_USER)) + + +@pytest.fixture +def unauthenticated_app(): + return _build_app(_StubProvider(user=None)) + + +async def test_authenticated_request_sets_user(authenticated_app): + transport = httpx.ASGITransport(app=authenticated_app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as c: + resp = await c.get("/some/page") + assert resp.status_code == 200 + assert resp.json()["user"]["email"] == "test@example.com" + + +async def test_unauthenticated_browser_redirects_to_login(unauthenticated_app): + transport = httpx.ASGITransport(app=unauthenticated_app) + async with httpx.AsyncClient( + transport=transport, base_url="http://test", follow_redirects=False + ) as c: + resp = await c.get("/protected/page") + assert resp.status_code == 302 + assert resp.headers["location"] == "/stub/login" + + +async def test_unauthenticated_api_returns_401(unauthenticated_app): + transport = httpx.ASGITransport(app=unauthenticated_app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as c: + resp = await c.get("/api/protected") + assert resp.status_code == 401 + assert resp.json()["detail"] == "Not authenticated" + + +async def test_unauthenticated_bearer_returns_401(unauthenticated_app): + transport = httpx.ASGITransport(app=unauthenticated_app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as c: + resp = await c.get("/some/page", headers={"Authorization": "Bearer bad"}) + assert resp.status_code == 401 + + +async def test_public_paths_skip_auth(unauthenticated_app): + transport = httpx.ASGITransport(app=unauthenticated_app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as c: + resp = await c.get("/stub/login") + assert resp.status_code == 200 + + +async def test_framework_public_paths_skip_auth(unauthenticated_app): + transport = httpx.ASGITransport(app=unauthenticated_app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as c: + resp = await c.get("/health") + assert resp.status_code == 200 + + +async def test_root_is_public(unauthenticated_app): + transport = httpx.ASGITransport(app=unauthenticated_app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as c: + resp = await c.get("/") + assert resp.status_code == 200 + + +async def test_resolver_chain_fallback(): + """When provider returns None, fall through to principal resolvers.""" + + async def fake_resolver(request): + auth = request.headers.get("authorization", "") + if auth == "Bearer good-token": + return _TEST_USER + return None + + app = _build_app(_StubProvider(user=None), principal_resolvers=[fake_resolver]) + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as c: + resp = await c.get("/protected", headers={"Authorization": "Bearer good-token"}) + assert resp.status_code == 200 + assert resp.json()["user"]["email"] == "test@example.com" + + +async def test_resolver_exception_is_logged_and_skipped(): + """A resolver that raises should be caught; middleware continues.""" + + async def bad_resolver(request): + raise RuntimeError("boom") + + app = _build_app(_StubProvider(user=None), principal_resolvers=[bad_resolver]) + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient( + transport=transport, base_url="http://test", follow_redirects=False + ) as c: + resp = await c.get("/protected/page") + assert resp.status_code == 302 diff --git a/modules/auth/tests/test_auth_provider_protocol.py b/modules/auth/tests/test_auth_provider_protocol.py new file mode 100644 index 00000000..5fde0863 --- /dev/null +++ b/modules/auth/tests/test_auth_provider_protocol.py @@ -0,0 +1,56 @@ +"""Tests for the AuthProvider protocol.""" + +from __future__ import annotations + +from auth.contracts.provider import AuthProvider +from auth.contracts.schemas import UserContext +from starlette.requests import Request + + +class _FakeProvider: + """Minimal implementation to verify protocol conformance.""" + + name = "fake" + + async def resolve_user(self, request: Request) -> UserContext | None: + return None + + def get_login_url(self, request: Request, next_url: str | None = None) -> str: + return "/fake/login" + + def get_logout_url(self, request: Request) -> str: + return "/fake/logout" + + def get_public_paths(self) -> tuple[tuple[str, ...], tuple[str, ...]]: + return (("/fake/login",), ()) + + def is_bearer_request(self, request: Request) -> bool: + return False + + +def test_fake_provider_satisfies_protocol(): + provider = _FakeProvider() + assert isinstance(provider, AuthProvider) + + +def test_protocol_rejects_incomplete_implementation(): + class _Incomplete: + name = "broken" + + assert not isinstance(_Incomplete(), AuthProvider) + + +def test_auth_package_reexports_auth_provider(): + import auth + + assert hasattr(auth, "AuthProvider") + assert "AuthProvider" in auth.__all__ + from auth.contracts.provider import AuthProvider as Canonical + + assert auth.AuthProvider is Canonical + + +def test_contracts_package_reexports_auth_provider(): + from auth.contracts import AuthProvider + + assert AuthProvider is not None diff --git a/modules/auth/tests/test_resolver_registry.py b/modules/auth/tests/test_resolver_registry.py index 86e71af9..0d0e0e14 100644 --- a/modules/auth/tests/test_resolver_registry.py +++ b/modules/auth/tests/test_resolver_registry.py @@ -79,3 +79,66 @@ def test_auth_package_reexports_public_surface(): assert auth.PrincipalResolver is PrincipalResolver assert auth.UserContext is UserContext + + +def test_auth_module_registers_middleware(): + """AuthModule.register_middleware should add AuthMiddleware.""" + from auth.module import AuthModule + from fastapi import FastAPI + + app = FastAPI() + AuthModule().register_middleware(app) + middleware_classes = [m.cls.__name__ for m in app.user_middleware] + assert "AuthMiddleware" in middleware_classes + + +def test_auth_module_registers_principal_serializer(): + """AuthModule.register_settings should set principal_serializer on app.state.""" + from auth.module import AuthModule + from fastapi import FastAPI + + app = FastAPI() + AuthModule().register_settings(app) + serializer = getattr(app.state, "principal_serializer", None) + assert serializer is not None + + from auth.contracts.schemas import UserContext + + ctx = UserContext(id="123", email="a@b.com", name="Test", roles=["admin"]) + result = serializer(ctx) + assert result == {"id": "123", "name": "Test", "email": "a@b.com", "roles": ["admin"]} + + +def test_auth_state_has_auth_provider_field(): + from auth.state import AuthState + + state = AuthState() + assert state.auth_provider is None + + +def test_auth_state_accepts_auth_provider(): + from auth.contracts.provider import AuthProvider + from auth.state import AuthState + + class FakeProvider: + name = "fake" + + async def resolve_user(self, request): + return None + + def get_login_url(self, request, next_url=None): + return "/login" + + def get_logout_url(self, request): + return "/logout" + + def get_public_paths(self): + return ((), ()) + + def is_bearer_request(self, request): + return False + + provider = FakeProvider() + state = AuthState(auth_provider=provider) + assert state.auth_provider is provider + assert isinstance(state.auth_provider, AuthProvider) diff --git a/modules/keycloak/README.md b/modules/keycloak/README.md new file mode 100644 index 00000000..44f74d33 --- /dev/null +++ b/modules/keycloak/README.md @@ -0,0 +1,30 @@ +# simple_module_keycloak + +Keycloak OIDC authentication provider for simple_module. Swap with `simple_module_users` for Keycloak-backed identity management — users, roles, and login all handled by Keycloak. + +## Install + +Add to your app's `pyproject.toml` dependencies instead of `simple_module_users`: + +```toml +dependencies = [ + "simple_module_keycloak==0.0.15", +] +``` + +Run `uv sync --all-packages` to install. + +## Usage + +Set the required environment variables (or configure via the settings admin UI after first boot): + +```bash +SM_KEYCLOAK_SERVER_URL=https://keycloak.example.com +SM_KEYCLOAK_REALM=my-realm +SM_KEYCLOAK_CLIENT_ID=my-app +SM_KEYCLOAK_CLIENT_SECRET=my-secret +``` + +The module auto-registers as the auth provider. Browser users are redirected to Keycloak's hosted login page. Mobile clients authenticate directly with Keycloak and send bearer tokens to the framework API. + +Keycloak realm roles are mapped to framework permissions via the `role_mapping` setting (default: `admin` and `user` map 1:1). diff --git a/modules/keycloak/keycloak/__init__.py b/modules/keycloak/keycloak/__init__.py new file mode 100644 index 00000000..0e52916f --- /dev/null +++ b/modules/keycloak/keycloak/__init__.py @@ -0,0 +1 @@ +"""Keycloak OIDC authentication provider for simple_module.""" diff --git a/modules/keycloak/keycloak/contracts/__init__.py b/modules/keycloak/keycloak/contracts/__init__.py new file mode 100644 index 00000000..e8bb4a19 --- /dev/null +++ b/modules/keycloak/keycloak/contracts/__init__.py @@ -0,0 +1 @@ +"""Keycloak module contracts.""" diff --git a/modules/keycloak/keycloak/endpoints/__init__.py b/modules/keycloak/keycloak/endpoints/__init__.py new file mode 100644 index 00000000..1ead3adf --- /dev/null +++ b/modules/keycloak/keycloak/endpoints/__init__.py @@ -0,0 +1 @@ +"""Keycloak endpoint routers.""" diff --git a/modules/keycloak/keycloak/endpoints/api.py b/modules/keycloak/keycloak/endpoints/api.py new file mode 100644 index 00000000..fe6281a5 --- /dev/null +++ b/modules/keycloak/keycloak/endpoints/api.py @@ -0,0 +1,92 @@ +"""Keycloak OIDC API endpoints — login redirect, callback.""" + +from __future__ import annotations + +import logging +import secrets +from typing import TYPE_CHECKING + +from fastapi import APIRouter, HTTPException, Request +from starlette.responses import RedirectResponse + +if TYPE_CHECKING: + from keycloak.settings import KeycloakSettings + +logger = logging.getLogger(__name__) +router = APIRouter(prefix="/auth", tags=["keycloak-auth"]) + +_SESSION_OIDC_STATE = "keycloak_oidc_state" +_SESSION_OIDC_NONCE = "keycloak_oidc_nonce" +_SESSION_USER_CTX = "user_ctx" +_SESSION_ID_TOKEN = "keycloak_id_token" +_SESSION_NEXT = "next" + + +def _get_settings(request: Request) -> KeycloakSettings: + return request.app.state.keycloak.settings + + +def _get_oidc_client(request: Request): + from keycloak.oidc import OIDCClient + + s = _get_settings(request) + return OIDCClient( + server_url=s.server_url, + realm=s.realm, + client_id=s.client_id, + client_secret=s.client_secret, + ) + + +@router.get("/login") +async def oidc_login(request: Request): + client = _get_oidc_client(request) + callback_url = str(request.url_for("oidc_callback")) + nonce = secrets.token_urlsafe(32) + url, state = client.build_authorization_url( + redirect_uri=callback_url, + nonce=nonce, + ) + request.session[_SESSION_OIDC_STATE] = state + request.session[_SESSION_OIDC_NONCE] = nonce + return RedirectResponse(url, status_code=302) + + +@router.get("/callback") +async def oidc_callback(request: Request): + code = request.query_params.get("code") + state = request.query_params.get("state") + + expected_state = request.session.pop(_SESSION_OIDC_STATE, None) + request.session.pop(_SESSION_OIDC_NONCE, None) + + if not code or not state or state != expected_state: + raise HTTPException(status_code=400, detail="Invalid OIDC state") + + client = _get_oidc_client(request) + callback_url = str(request.url_for("oidc_callback")) + + try: + tokens = await client.exchange_code(code=code, redirect_uri=callback_url) + except Exception: + logger.exception("Token exchange failed") + raise HTTPException(status_code=502, detail="Token exchange failed") from None + + id_token = tokens.get("id_token", "") + access_token = tokens.get("access_token", "") + + jwks_cache = request.app.state.keycloak.jwks_cache + claims = await jwks_cache.validate_jwt(access_token) if jwks_cache else None + if claims is None: + raise HTTPException(status_code=401, detail="Token validation failed") + + provider = request.app.state.auth.auth_provider + cache_id = await provider._upsert_user_cache(request, claims) + user_ctx = provider._claims_to_user_context(claims, cache_id=cache_id) + + request.session[_SESSION_USER_CTX] = user_ctx.to_session_dict() + request.session[_SESSION_ID_TOKEN] = id_token + + s = _get_settings(request) + next_url = request.session.pop(_SESSION_NEXT, None) or s.login_redirect_url + return RedirectResponse(next_url, status_code=303) diff --git a/modules/keycloak/keycloak/endpoints/views.py b/modules/keycloak/keycloak/endpoints/views.py new file mode 100644 index 00000000..b40be653 --- /dev/null +++ b/modules/keycloak/keycloak/endpoints/views.py @@ -0,0 +1,40 @@ +"""Keycloak Inertia view routes — login page, logout.""" + +from __future__ import annotations + +from fastapi import APIRouter, Request +from simple_module_hosting.inertia_deps import InertiaDep +from starlette.responses import RedirectResponse + +router = APIRouter(tags=["keycloak-views"]) + +_SESSION_ID_TOKEN = "keycloak_id_token" +_PAGE_LOGIN = "Keycloak/Login" + + +@router.get("/login") +async def login_page(request: Request, inertia: InertiaDep): + return await inertia.render(_PAGE_LOGIN) + + +@router.post("/logout") +async def logout(request: Request): + from keycloak.oidc import OIDCClient + + s = request.app.state.keycloak.settings + id_token = request.session.get(_SESSION_ID_TOKEN) + + request.session.clear() + + client = OIDCClient( + server_url=s.server_url, + realm=s.realm, + client_id=s.client_id, + client_secret=s.client_secret, + ) + base_url = str(request.base_url).rstrip("/") + logout_url = client.build_logout_url( + post_logout_redirect_uri=f"{base_url}/keycloak/login", + id_token_hint=id_token, + ) + return RedirectResponse(logout_url, status_code=303) diff --git a/modules/keycloak/keycloak/jwks.py b/modules/keycloak/keycloak/jwks.py new file mode 100644 index 00000000..1c6af3d2 --- /dev/null +++ b/modules/keycloak/keycloak/jwks.py @@ -0,0 +1,113 @@ +"""JWKS key cache and JWT validation for Keycloak tokens.""" + +from __future__ import annotations + +import logging +import time +from typing import Any + +import httpx +import jwt +from jwt.algorithms import RSAAlgorithm + +logger = logging.getLogger(__name__) + + +class JWKSCache: + """Caches Keycloak's public signing keys and validates JWTs. + + On validation failure with cached keys, refetches JWKS once before + rejecting -- this handles Keycloak key rotation gracefully. + """ + + def __init__( + self, + jwks_url: str, + ttl_seconds: int = 3600, + *, + issuer: str, + audience: str, + ) -> None: + if not issuer: + raise ValueError("issuer is required for JWT validation") + if not audience: + raise ValueError("audience is required for JWT validation") + self._jwks_url = jwks_url + self._ttl = ttl_seconds + self._issuer = issuer + self._audience = audience + self._keys: dict[str, Any] = {} + self._fetched_at: float = 0 + + async def validate_jwt(self, token: str) -> dict[str, Any] | None: + """Decode and validate a JWT. Returns claims dict or None.""" + try: + unverified = jwt.get_unverified_header(token) + except jwt.exceptions.DecodeError: + return None + + kid = unverified.get("kid") + if kid is None: + return None + + key = await self._get_key(kid) + if key is None: + return None + + return self._decode(token, key) + + def _decode(self, token: str, key: Any) -> dict[str, Any] | None: + try: + return jwt.decode( + token, + key, + algorithms=["RS256"], + issuer=self._issuer, + audience=self._audience, + ) + except ( + jwt.ExpiredSignatureError, + jwt.InvalidIssuerError, + jwt.InvalidAudienceError, + ): + return None + except jwt.PyJWTError: + logger.exception("JWT validation failed") + return None + + async def _get_key(self, kid: str) -> Any | None: + if self._is_stale() or kid not in self._keys: + await self._fetch_keys() + + if kid in self._keys: + return self._keys[kid] + + await self._fetch_keys(force=True) + return self._keys.get(kid) + + def _is_stale(self) -> bool: + return time.monotonic() - self._fetched_at > self._ttl + + async def _fetch_keys(self, *, force: bool = False) -> None: + if not force and not self._is_stale(): + return + try: + async with httpx.AsyncClient() as client: + resp = await client.get(self._jwks_url, timeout=10) + resp.raise_for_status() + jwks_data = resp.json() + except Exception: + logger.exception("Failed to fetch JWKS from %s", self._jwks_url) + return + + new_keys: dict[str, Any] = {} + for key_data in jwks_data.get("keys", []): + kid = key_data.get("kid") + if kid and key_data.get("alg") == "RS256": + try: + public_key = RSAAlgorithm.from_jwk(key_data) + new_keys[kid] = public_key + except Exception: + logger.warning("Failed to parse JWK kid=%s", kid) + self._keys = new_keys + self._fetched_at = time.monotonic() diff --git a/modules/keycloak/keycloak/locales/en.json b/modules/keycloak/keycloak/locales/en.json new file mode 100644 index 00000000..a5d10d40 --- /dev/null +++ b/modules/keycloak/keycloak/locales/en.json @@ -0,0 +1,15 @@ +{ + "login": { + "redirecting": "Redirecting to identity provider…", + "title": "Sign In" + }, + "logout": { + "title": "Signed Out", + "message": "You have been signed out successfully." + }, + "errors": { + "callback_failed": "Authentication failed. Please try again.", + "invalid_state": "Invalid authentication state. Please try again.", + "token_validation_failed": "Token validation failed." + } +} diff --git a/modules/keycloak/keycloak/models.py b/modules/keycloak/keycloak/models.py new file mode 100644 index 00000000..0183b42e --- /dev/null +++ b/modules/keycloak/keycloak/models.py @@ -0,0 +1,21 @@ +"""Keycloak user cache -- maps Keycloak sub to a stable framework UUID.""" + +from __future__ import annotations + +import uuid as uuid_mod +from datetime import datetime + +from simple_module_db.base import create_module_base +from sqlmodel import Field + +Base = create_module_base("keycloak") + + +class KeycloakUserCache(Base, table=True): + __tablename__ = "keycloak_user_cache" + + id: uuid_mod.UUID = Field(default_factory=uuid_mod.uuid4, primary_key=True) + keycloak_sub: str = Field(unique=True, index=True) + email: str = "" + full_name: str | None = None + last_login_at: datetime | None = None diff --git a/modules/keycloak/keycloak/module.py b/modules/keycloak/keycloak/module.py new file mode 100644 index 00000000..0576cd20 --- /dev/null +++ b/modules/keycloak/keycloak/module.py @@ -0,0 +1,83 @@ +"""Keycloak OIDC authentication module.""" + +from __future__ import annotations + +import importlib.resources +from pathlib import Path +from typing import TYPE_CHECKING + +from simple_module_core.menu import MenuItem, MenuRegistry, MenuSection +from simple_module_core.module import ModuleBase, ModuleMeta + +if TYPE_CHECKING: + from fastapi import APIRouter, FastAPI + +_MODULE_DEPENDENCY_AUTH = "Auth" +_MODULE_DEPENDENCY_SETTINGS = "Settings" + + +class KeycloakModule(ModuleBase): + meta = ModuleMeta( + name="Keycloak", + route_prefix="/api/keycloak", + view_prefix="/keycloak", + depends_on=[_MODULE_DEPENDENCY_AUTH, _MODULE_DEPENDENCY_SETTINGS], + ) + _is_auth_provider = True + + def register_settings(self, app: FastAPI) -> None: + import importlib + + from keycloak.provider import KeycloakAuthProvider + from keycloak.settings import KeycloakSettings + from keycloak.state import KeycloakState + + register_module_settings = importlib.import_module( + "settings.registration" + ).register_module_settings + + register_module_settings( + app, + "keycloak", + KeycloakSettings, + lambda s: KeycloakState(settings=s), + ) + + app.state.auth.auth_provider = KeycloakAuthProvider(app.state.keycloak.settings) + + def register_menu_items(self, registry: MenuRegistry) -> None: + registry.add( + MenuItem( + label="Logout", + url="/keycloak/logout", + icon="log-out", + order=999, + section=MenuSection.USER_DROPDOWN, + method="post", + ) + ) + + def register_routes(self, api_router: APIRouter, view_router: APIRouter) -> None: + from keycloak.endpoints.api import router as api + from keycloak.endpoints.views import router as views + + api_router.include_router(api) + view_router.include_router(views) + + async def on_startup(self, app: FastAPI) -> None: + from keycloak.jwks import JWKSCache + + state = app.state.keycloak + s = state.settings + if s.server_url and s.realm: + state.jwks_cache = JWKSCache( + jwks_url=(f"{s.server_url}/realms/{s.realm}/protocol/openid-connect/certs"), + ttl_seconds=s.jwks_cache_ttl_seconds, + issuer=f"{s.server_url}/realms/{s.realm}", + audience=s.client_id, + ) + provider = app.state.auth.auth_provider + provider.jwks_cache = state.jwks_cache + + def locale_dirs(self) -> dict[str, Path]: + return {"keycloak": Path(str(importlib.resources.files(__package__) / "locales"))} diff --git a/modules/keycloak/keycloak/oidc.py b/modules/keycloak/keycloak/oidc.py new file mode 100644 index 00000000..6439da67 --- /dev/null +++ b/modules/keycloak/keycloak/oidc.py @@ -0,0 +1,85 @@ +"""OIDC helpers for Keycloak -- authorization URL, token exchange, logout.""" + +from __future__ import annotations + +import secrets +from typing import Any +from urllib.parse import urlencode + +import httpx + + +class OIDCClient: + """Thin wrapper around Keycloak's OIDC endpoints.""" + + def __init__( + self, + server_url: str, + realm: str, + client_id: str, + client_secret: str, + ) -> None: + self._base = f"{server_url.rstrip('/')}/realms/{realm}/protocol/openid-connect" + self._client_id = client_id + self._client_secret = client_secret + self._server_url = server_url.rstrip("/") + self._realm = realm + + @property + def issuer(self) -> str: + return f"{self._server_url}/realms/{self._realm}" + + @property + def token_endpoint(self) -> str: + return f"{self._base}/token" + + @property + def jwks_url(self) -> str: + return f"{self._base}/certs" + + def build_authorization_url( + self, + redirect_uri: str, + nonce: str, + scope: str = "openid email profile", + ) -> tuple[str, str]: + state = secrets.token_urlsafe(32) + params = { + "client_id": self._client_id, + "redirect_uri": redirect_uri, + "response_type": "code", + "scope": scope, + "state": state, + "nonce": nonce, + } + url = f"{self._base}/auth?{urlencode(params)}" + return url, state + + async def exchange_code( + self, + code: str, + redirect_uri: str, + ) -> dict[str, Any]: + data = { + "grant_type": "authorization_code", + "code": code, + "redirect_uri": redirect_uri, + "client_id": self._client_id, + "client_secret": self._client_secret, + } + async with httpx.AsyncClient() as client: + resp = await client.post(self.token_endpoint, data=data, timeout=10) + resp.raise_for_status() + return resp.json() + + def build_logout_url( + self, + post_logout_redirect_uri: str, + id_token_hint: str | None = None, + ) -> str: + params: dict[str, str] = { + "post_logout_redirect_uri": post_logout_redirect_uri, + } + if id_token_hint: + params["id_token_hint"] = id_token_hint + return f"{self._base}/logout?{urlencode(params)}" diff --git a/modules/keycloak/keycloak/pages/LoggedOut.tsx b/modules/keycloak/keycloak/pages/LoggedOut.tsx new file mode 100644 index 00000000..8e1b961e --- /dev/null +++ b/modules/keycloak/keycloak/pages/LoggedOut.tsx @@ -0,0 +1,13 @@ +import { Link } from '@inertiajs/react'; + +export default function LoggedOut() { + return ( +
+

Signed Out

+

You have been signed out successfully.

+ + Sign in again + +
+ ); +} diff --git a/modules/keycloak/keycloak/pages/Login.tsx b/modules/keycloak/keycloak/pages/Login.tsx new file mode 100644 index 00000000..4e4641ee --- /dev/null +++ b/modules/keycloak/keycloak/pages/Login.tsx @@ -0,0 +1,14 @@ +import { router } from '@inertiajs/react'; +import { useEffect } from 'react'; + +export default function Login() { + useEffect(() => { + router.get('/api/keycloak/auth/login'); + }, []); + + return ( +
+

Redirecting to identity provider...

+
+ ); +} diff --git a/modules/keycloak/keycloak/provider.py b/modules/keycloak/keycloak/provider.py new file mode 100644 index 00000000..bc705aad --- /dev/null +++ b/modules/keycloak/keycloak/provider.py @@ -0,0 +1,138 @@ +"""KeycloakAuthProvider — resolves users from Keycloak JWTs or session.""" + +from __future__ import annotations + +import logging +from datetime import UTC +from typing import TYPE_CHECKING, Any + +from auth.contracts.schemas import UserContext +from starlette.requests import Request + +if TYPE_CHECKING: + from keycloak.jwks import JWKSCache + from keycloak.settings import KeycloakSettings + +logger = logging.getLogger(__name__) + +_SESSION_USER_CTX_KEY = "user_ctx" + + +class KeycloakAuthProvider: + """OIDC auth provider backed by Keycloak.""" + + name = "keycloak" + _is_auth_provider = True + + def __init__(self, settings: KeycloakSettings | None = None) -> None: + self._settings = settings + self.jwks_cache: JWKSCache | None = None + + async def resolve_user(self, request: Request) -> UserContext | None: + auth_header = request.headers.get("authorization", "") + if auth_header.startswith("Bearer "): + return await self._resolve_bearer(request, auth_header[7:]) + + session = request.scope.get("session", {}) + return UserContext.from_session_dict(session.get(_SESSION_USER_CTX_KEY)) + + def get_login_url(self, request: Request | None, next_url: str | None = None) -> str: + return "/keycloak/login" + + def get_logout_url(self, request: Request | None) -> str: + return "/keycloak/logout" + + def get_public_paths(self) -> tuple[tuple[str, ...], tuple[str, ...]]: + return ( + ("/keycloak/login", "/keycloak/logout", "/api/keycloak/auth/"), + (), + ) + + def is_bearer_request(self, request: Request | None) -> bool: + if request is None: + return False + return request.headers.get("authorization", "").startswith("Bearer ") + + async def _resolve_bearer(self, request: Request, token: str) -> UserContext | None: + if self.jwks_cache is None: + logger.warning("JWKS cache not initialized; rejecting bearer token") + return None + claims = await self.jwks_cache.validate_jwt(token) + if claims is None: + return None + + cache_id = await self._upsert_user_cache(request, claims) + return self._claims_to_user_context(claims, cache_id=cache_id) + + def _claims_to_user_context( + self, + claims: dict[str, Any], + *, + cache_id: str, + ) -> UserContext: + roles_raw = ( + _extract_nested(claims, self._settings.roles_claim_path) if self._settings else None + ) + mapped = [ + self._settings.role_mapping[r] + for r in (roles_raw or []) + if self._settings and r in self._settings.role_mapping + ] + return UserContext( + id=cache_id, + email=claims.get("email", ""), + name=(claims.get("preferred_username") or claims.get("name", "")), + roles=mapped, + tenant_id=claims.get("tenant_id"), + ) + + async def _upsert_user_cache(self, request: Request, claims: dict) -> str: + try: + from sqlalchemy import select + + from keycloak.models import KeycloakUserCache + + session_factory = request.app.state.sm.db.session_factory + sub = claims["sub"] + async with session_factory() as db: + stmt = select(KeycloakUserCache).where(KeycloakUserCache.keycloak_sub == sub) + row = (await db.execute(stmt)).scalar_one_or_none() + if row is None: + import uuid as uuid_mod + from datetime import datetime + + row = KeycloakUserCache( + id=uuid_mod.uuid4(), + keycloak_sub=sub, + email=claims.get("email", ""), + full_name=claims.get("preferred_username"), + last_login_at=datetime.now(UTC), + ) + db.add(row) + await db.flush() + else: + from datetime import datetime + + row.email = claims.get("email", row.email) + row.full_name = claims.get("preferred_username", row.full_name) + row.last_login_at = datetime.now(UTC) + await db.flush() + return str(row.id) + except Exception: + logger.exception( + "Failed to upsert KeycloakUserCache for sub=%s", + claims.get("sub"), + ) + return claims.get("sub", "unknown") + + +def _extract_nested(data: dict, path: str) -> list[str] | None: + parts = path.split(".") + current: Any = data + for part in parts: + if not isinstance(current, dict): + return None + current = current.get(part) + if current is None: + return None + return current if isinstance(current, list) else None diff --git a/modules/keycloak/keycloak/settings.py b/modules/keycloak/keycloak/settings.py new file mode 100644 index 00000000..e9e9745c --- /dev/null +++ b/modules/keycloak/keycloak/settings.py @@ -0,0 +1,49 @@ +"""Keycloak module settings -- DB-backed via ``register_module_settings``.""" + +from __future__ import annotations + +from pydantic import Field, model_validator +from pydantic_settings import BaseSettings, SettingsConfigDict +from simple_module_core.dotenv import env_str +from simple_module_core.environments import NON_PROD_ENVIRONMENTS + + +class KeycloakSettings(BaseSettings): + """Keycloak OIDC configuration.""" + + model_config = SettingsConfigDict(extra="ignore") + + server_url: str = env_str("SM_KEYCLOAK_SERVER_URL", "") + realm: str = env_str("SM_KEYCLOAK_REALM", "") + client_id: str = env_str("SM_KEYCLOAK_CLIENT_ID", "") + client_secret: str = env_str("SM_KEYCLOAK_CLIENT_SECRET", "") + + roles_claim_path: str = "realm_access.roles" + admin_role: str = "admin" + login_redirect_url: str = "/dashboard/" + jwks_cache_ttl_seconds: int = 3600 + + role_mapping: dict[str, str] = Field( + default_factory=lambda: {"admin": "admin", "user": "user"}, + ) + + @model_validator(mode="after") + def _check_required_in_production(self) -> KeycloakSettings: + import os + + env = os.environ.get("SM_ENVIRONMENT", "development") + if env in NON_PROD_ENVIRONMENTS: + return self + missing = [] + if not self.server_url: + missing.append("SM_KEYCLOAK_SERVER_URL") + if not self.realm: + missing.append("SM_KEYCLOAK_REALM") + if not self.client_id: + missing.append("SM_KEYCLOAK_CLIENT_ID") + if not self.client_secret: + missing.append("SM_KEYCLOAK_CLIENT_SECRET") + if missing: + msg = f"Keycloak settings required in production: {', '.join(missing)}" + raise ValueError(msg) + return self diff --git a/modules/keycloak/keycloak/state.py b/modules/keycloak/keycloak/state.py new file mode 100644 index 00000000..024449af --- /dev/null +++ b/modules/keycloak/keycloak/state.py @@ -0,0 +1,18 @@ +"""Module-scoped state container for the keycloak module.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from keycloak.jwks import JWKSCache + from keycloak.settings import KeycloakSettings + + +@dataclass +class KeycloakState: + """Keycloak-module singletons. Single slot at ``app.state.keycloak``.""" + + settings: KeycloakSettings + jwks_cache: JWKSCache | None = None diff --git a/modules/keycloak/package.json b/modules/keycloak/package.json new file mode 100644 index 00000000..cc4f33f4 --- /dev/null +++ b/modules/keycloak/package.json @@ -0,0 +1,7 @@ +{ + "name": "@simple-module/keycloak", + "private": true, + "version": "0.0.0", + "type": "module", + "dependencies": {} +} diff --git a/modules/keycloak/pyproject.toml b/modules/keycloak/pyproject.toml new file mode 100644 index 00000000..27d9da86 --- /dev/null +++ b/modules/keycloak/pyproject.toml @@ -0,0 +1,54 @@ +[project] +name = "simple_module_keycloak" +version = "0.0.15" +description = "Keycloak OIDC authentication provider for simple_module — swap with simple_module_users" +readme = "README.md" +license = "MIT" +requires-python = ">=3.12" +authors = [{ name = "Anto Subash", email = "antosubash@live.com" }] +keywords = ["simple-module", "keycloak", "oidc", "authentication", "fastapi"] +classifiers = [ + "Development Status :: 3 - Alpha", + "Framework :: FastAPI", + "Intended Audience :: Developers", + "License :: OSI Approved :: MIT License", + "Operating System :: OS Independent", + "Programming Language :: Python :: 3", + "Programming Language :: Python :: 3.12", + "Topic :: Internet :: WWW/HTTP", + "Topic :: Software Development :: Libraries :: Application Frameworks", + "Typing :: Typed", +] +dependencies = [ + "simple_module_core==0.0.15", + "simple_module_db==0.0.15", + "simple_module_hosting==0.0.15", + "simple_module_settings==0.0.15", + "simple_module_auth==0.0.15", + "PyJWT[crypto]>=2.8", + "httpx>=0.27", +] + +[project.entry-points.simple_module] +keycloak = "keycloak.module:KeycloakModule" + +[project.urls] +Homepage = "https://github.com/antosubash/simple_module_python" +Repository = "https://github.com/antosubash/simple_module_python" + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[tool.hatch.build.targets.wheel] +packages = ["keycloak"] + +[tool.hatch.build.targets.wheel.force-include] +"package.json" = "keycloak/package.json" + +[tool.uv.sources] +simple_module_core = { workspace = true } +simple_module_db = { workspace = true } +simple_module_hosting = { workspace = true } +simple_module_settings = { workspace = true } +simple_module_auth = { workspace = true } diff --git a/modules/keycloak/tests/conftest.py b/modules/keycloak/tests/conftest.py new file mode 100644 index 00000000..727ed731 --- /dev/null +++ b/modules/keycloak/tests/conftest.py @@ -0,0 +1 @@ +"""Keycloak module test fixtures.""" diff --git a/modules/keycloak/tests/test_jwks.py b/modules/keycloak/tests/test_jwks.py new file mode 100644 index 00000000..d9b6a0df --- /dev/null +++ b/modules/keycloak/tests/test_jwks.py @@ -0,0 +1,157 @@ +"""Tests for JWKS key cache and JWT validation.""" + +from __future__ import annotations + +import json +import time + +import jwt +import pytest +from cryptography.hazmat.primitives.asymmetric import rsa +from keycloak.jwks import JWKSCache + + +def _generate_rsa_keypair(): + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_key = private_key.public_key() + return private_key, public_key + + +def _make_jwks_response(public_key, kid="test-key-1"): + from jwt.algorithms import RSAAlgorithm + + jwk = json.loads(RSAAlgorithm.to_jwk(public_key)) + jwk["kid"] = kid + jwk["use"] = "sig" + jwk["alg"] = "RS256" + return {"keys": [jwk]} + + +def _sign_token(private_key, payload, kid="test-key-1"): + return jwt.encode(payload, private_key, algorithm="RS256", headers={"kid": kid}) + + +@pytest.fixture +def rsa_keys(): + return _generate_rsa_keypair() + + +@pytest.fixture +def valid_payload(): + now = int(time.time()) + return { + "sub": "user-123", + "email": "test@example.com", + "preferred_username": "testuser", + "iss": "https://auth.example.com/realms/test", + "aud": "my-client", + "exp": now + 3600, + "iat": now, + "realm_access": {"roles": ["admin", "user"]}, + } + + +async def test_validate_jwt_valid_token(rsa_keys, valid_payload, httpx_mock): + private_key, public_key = rsa_keys + jwks_data = _make_jwks_response(public_key) + httpx_mock.add_response(url="https://auth.example.com/jwks", json=jwks_data) + + cache = JWKSCache( + jwks_url="https://auth.example.com/jwks", + ttl_seconds=3600, + issuer="https://auth.example.com/realms/test", + audience="my-client", + ) + + token = _sign_token(private_key, valid_payload) + claims = await cache.validate_jwt(token) + assert claims is not None + assert claims["sub"] == "user-123" + assert claims["email"] == "test@example.com" + + +async def test_validate_jwt_expired_token(rsa_keys, valid_payload, httpx_mock): + private_key, public_key = rsa_keys + valid_payload["exp"] = int(time.time()) - 100 + jwks_data = _make_jwks_response(public_key) + httpx_mock.add_response(url="https://auth.example.com/jwks", json=jwks_data) + + cache = JWKSCache( + jwks_url="https://auth.example.com/jwks", + ttl_seconds=3600, + issuer="https://auth.example.com/realms/test", + audience="my-client", + ) + + token = _sign_token(private_key, valid_payload) + claims = await cache.validate_jwt(token) + assert claims is None + + +async def test_validate_jwt_wrong_issuer(rsa_keys, valid_payload, httpx_mock): + private_key, public_key = rsa_keys + jwks_data = _make_jwks_response(public_key) + httpx_mock.add_response(url="https://auth.example.com/jwks", json=jwks_data) + + cache = JWKSCache( + jwks_url="https://auth.example.com/jwks", + ttl_seconds=3600, + issuer="https://wrong-issuer.com/realms/test", + audience="my-client", + ) + + token = _sign_token(private_key, valid_payload) + claims = await cache.validate_jwt(token) + assert claims is None + + +async def test_validate_jwt_wrong_audience(rsa_keys, valid_payload, httpx_mock): + private_key, public_key = rsa_keys + jwks_data = _make_jwks_response(public_key) + httpx_mock.add_response(url="https://auth.example.com/jwks", json=jwks_data) + + cache = JWKSCache( + jwks_url="https://auth.example.com/jwks", + ttl_seconds=3600, + issuer="https://auth.example.com/realms/test", + audience="wrong-client", + ) + + token = _sign_token(private_key, valid_payload) + claims = await cache.validate_jwt(token) + assert claims is None + + +async def test_jwks_cache_refetches_on_unknown_kid(rsa_keys, valid_payload): + """When a token has a kid not in cache, refetch JWKS once before rejecting.""" + private_key, public_key = rsa_keys + jwks_data = _make_jwks_response(public_key, kid="rotated-key") + + call_count = 0 + + async def counting_fetch(self, *, force=False): + nonlocal call_count + call_count += 1 + if call_count == 1: + self._keys = {} + self._fetched_at = __import__("time").monotonic() + else: + from jwt.algorithms import RSAAlgorithm + + for kd in jwks_data["keys"]: + self._keys[kd["kid"]] = RSAAlgorithm.from_jwk(kd) + self._fetched_at = __import__("time").monotonic() + + cache = JWKSCache( + jwks_url="https://auth.example.com/jwks", + ttl_seconds=3600, + issuer="https://auth.example.com/realms/test", + audience="my-client", + ) + cache._fetch_keys = counting_fetch.__get__(cache, JWKSCache) + + token = _sign_token(private_key, valid_payload, kid="rotated-key") + claims = await cache.validate_jwt(token) + assert claims is not None + assert claims["sub"] == "user-123" + assert call_count == 2 diff --git a/modules/keycloak/tests/test_keycloak_module.py b/modules/keycloak/tests/test_keycloak_module.py new file mode 100644 index 00000000..7b073d11 --- /dev/null +++ b/modules/keycloak/tests/test_keycloak_module.py @@ -0,0 +1,22 @@ +"""Tests for KeycloakModule lifecycle.""" + +from __future__ import annotations + +from auth.contracts.provider import AuthProvider + + +def test_keycloak_module_meta(): + from keycloak.module import KeycloakModule + + mod = KeycloakModule() + assert mod.meta.name == "Keycloak" + assert mod.meta.depends_on == ["Auth", "Settings"] + assert mod._is_auth_provider is True + + +def test_keycloak_provider_satisfies_protocol(): + from keycloak.provider import KeycloakAuthProvider + + provider = KeycloakAuthProvider() + assert isinstance(provider, AuthProvider) + assert provider.name == "keycloak" diff --git a/modules/keycloak/tests/test_keycloak_provider.py b/modules/keycloak/tests/test_keycloak_provider.py new file mode 100644 index 00000000..4b2982a4 --- /dev/null +++ b/modules/keycloak/tests/test_keycloak_provider.py @@ -0,0 +1,94 @@ +"""Tests for KeycloakAuthProvider.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest +from auth.contracts.provider import AuthProvider +from auth.contracts.schemas import UserContext +from keycloak.provider import KeycloakAuthProvider +from keycloak.settings import KeycloakSettings + + +@pytest.fixture +def settings(): + return KeycloakSettings( + server_url="https://auth.example.com", + realm="test", + client_id="my-app", + client_secret="secret", + role_mapping={ + "admin": "admin", + "user": "user", + "editor": "editor", + }, + ) + + +@pytest.fixture +def provider(settings): + return KeycloakAuthProvider(settings) + + +def test_satisfies_protocol(provider): + assert isinstance(provider, AuthProvider) + + +def test_name(provider): + assert provider.name == "keycloak" + + +def test_login_url(provider): + assert provider.get_login_url(None) == "/keycloak/login" + + +def test_logout_url(provider): + assert provider.get_logout_url(None) == "/keycloak/logout" + + +def test_public_paths(provider): + prefixes, _exact = provider.get_public_paths() + assert "/keycloak/login" in prefixes + assert "/api/keycloak/auth/" in prefixes + + +def test_is_bearer_request(provider): + req = MagicMock() + req.headers = {"authorization": "Bearer abc"} + assert provider.is_bearer_request(req) is True + + req.headers = {} + assert provider.is_bearer_request(req) is False + + +def test_claims_to_user_context(provider): + claims = { + "sub": "kc-user-123", + "email": "test@example.com", + "preferred_username": "testuser", + "realm_access": {"roles": ["admin", "unknown_role", "user"]}, + } + ctx = provider._claims_to_user_context(claims, cache_id="aaaaaaaa-0000-0000-0000-000000000001") + assert isinstance(ctx, UserContext) + assert ctx.id == "aaaaaaaa-0000-0000-0000-000000000001" + assert ctx.email == "test@example.com" + assert ctx.name == "testuser" + assert sorted(ctx.roles) == ["admin", "user"] + + +def test_claims_to_user_context_no_roles(provider): + claims = {"sub": "kc-user-456", "email": "noroles@example.com"} + ctx = provider._claims_to_user_context(claims, cache_id="bbbb") + assert ctx.roles == [] + + +def test_extract_roles_custom_claim_path(settings): + settings.roles_claim_path = "resource_access.my-app.roles" + provider = KeycloakAuthProvider(settings) + claims = { + "sub": "x", + "resource_access": {"my-app": {"roles": ["admin"]}}, + } + ctx = provider._claims_to_user_context(claims, cache_id="cccc") + assert ctx.roles == ["admin"] diff --git a/modules/keycloak/tests/test_oidc.py b/modules/keycloak/tests/test_oidc.py new file mode 100644 index 00000000..c278afea --- /dev/null +++ b/modules/keycloak/tests/test_oidc.py @@ -0,0 +1,70 @@ +"""Tests for OIDC discovery and token exchange helpers.""" + +from __future__ import annotations + +import pytest +from keycloak.oidc import OIDCClient + + +@pytest.fixture +def oidc_client(): + return OIDCClient( + server_url="https://auth.example.com", + realm="test", + client_id="my-app", + client_secret="secret123", + ) + + +def test_authorization_url(oidc_client): + url, state = oidc_client.build_authorization_url( + redirect_uri="https://app.example.com/callback", + nonce="test-nonce", + ) + assert "auth.example.com/realms/test/protocol/openid-connect/auth" in url + assert "client_id=my-app" in url + assert "redirect_uri=" in url + assert "response_type=code" in url + assert "scope=openid" in url + assert "nonce=test-nonce" in url + assert state is not None + assert len(state) > 0 + + +def test_token_endpoint_url(oidc_client): + assert oidc_client.token_endpoint == ( + "https://auth.example.com/realms/test/protocol/openid-connect/token" + ) + + +def test_logout_url(oidc_client): + url = oidc_client.build_logout_url( + post_logout_redirect_uri="https://app.example.com/login", + id_token_hint="token123", + ) + assert "auth.example.com/realms/test/protocol/openid-connect/logout" in url + assert "post_logout_redirect_uri=" in url + assert "id_token_hint=token123" in url + + +def test_issuer(oidc_client): + assert oidc_client.issuer == "https://auth.example.com/realms/test" + + +async def test_exchange_code(oidc_client, httpx_mock): + httpx_mock.add_response( + url=oidc_client.token_endpoint, + json={ + "access_token": "at-123", + "id_token": "id-123", + "refresh_token": "rt-123", + "token_type": "Bearer", + "expires_in": 300, + }, + ) + tokens = await oidc_client.exchange_code( + code="auth-code-xyz", + redirect_uri="https://app.example.com/callback", + ) + assert tokens["access_token"] == "at-123" + assert tokens["id_token"] == "id-123" diff --git a/modules/keycloak/tsconfig.json b/modules/keycloak/tsconfig.json new file mode 100644 index 00000000..d479e6d7 --- /dev/null +++ b/modules/keycloak/tsconfig.json @@ -0,0 +1,4 @@ +{ + "extends": "../../host/client_app/tsconfig.json", + "include": ["keycloak/**/*.ts", "keycloak/**/*.tsx"] +} diff --git a/modules/settings/settings/_module_settings.py b/modules/settings/settings/_module_settings.py index 39c7e148..ad3c8176 100644 --- a/modules/settings/settings/_module_settings.py +++ b/modules/settings/settings/_module_settings.py @@ -77,17 +77,34 @@ def _extract_settings(app: FastAPI, package: str) -> BaseSettings | None: return inner if isinstance(inner, BaseSettings) else None +def _resolve_default(info) -> Any: + """Return the effective default, handling ``default_factory`` fields. + + Pydantic sets ``info.default`` to ``PydanticUndefined`` when + ``default_factory`` is used. We call the factory to get the concrete + default so the settings UI can serialize it. + """ + from pydantic_core import PydanticUndefined + + if info.default is not PydanticUndefined: + return info.default + if info.default_factory is not None: + return info.default_factory() + return None + + def _field_view(name: str, settings: BaseSettings, prefix: str) -> ModuleSettingField: cls = type(settings) info = cls.model_fields[name] raw_value = getattr(settings, name) secret = is_secret_field(name) extra = info.json_schema_extra if isinstance(info.json_schema_extra, dict) else {} + default = _resolve_default(info) return ModuleSettingField( name=name, env_var=f"{prefix}{name.upper()}", value=_mask(raw_value) if secret else raw_value, - default=_mask(info.default) if secret else info.default, + default=_mask(default) if secret else default, description=info.description or "", is_secret=secret, type=value_type_for_field(cls, name), diff --git a/modules/users/tests/_middleware_support.py b/modules/users/tests/_middleware_support.py index e4454f41..0744746a 100644 --- a/modules/users/tests/_middleware_support.py +++ b/modules/users/tests/_middleware_support.py @@ -14,12 +14,12 @@ from typing import Any import pytest +from auth.middleware import AuthMiddleware from fastapi import FastAPI, Request from simple_module_test import forge_session_cookie from starlette.middleware.sessions import SessionMiddleware from starlette.responses import JSONResponse from users.constants import ADMIN_ROLE_ID, USER_ROLE_ID -from users.middleware import AuthMiddleware SECRET_KEY = "test-secret-key-for-session-middleware" @@ -61,7 +61,10 @@ async def _default_handler(request: Request): app = FastAPI() app.state.sm = SimpleNamespace(db=db_state) + from users.provider import UsersAuthProvider + app.state.auth = AuthState( + auth_provider=UsersAuthProvider(), principal_resolvers=list(principal_resolvers or []), ) diff --git a/modules/users/tests/test_token_api.py b/modules/users/tests/test_token_api.py new file mode 100644 index 00000000..0c012eca --- /dev/null +++ b/modules/users/tests/test_token_api.py @@ -0,0 +1,191 @@ +"""Tests for bearer token endpoints: /api/users/auth/token*. + +Covers: login via email+password, refresh rotation, revoke, and error paths. +""" + +from __future__ import annotations + +import uuid + +import pytest +from fastapi_users.password import PasswordHelper +from users.models import User + +_pw = PasswordHelper() + + +def _hash(plain: str) -> str: + return _pw.hash(plain) + + +async def _seed_user(session, email="api@example.com", password="SecurePass1!"): + """Create a verified, active user for token tests.""" + user = User( + id=uuid.uuid4(), + email=email, + hashed_password=_hash(password), + is_active=True, + is_superuser=False, + is_verified=True, + ) + session.add(user) + await session.commit() + await session.refresh(user) + return user + + +# --------------------------------------------------------------------------- +# POST /api/users/auth/token — login +# --------------------------------------------------------------------------- + + +class TestTokenLogin: + @pytest.mark.anyio + async def test_invalid_email_returns_401(self, anon_client): + resp = await anon_client.post( + "/api/users/auth/token", + json={"email": "nobody@example.com", "password": "whatever"}, + ) + assert resp.status_code == 401 + assert resp.json()["detail"] == "Invalid credentials" + + @pytest.mark.anyio + async def test_wrong_password_returns_401(self, anon_client, users_db): + await _seed_user(users_db) + resp = await anon_client.post( + "/api/users/auth/token", + json={"email": "api@example.com", "password": "WRONG"}, + ) + assert resp.status_code == 401 + assert resp.json()["detail"] == "Invalid credentials" + + @pytest.mark.anyio + async def test_inactive_user_returns_401(self, anon_client, users_db): + user = await _seed_user(users_db, email="inactive@example.com") + user.is_active = False + await users_db.commit() + + resp = await anon_client.post( + "/api/users/auth/token", + json={"email": "inactive@example.com", "password": "SecurePass1!"}, + ) + assert resp.status_code == 401 + + @pytest.mark.anyio + async def test_valid_credentials_returns_token_pair(self, anon_client, users_db): + await _seed_user(users_db) + resp = await anon_client.post( + "/api/users/auth/token", + json={"email": "api@example.com", "password": "SecurePass1!"}, + ) + assert resp.status_code == 200 + body = resp.json() + assert body["token_type"] == "bearer" + assert body["access_token"] + assert body["refresh_token"] + assert body["expires_in"] > 0 + + +# --------------------------------------------------------------------------- +# POST /api/users/auth/token/refresh +# --------------------------------------------------------------------------- + + +class TestTokenRefresh: + @pytest.mark.anyio + async def test_invalid_uuid_returns_401(self, anon_client): + resp = await anon_client.post( + "/api/users/auth/token/refresh", + json={"refresh_token": "not-a-uuid"}, + ) + assert resp.status_code == 401 + assert resp.json()["detail"] == "Invalid refresh token" + + @pytest.mark.anyio + async def test_nonexistent_token_returns_401(self, anon_client): + resp = await anon_client.post( + "/api/users/auth/token/refresh", + json={"refresh_token": str(uuid.uuid4())}, + ) + assert resp.status_code == 401 + assert resp.json()["detail"] == "Invalid or expired refresh token" + + @pytest.mark.anyio + async def test_valid_refresh_rotates_tokens(self, anon_client, users_db): + """Login, then refresh — old refresh revoked, new pair returned.""" + await _seed_user(users_db) + login = await anon_client.post( + "/api/users/auth/token", + json={"email": "api@example.com", "password": "SecurePass1!"}, + ) + assert login.status_code == 200 + old_refresh = login.json()["refresh_token"] + + refresh_resp = await anon_client.post( + "/api/users/auth/token/refresh", + json={"refresh_token": old_refresh}, + ) + assert refresh_resp.status_code == 200 + new_body = refresh_resp.json() + assert new_body["access_token"] + assert new_body["refresh_token"] != old_refresh + + # Old refresh token should now be revoked + reuse = await anon_client.post( + "/api/users/auth/token/refresh", + json={"refresh_token": old_refresh}, + ) + assert reuse.status_code == 401 + + +# --------------------------------------------------------------------------- +# DELETE /api/users/auth/token — revoke +# --------------------------------------------------------------------------- + + +class TestTokenRevoke: + @pytest.mark.anyio + async def test_invalid_format_returns_400(self, anon_client): + resp = await anon_client.request( + "DELETE", + "/api/users/auth/token", + json={"refresh_token": "garbage"}, + ) + assert resp.status_code == 400 + assert resp.json()["detail"] == "Invalid token format" + + @pytest.mark.anyio + async def test_nonexistent_token_returns_ok(self, anon_client): + """Revoking a non-existent token is idempotent — still returns ok.""" + resp = await anon_client.request( + "DELETE", + "/api/users/auth/token", + json={"refresh_token": str(uuid.uuid4())}, + ) + assert resp.status_code == 200 + assert resp.json()["status"] == "ok" + + @pytest.mark.anyio + async def test_revoke_makes_refresh_unusable(self, anon_client, users_db): + """After revoking, the refresh token can no longer be used.""" + await _seed_user(users_db) + login = await anon_client.post( + "/api/users/auth/token", + json={"email": "api@example.com", "password": "SecurePass1!"}, + ) + rt = login.json()["refresh_token"] + + # Revoke + revoke_resp = await anon_client.request( + "DELETE", + "/api/users/auth/token", + json={"refresh_token": rt}, + ) + assert revoke_resp.status_code == 200 + + # Attempt refresh — should fail + refresh_resp = await anon_client.post( + "/api/users/auth/token/refresh", + json={"refresh_token": rt}, + ) + assert refresh_resp.status_code == 401 diff --git a/modules/users/tests/test_users_provider.py b/modules/users/tests/test_users_provider.py new file mode 100644 index 00000000..5772f974 --- /dev/null +++ b/modules/users/tests/test_users_provider.py @@ -0,0 +1,46 @@ +"""Tests for UsersAuthProvider.""" + +from __future__ import annotations + +from auth.contracts.provider import AuthProvider +from users.provider import UsersAuthProvider + + +def test_users_provider_satisfies_protocol(): + provider = UsersAuthProvider() + assert isinstance(provider, AuthProvider) + + +def test_login_url(): + provider = UsersAuthProvider() + assert provider.get_login_url(None) == "/users/login" + + +def test_logout_url(): + provider = UsersAuthProvider() + assert provider.get_logout_url(None) == "/users/logout" + + +def test_public_paths(): + provider = UsersAuthProvider() + prefixes, _exact = provider.get_public_paths() + assert "/users/login" in prefixes + assert "/api/users/auth/" in prefixes + + +def test_is_bearer_request_true(): + from unittest.mock import MagicMock + + request = MagicMock() + request.headers = {"authorization": "Bearer abc123"} + provider = UsersAuthProvider() + assert provider.is_bearer_request(request) is True + + +def test_is_bearer_request_false(): + from unittest.mock import MagicMock + + request = MagicMock() + request.headers = {} + provider = UsersAuthProvider() + assert provider.is_bearer_request(request) is False diff --git a/modules/users/users/auth_local/api.py b/modules/users/users/auth_local/api.py index 77ba93eb..209b4e99 100644 --- a/modules/users/users/auth_local/api.py +++ b/modules/users/users/auth_local/api.py @@ -117,12 +117,22 @@ async def login( # ── Mount fastapi-users stock routers ──────────────────────────────────────── # The stock auth router (login + logout) is mounted at /auth-inner so its -# logout and other endpoints remain accessible. Our wrapper at /auth/login -# shadows the stock login endpoint. Logout is exposed via /auth-inner/logout. +# endpoints remain accessible. Our wrappers at /auth/login and /auth/logout +# shadow the stock endpoints to also manage the session cookie. auth_inner = fastapi_users.get_auth_router(auth_backend, requires_verification=True) router.include_router(auth_inner, prefix="/auth-inner") +@router.post("/auth/logout", status_code=204) +async def api_logout(request: Request): + """API logout — clears both the access-token cookie and the session.""" + request.session.clear() + cookie_name = request.app.state.users.settings.cookie_name + response = Response(status_code=204) + response.delete_cookie(cookie_name, path="/") + return response + + # ── Accept-invite (verify + set password + login, one shot) ───────────────── diff --git a/modules/users/users/auth_local/token_api.py b/modules/users/users/auth_local/token_api.py new file mode 100644 index 00000000..3b333edf --- /dev/null +++ b/modules/users/users/auth_local/token_api.py @@ -0,0 +1,159 @@ +"""Bearer token endpoints for mobile/API clients. + +Provides email+password → access_token + refresh_token, refresh, and revoke +flows for clients that cannot use browser cookies (mobile apps, CLI tools, +third-party API consumers). +""" + +from __future__ import annotations + +import uuid as uuid_mod +from datetime import UTC, datetime, timedelta + +from fastapi import APIRouter, Depends, HTTPException, Request +from simple_module_db.deps import get_db +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession +from sqlmodel import SQLModel + +from users.models import User +from users.models.refresh_token import RefreshToken + +# Pre-computed bcrypt hash used when a login attempt targets a non-existent +# user. Running verify against this takes the same time as a real check, +# preventing timing-based email enumeration. +_DUMMY_HASH = "$2b$12$LJ3m4ys3Lg/PFgWCZxEzR.ZVxFMz3yeqHEhSYmiJ9gJOPG7W3Cq2G" + +router = APIRouter(prefix="/auth", tags=["users-token"]) + + +class TokenRequest(SQLModel): + """Email + password login for bearer-token auth.""" + + email: str + password: str + + +class TokenResponse(SQLModel): + """Access + refresh token pair returned on successful auth.""" + + access_token: str + refresh_token: str + token_type: str = "bearer" + expires_in: int + + +class RefreshRequest(SQLModel): + """Body for refresh and revoke endpoints.""" + + refresh_token: str + + +@router.post("/token", response_model=TokenResponse) +async def token_login( + body: TokenRequest, + request: Request, + db: AsyncSession = Depends(get_db), +): + """Exchange email + password for an access/refresh token pair.""" + from fastapi_users.password import PasswordHelper + + helper = PasswordHelper() + + stmt = select(User).where(User.email == body.email) + user = (await db.execute(stmt)).scalar_one_or_none() + if user is None or not user.is_active or user.disabled_at is not None: + # Constant-time: run bcrypt on a dummy hash to prevent timing-based + # email enumeration (existing user + wrong password takes ~50ms for + # bcrypt; missing user would be instant without this). + helper.verify_and_update(body.password, _DUMMY_HASH) + raise HTTPException(status_code=401, detail="Invalid credentials") + + verified, _ = helper.verify_and_update(body.password, user.hashed_password) + if not verified: + raise HTTPException(status_code=401, detail="Invalid credentials") + + settings = request.app.state.users.settings + return await _create_token_pair(db, user.id, settings) + + +@router.post("/token/refresh", response_model=TokenResponse) +async def token_refresh( + body: RefreshRequest, + request: Request, + db: AsyncSession = Depends(get_db), +): + """Rotate a refresh token into a new access/refresh pair.""" + try: + token_uuid = uuid_mod.UUID(body.refresh_token) + except (ValueError, TypeError): + raise HTTPException(status_code=401, detail="Invalid refresh token") from None + + now = datetime.now(UTC) + stmt = select(RefreshToken).where( + RefreshToken.token == token_uuid, + RefreshToken.revoked_at.is_(None), # type: ignore[union-attr] + RefreshToken.expires_at > now, + ) + rt = (await db.execute(stmt)).scalar_one_or_none() + if rt is None: + raise HTTPException(status_code=401, detail="Invalid or expired refresh token") + + rt.revoked_at = now + await db.flush() + + settings = request.app.state.users.settings + return await _create_token_pair(db, rt.user_id, settings) + + +@router.delete("/token") +async def token_revoke( + body: RefreshRequest, + db: AsyncSession = Depends(get_db), +): + """Revoke a refresh token (idempotent).""" + try: + token_uuid = uuid_mod.UUID(body.refresh_token) + except (ValueError, TypeError): + raise HTTPException(status_code=400, detail="Invalid token format") from None + + stmt = select(RefreshToken).where(RefreshToken.token == token_uuid) + rt = (await db.execute(stmt)).scalar_one_or_none() + if rt and rt.revoked_at is None: + rt.revoked_at = datetime.now(UTC) + await db.flush() + return {"status": "ok"} + + +async def _create_token_pair( + db: AsyncSession, + user_id: uuid_mod.UUID, + settings, +) -> TokenResponse: + """Mint a new access token + refresh token pair and persist both.""" + from users.models import UserAccessToken + + now = datetime.now(UTC) + + access_token = UserAccessToken( + token=str(uuid_mod.uuid4()), + user_id=user_id, + created_at=now, + ) + db.add(access_token) + + refresh = RefreshToken( + token=uuid_mod.uuid4(), + user_id=user_id, + created_at=now, + expires_at=now + timedelta(seconds=settings.refresh_token_lifetime_seconds), + ) + db.add(refresh) + await db.flush() + + return TokenResponse( + access_token=access_token.token, + refresh_token=str(refresh.token), + token_type="bearer", + expires_in=settings.bearer_token_lifetime_seconds, + ) diff --git a/modules/users/users/middleware.py b/modules/users/users/middleware.py index f78d01e7..67bfdce3 100644 --- a/modules/users/users/middleware.py +++ b/modules/users/users/middleware.py @@ -1,168 +1,10 @@ -"""Local-user auth middleware — replaces the Keycloak session reader. +"""Backwards-compatibility re-export. -Reads ``session["user_id"]``, loads the User row with eagerly-loaded roles, -builds a UserContext, and sets ``request.state.user`` + the -``current_user_id`` ContextVar consumed by DB audit listeners. - -The resolved ``UserContext`` is cached in the signed session cookie under -``session["user_ctx"]`` so subsequent requests skip the DB lookup. The cache -is refreshed when the session is cleared (logout / rotation) or when the -cached payload is missing/invalid. Trade-off: admin-side changes (role -assignment, disable/enable) do not take effect until the affected user's -session is recreated (re-login or session expiry); acceptable for this app. - -Registered via ``UsersModule.register_middleware``. +The canonical AuthMiddleware now lives in ``auth.middleware``. This shim +exists only to avoid breaking imports in downstream apps that referenced +``users.middleware.AuthMiddleware`` directly. """ -from __future__ import annotations - -import logging -import uuid - -from auth.contracts.schemas import UserContext -from simple_module_db.listeners import current_user_id -from sqlalchemy import select -from sqlalchemy.orm import selectinload -from starlette.requests import Request -from starlette.responses import JSONResponse, RedirectResponse -from starlette.types import ASGIApp, Receive, Scope, Send - -from users.constants import SESSION_USER_ID_KEY -from users.models import User - -logger = logging.getLogger(__name__) - -SESSION_USER_CTX_KEY = "user_ctx" -_SESSION_USER_ID_KEY = SESSION_USER_ID_KEY -_SESSION_NEXT_KEY = "next" -_SCOPE_HTTP = "http" -_LOGIN_REDIRECT = "/users/login" -_API_PATH_PREFIX = "/api/" -_UNAUTH_DETAIL = "Not authenticated" - -# Paths that don't require authentication. -PUBLIC_PATHS = ( - "/users/login", - "/users/register", - "/users/forgot-password", - "/users/reset-password", - "/users/verify", - "/users/invite/accept", - "/api/users/auth/", - "/api/users/register", - "/health", - "/static/", - "/api/docs", - "/api/redoc", - "/openapi.json", - "/i18n/", -) -EXACT_PUBLIC_PATHS = ("/",) - - -class AuthMiddleware: - """Redirect unauthenticated users to /users/login. - - On cache hit (``session["user_ctx"]`` present), skips the DB entirely. - On cache miss, loads the user with roles, validates active/enabled, and - writes the resolved context back to the session. - """ - - def __init__(self, app: ASGIApp) -> None: - self.app = app - - async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: - if scope["type"] != _SCOPE_HTTP: - await self.app(scope, receive, send) - return - - path = scope["path"] - is_public = any(path.startswith(p) for p in PUBLIC_PATHS) or path in EXACT_PUBLIC_PATHS - - session = scope["session"] - raw_user_id = session.get(_SESSION_USER_ID_KEY) - - user_ctx: UserContext | None = None - if raw_user_id: - user_id_str = str(raw_user_id) - # Fast path — rebuild from the signed session cookie. - user_ctx = UserContext.from_session_dict(session.get(SESSION_USER_CTX_KEY)) - if user_ctx is None or user_ctx.id != user_id_str: - try: - user_uuid = uuid.UUID(user_id_str) - except (ValueError, TypeError): - logger.warning("Invalid user_id in session: %r", raw_user_id) - session.pop(_SESSION_USER_ID_KEY, None) - session.pop(SESSION_USER_CTX_KEY, None) - user_ctx = None - else: - user_ctx = await self._load_user(scope, user_uuid) - if user_ctx is None: - # User was deleted / disabled since session creation. - session.pop(_SESSION_USER_ID_KEY, None) - session.pop(SESSION_USER_CTX_KEY, None) - else: - session[SESSION_USER_CTX_KEY] = user_ctx.to_session_dict() - - # Fall-through: registered principal resolvers (PAT, API key, ...). - # The session-cookie path above is authoritative; resolvers only run - # when no session-authenticated user was resolved. - if user_ctx is None: - auth_state = getattr(scope["app"].state, "auth", None) - resolvers = getattr(auth_state, "principal_resolvers", ()) if auth_state else () - if resolvers: - request = Request(scope) - for resolver in resolvers: - try: - user_ctx = await resolver(request) - except Exception: - logger.exception( - "Principal resolver %r raised; treating as no-match", - resolver, - ) - continue - if user_ctx is not None: - break - - if user_ctx is None and not is_public: - if path.startswith(_API_PATH_PREFIX): - response = JSONResponse({"detail": _UNAUTH_DETAIL}, status_code=401) - else: - request = Request(scope) - session[_SESSION_NEXT_KEY] = str(request.url) - response = RedirectResponse(_LOGIN_REDIRECT, status_code=302) - await response(scope, receive, send) - return - - if user_ctx is not None: - request = Request(scope) - request.state.user = user_ctx - token = current_user_id.set(user_ctx.id) - try: - await self.app(scope, receive, send) - finally: - current_user_id.reset(token) - return - - await self.app(scope, receive, send) - - async def _load_user(self, scope: Scope, user_id: uuid.UUID) -> UserContext | None: - """Open a fresh session from app.state.db and load the User + roles. +from auth.middleware import AuthMiddleware - Returns a UserContext, or None if the user doesn't exist or is - disabled/inactive. The session is closed on exit; we never commit - (read-only). - """ - try: - session_factory = scope["app"].state.sm.db.session_factory - async with session_factory() as db_session: - stmt = select(User).where(User.id == user_id).options(selectinload(User.roles)) - user = (await db_session.execute(stmt)).scalar_one_or_none() - if user is None: - return None - if not user.is_active or user.disabled_at is not None: - return None - return UserContext.from_user(user) - except Exception: - logger.exception("Failed to load user %s from DB; treating as unauthenticated", user_id) - return None +__all__ = ["AuthMiddleware"] diff --git a/modules/users/users/models/__init__.py b/modules/users/users/models/__init__.py index bbfdd733..d92f31f9 100644 --- a/modules/users/users/models/__init__.py +++ b/modules/users/users/models/__init__.py @@ -10,6 +10,7 @@ from users.models._base import Base from users.models.access_token import UserAccessToken from users.models.oauth_account import OAuthAccount +from users.models.refresh_token import RefreshToken from users.models.role import Role from users.models.user import User from users.models.user_role import UserRole @@ -17,6 +18,7 @@ __all__ = [ "Base", "OAuthAccount", + "RefreshToken", "Role", "SQLAlchemyAccessTokenDatabase", "SQLAlchemyUserDatabase", diff --git a/modules/users/users/models/refresh_token.py b/modules/users/users/models/refresh_token.py new file mode 100644 index 00000000..3380464d --- /dev/null +++ b/modules/users/users/models/refresh_token.py @@ -0,0 +1,22 @@ +"""Refresh token for mobile/API bearer auth.""" + +from __future__ import annotations + +import uuid as uuid_mod +from datetime import UTC, datetime + +from sqlmodel import Field + +from users.models._base import Base + + +class RefreshToken(Base, table=True): # ty: ignore[unsupported-base] + """Opaque refresh token exchanged for a new access + refresh pair.""" + + __tablename__ = "users_refresh_token" + + token: uuid_mod.UUID = Field(default_factory=uuid_mod.uuid4, primary_key=True) + user_id: uuid_mod.UUID = Field(foreign_key="users_user.id", index=True) + created_at: datetime = Field(default_factory=lambda: datetime.now(UTC)) + expires_at: datetime + revoked_at: datetime | None = None diff --git a/modules/users/users/module.py b/modules/users/users/module.py index 2abc7404..e9049e59 100644 --- a/modules/users/users/module.py +++ b/modules/users/users/module.py @@ -39,12 +39,11 @@ class UsersModule(ModuleBase): view_prefix="/users", depends_on=[_MODULE_DEPENDENCY_AUTH], ) + _is_auth_provider = True def register_settings(self, app: FastAPI) -> None: import importlib - from auth.contracts.schemas import UserContext - from users.settings import UsersSettings from users.state import UsersState @@ -58,15 +57,9 @@ def register_settings(self, app: FastAPI) -> None: register_module_settings(app, "users", UsersSettings, lambda s: UsersState(settings=s)) - def serialize_principal(user: UserContext) -> dict: - return { - "id": user.id, - "name": user.name, - "email": user.email, - "roles": user.roles, - } + from users.provider import UsersAuthProvider - app.state.principal_serializer = serialize_principal + app.state.auth.auth_provider = UsersAuthProvider() def register_permissions(self, registry: PermissionRegistry) -> None: registry.add_group( @@ -113,6 +106,7 @@ def register_routes(self, api_router: APIRouter, view_router: APIRouter) -> None from users.admin.api import admin_router from users.admin.views import router as admin_views from users.auth_local import api as auth_local_api + from users.auth_local.token_api import router as token_router from users.auth_local.views import router as auth_views from users.contracts.schemas import UserCreate, UserRead from users.deps import fastapi_users @@ -127,6 +121,7 @@ def register_routes(self, api_router: APIRouter, view_router: APIRouter) -> None settings = UsersSettings() api_router.include_router(auth_local_api.router) + api_router.include_router(token_router) api_router.include_router(admin_router) # Throughput-wrap the stock fastapi-users routers; ``require_signup_enabled`` # gates /register at request time so ``allow_signup`` is hot-reloadable. @@ -156,11 +151,6 @@ def register_routes(self, api_router: APIRouter, view_router: APIRouter) -> None view_router.include_router(auth_views) view_router.include_router(admin_views) - def register_middleware(self, app: FastAPI) -> None: - from users.middleware import AuthMiddleware - - app.add_middleware(AuthMiddleware) - async def on_startup(self, app: FastAPI) -> None: """Build the mailer, rate limiter, and apply production cookie params.""" import asyncio diff --git a/modules/users/users/provider.py b/modules/users/users/provider.py new file mode 100644 index 00000000..e956c77c --- /dev/null +++ b/modules/users/users/provider.py @@ -0,0 +1,130 @@ +"""UsersAuthProvider — AuthProvider implementation for the users module. + +Resolves users from session cookies (browser) or the principal-resolver chain +(bearer tokens, PATs). Session handling mirrors the original AuthMiddleware +logic: fast path from ``session["user_ctx"]``, slow path via DB lookup. +""" + +from __future__ import annotations + +import logging +import uuid as uuid_mod + +from auth.contracts.schemas import UserContext +from starlette.requests import Request + +logger = logging.getLogger(__name__) + +_SESSION_USER_ID_KEY = "user_id" +_SESSION_USER_CTX_KEY = "user_ctx" + + +class UsersAuthProvider: + """Cookie + bearer auth provider using fastapi-users' DatabaseStrategy.""" + + name = "users" + _is_auth_provider = True + + async def resolve_user(self, request: Request) -> UserContext | None: + auth_header = request.headers.get("authorization", "") + if auth_header.startswith("Bearer "): + return await self._resolve_bearer(request.scope, auth_header[7:]) + + session = request.scope.get("session", {}) + raw_user_id = session.get(_SESSION_USER_ID_KEY) + if not raw_user_id: + return None + + user_id_str = str(raw_user_id) + + cached = UserContext.from_session_dict(session.get(_SESSION_USER_CTX_KEY)) + if cached is not None and cached.id == user_id_str: + return cached + + try: + user_uuid = uuid_mod.UUID(user_id_str) + except (ValueError, TypeError): + logger.warning("Invalid user_id in session: %r", raw_user_id) + session.pop(_SESSION_USER_ID_KEY, None) + session.pop(_SESSION_USER_CTX_KEY, None) + return None + + user_ctx = await self._load_user(request.scope, user_uuid) + if user_ctx is None: + session.pop(_SESSION_USER_ID_KEY, None) + session.pop(_SESSION_USER_CTX_KEY, None) + else: + session[_SESSION_USER_CTX_KEY] = user_ctx.to_session_dict() + return user_ctx + + def get_login_url(self, request: Request | None, next_url: str | None = None) -> str: + return "/users/login" + + def get_logout_url(self, request: Request | None) -> str: + return "/users/logout" + + def get_public_paths(self) -> tuple[tuple[str, ...], tuple[str, ...]]: + return ( + ( + "/users/login", + "/users/register", + "/users/forgot-password", + "/users/reset-password", + "/users/verify", + "/users/invite/accept", + "/api/users/auth/", + "/api/users/register", + ), + (), + ) + + def is_bearer_request(self, request: Request | None) -> bool: + if request is None: + return False + return request.headers.get("authorization", "").startswith("Bearer ") + + async def _resolve_bearer(self, scope, token: str) -> UserContext | None: + """Look up an access token in users_access_token and return the user.""" + try: + from sqlalchemy import select + from sqlalchemy.orm import selectinload + + from users.models import User, UserAccessToken + + session_factory = scope["app"].state.sm.db.session_factory + async with session_factory() as db_session: + stmt = select(UserAccessToken).where(UserAccessToken.token == token) + access = (await db_session.execute(stmt)).scalar_one_or_none() + if access is None: + return None + stmt = ( + select(User).where(User.id == access.user_id).options(selectinload(User.roles)) + ) + user = (await db_session.execute(stmt)).scalar_one_or_none() + if user is None or not user.is_active or user.disabled_at is not None: + return None + return UserContext.from_user(user) + except Exception: + logger.exception("Bearer token resolution failed") + return None + + async def _load_user(self, scope, user_id: uuid_mod.UUID) -> UserContext | None: + try: + from sqlalchemy import select + from sqlalchemy.orm import selectinload + + from users.models import User + + session_factory = scope["app"].state.sm.db.session_factory + async with session_factory() as db_session: + stmt = select(User).where(User.id == user_id).options(selectinload(User.roles)) + user = (await db_session.execute(stmt)).scalar_one_or_none() + if user is None or not user.is_active or user.disabled_at is not None: + return None + return UserContext.from_user(user) + except Exception: + logger.exception( + "Failed to load user %s from DB; treating as unauthenticated", + user_id, + ) + return None diff --git a/modules/users/users/settings.py b/modules/users/users/settings.py index bc644d27..5d955fd8 100644 --- a/modules/users/users/settings.py +++ b/modules/users/users/settings.py @@ -50,6 +50,10 @@ class UsersSettings(BaseSettings): reset_password_token_lifetime_seconds: int = 60 * 60 # 1 hour verification_token_lifetime_seconds: int = 60 * 60 * 24 * 7 # 7 days + # Bearer token (mobile / API clients) + bearer_token_lifetime_seconds: int = 60 * 15 # 15 minutes + refresh_token_lifetime_seconds: int = 60 * 60 * 24 * 30 # 30 days + # Cookie (fastapi-users AuthenticationBackend) cookie_name: str = "sm_auth" cookie_max_age_seconds: int = 60 * 60 * 24 * 14 # 14 days diff --git a/package-lock.json b/package-lock.json index 8a580ea7..2d546a1a 100644 --- a/package-lock.json +++ b/package-lock.json @@ -188,6 +188,10 @@ "react-dom": "^19.0.0" } }, + "modules/keycloak": { + "name": "@simple-module/keycloak", + "version": "0.0.0" + }, "modules/permissions": { "name": "@simple-module-py/permissions", "version": "0.1.0", @@ -4545,6 +4549,10 @@ "resolved": "modules/users", "link": true }, + "node_modules/@simple-module/keycloak": { + "resolved": "modules/keycloak", + "link": true + }, "node_modules/@standard-schema/spec": { "version": "1.1.0", "resolved": "https://registry.npmjs.org/@standard-schema/spec/-/spec-1.1.0.tgz", diff --git a/pyproject.toml b/pyproject.toml index d84326d5..ceb8a21b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -27,6 +27,7 @@ dev = [ "aioboto3>=13", "moto[s3]>=5", "tomlkit>=0.13", + "pytest-httpx>=0.36.2", ] [tool.ruff] @@ -71,6 +72,7 @@ extra-paths = [ "modules/file_storage", "modules/settings", "modules/feature_flags", + "modules/keycloak", "host", "scripts", ] @@ -100,10 +102,13 @@ no-matching-overload = "ignore" # The false positives massively outnumber the real bugs and the latter still # get caught at runtime by tests, so we ignore the rule globally. invalid-argument-type = "ignore" +# SQLModel's ``model_config = ConfigDict(...)`` clashes with the internal +# ``SQLModelConfig`` type in newer ty versions. Same root cause as above. +invalid-assignment = "ignore" [tool.pytest.ini_options] asyncio_mode = "auto" -testpaths = ["framework/cli/tests", "framework/core/tests", "framework/db/tests", "framework/hosting/tests", "framework/testing/tests", "host/tests", "modules/auth/tests", "modules/dashboard/tests", "modules/users/tests", "modules/permissions/tests", "modules/background_tasks/tests", "modules/file_storage/tests", "modules/settings/tests", "modules/feature_flags/tests", "scripts/tests", "tests/integration", "tests/e2e", "tests/benchmarks"] +testpaths = ["framework/cli/tests", "framework/core/tests", "framework/db/tests", "framework/hosting/tests", "framework/testing/tests", "host/tests", "modules/auth/tests", "modules/dashboard/tests", "modules/users/tests", "modules/permissions/tests", "modules/background_tasks/tests", "modules/file_storage/tests", "modules/settings/tests", "modules/feature_flags/tests", "modules/keycloak/tests", "scripts/tests", "tests/integration", "tests/e2e", "tests/benchmarks"] markers = [ "e2e: end-to-end tests requiring a live browser", "perf: performance benchmarks (opt-in; run via `make bench`)", diff --git a/tests/integration/test_pluggable_auth.py b/tests/integration/test_pluggable_auth.py new file mode 100644 index 00000000..834ae83e --- /dev/null +++ b/tests/integration/test_pluggable_auth.py @@ -0,0 +1,58 @@ +"""Integration tests for pluggable auth — verifying both providers work.""" + +from __future__ import annotations + +from auth.contracts.provider import AuthProvider + + +def test_users_module_is_auth_provider(): + from users.module import UsersModule + + assert UsersModule._is_auth_provider is True + + +def test_keycloak_module_is_auth_provider(): + from keycloak.module import KeycloakModule + + assert KeycloakModule._is_auth_provider is True + + +def test_sm020_fires_with_both_modules(): + from keycloak.module import KeycloakModule + from simple_module_core.diagnostics._module import ModuleDiagnostics + from users.module import UsersModule + + diags = ModuleDiagnostics() + results = diags._check_auth_provider_conflict([UsersModule(), KeycloakModule()]) + assert any(d.code == "SM020" for d in results) + + +def test_sm021_fires_with_neither(): + from simple_module_core.diagnostics._module import ModuleDiagnostics + from simple_module_core.module import ModuleBase, ModuleMeta + + class StubModule(ModuleBase): + meta = ModuleMeta(name="Stub") + + diags = ModuleDiagnostics() + results = diags._check_auth_provider_conflict([StubModule()]) + assert any(d.code == "SM021" for d in results) + + +def test_auth_provider_protocol_satisfied_by_users(): + from users.provider import UsersAuthProvider + + assert isinstance(UsersAuthProvider(), AuthProvider) + + +def test_auth_provider_protocol_satisfied_by_keycloak(): + from keycloak.provider import KeycloakAuthProvider + from keycloak.settings import KeycloakSettings + + settings = KeycloakSettings( + server_url="https://example.com", + realm="test", + client_id="app", + client_secret="secret", + ) + assert isinstance(KeycloakAuthProvider(settings), AuthProvider)