diff --git a/conftest.py b/conftest.py index 03e6a651..7294cbb8 100644 --- a/conftest.py +++ b/conftest.py @@ -12,7 +12,9 @@ from simple_module_core.discovery import discover_modules from simple_module_db.base import all_module_bases from simple_module_db.session import DatabaseState, init_db +from simple_module_hosting.csrf import SESSION_CSRF_TOKEN_KEY from simple_module_hosting.settings import Settings +from simple_module_testing import forge_session_cookie from sqlalchemy.ext.asyncio import ( AsyncEngine, AsyncSession, @@ -135,21 +137,30 @@ async def app(settings: Settings): await ctx.__aexit__(None, None, None) +_TEST_CSRF_TOKEN = "test-csrf-token" + + @pytest.fixture async def client(app) -> AsyncGenerator[httpx.AsyncClient, None]: - """Unauthenticated async HTTP client.""" + """Unauthenticated async HTTP client — still CSRF-equipped so POST/PATCH/DELETE + requests from anonymous test flows (login, accept-invite, register) pass + validation without first making a GET to mint a token.""" + signed = forge_session_cookie( + app.state.settings.secret_key, {SESSION_CSRF_TOKEN_KEY: _TEST_CSRF_TOKEN} + ) transport = httpx.ASGITransport(app=app) - async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as c: + async with httpx.AsyncClient( + transport=transport, + base_url="http://testserver", + cookies={"session": signed}, + headers={"X-CSRF-Token": _TEST_CSRF_TOKEN}, + ) as c: yield c @pytest.fixture async def authenticated_client(app) -> AsyncGenerator[httpx.AsyncClient, None]: """HTTPX client with a signed session cookie carrying a seeded admin user's id.""" - import json - from base64 import b64encode - - from itsdangerous import TimestampSigner from users.bootstrap import create_admin async with app.state.db.session_factory() as session: @@ -161,15 +172,16 @@ async def authenticated_client(app) -> AsyncGenerator[httpx.AsyncClient, None]: ) user_id = str(result.user.id) - session_data = {"user_id": user_id} - data = b64encode(json.dumps(session_data).encode()) - signer = TimestampSigner(str(app.state.settings.secret_key)) - signed = signer.sign(data).decode("utf-8") + signed = forge_session_cookie( + app.state.settings.secret_key, + {"user_id": user_id, SESSION_CSRF_TOKEN_KEY: _TEST_CSRF_TOKEN}, + ) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( transport=transport, base_url="http://testserver", cookies={"session": signed}, + headers={"X-CSRF-Token": _TEST_CSRF_TOKEN}, ) as c: yield c diff --git a/framework/core/simple_module_core/environments.py b/framework/core/simple_module_core/environments.py new file mode 100644 index 00000000..3bcb53fb --- /dev/null +++ b/framework/core/simple_module_core/environments.py @@ -0,0 +1,15 @@ +"""Shared environment-classification constants. + +Both the host and module settings validators need to know which ``SM_ENVIRONMENT`` +values are "non-prod" — anything outside this set is treated as production +and subject to stricter defaults (e.g. placeholder-secret rejection). + +Duplicating this constant in each settings module would mean an operator who +adds ``"staging"`` to one list but forgets the other gets inconsistent +validation. Lives in ``simple_module_core`` because both the host package and +module settings already depend on it. +""" + +from __future__ import annotations + +NON_PROD_ENVIRONMENTS: frozenset[str] = frozenset({"development", "testing"}) diff --git a/framework/core/simple_module_core/i18n.py b/framework/core/simple_module_core/i18n.py index 1faf15fa..9c8162bf 100644 --- a/framework/core/simple_module_core/i18n.py +++ b/framework/core/simple_module_core/i18n.py @@ -4,8 +4,10 @@ import json import logging +from collections.abc import Mapping from functools import lru_cache from pathlib import Path +from types import MappingProxyType from typing import Any from babel import Locale @@ -78,6 +80,13 @@ def __init__(self, default_locale: str, supported_locales: list[str]) -> None: self.supported_locales = list(supported_locales) self._sources: list[tuple[str, Path]] = [] self._messages: dict[str, dict[str, str]] = {} + # Immutable views into ``_messages`` — handed out by ``messages()`` to + # avoid a per-call dict copy. Rebuilt whenever ``load()`` runs. + self._message_views: dict[str, MappingProxyType[str, str]] = {} + self._available_locales: tuple[str, ...] = () + self._available_locales_list: list[str] = [] + self._empty_view: MappingProxyType[str, str] = MappingProxyType({}) + self._loaded = False def add_source(self, namespace: str, locale_dir: Path) -> None: """Queue a module's locale directory for loading under a namespace.""" @@ -110,13 +119,43 @@ def load(self) -> None: flat = flatten_messages(raw, prefix=namespace) self._messages[locale].update(flat) + # Cache the derived views now that loading is complete. Downstream + # (middleware, translator, switcher) reads these on every request. + self._message_views = { + locale: MappingProxyType(msgs) for locale, msgs in self._messages.items() + } + self._available_locales = tuple(locale for locale, msgs in self._messages.items() if msgs) + self._available_locales_list = list(self._available_locales) + self._loaded = True + def available_locales(self) -> list[str]: - """Locales that have at least one loaded message.""" + """Locales that have at least one loaded message. + + The list is cached at ``load()`` time; if ``load()`` hasn't run but + tests populated ``_messages`` directly, a one-off scan returns the + derived list without caching it (the test is outside the normal flow). + """ + if self._loaded: + return self._available_locales_list return [locale for locale, msgs in self._messages.items() if msgs] - def messages(self, locale: str) -> dict[str, str]: - """Flat dotted-key map for the given locale. Empty dict if unknown.""" - return dict(self._messages.get(locale, {})) + def messages(self, locale: str) -> Mapping[str, str]: + """Flat dotted-key map for the given locale. Empty mapping if unknown. + + Returns an immutable view (``MappingProxyType``) into the cached + message dict — zero-copy. Callers that JSON-serialize the result + (e.g. Inertia shared props) should wrap with ``dict(...)`` at the + boundary. + """ + view = self._message_views.get(locale) + if view is not None: + return view + # Fallback: ``load()`` wasn't called (tests may populate _messages + # directly). Expose the raw dict as a proxy so Translator still works. + raw = self._messages.get(locale) + if raw is None: + return self._empty_view + return MappingProxyType(raw) class _SafeFormatDict(dict): diff --git a/framework/core/simple_module_core/menu.py b/framework/core/simple_module_core/menu.py index 39b88d9b..f4b89806 100644 --- a/framework/core/simple_module_core/menu.py +++ b/framework/core/simple_module_core/menu.py @@ -4,6 +4,9 @@ from dataclasses import dataclass, field from enum import StrEnum +from typing import Literal + +MenuItemMethod = Literal["get", "post"] class MenuSection(StrEnum): @@ -27,6 +30,9 @@ class MenuItem: requires_auth: bool = True roles: list[str] = field(default_factory=list) """Empty list = visible to all authenticated users.""" + method: MenuItemMethod = "get" + """HTTP method used when the item is activated. ``"post"`` renders as an + Inertia form submission so the target endpoint can be POST-only (e.g. logout).""" class MenuRegistry: @@ -34,16 +40,24 @@ class MenuRegistry: def __init__(self) -> None: self._items: list[MenuItem] = [] + self._sorted: list[MenuItem] | None = None + + def _invalidate(self) -> None: + self._sorted = None def add(self, item: MenuItem) -> None: self._items.append(item) + self._invalidate() def add_many(self, items: list[MenuItem]) -> None: self._items.extend(items) + self._invalidate() @property def all_items(self) -> list[MenuItem]: - return sorted(self._items, key=lambda i: i.order) + if self._sorted is None: + self._sorted = sorted(self._items, key=lambda i: i.order) + return self._sorted def get_for_user( self, @@ -68,6 +82,7 @@ def get_for_user( "label": item.label, "url": item.url, "icon": item.icon, + "method": item.method, } ) diff --git a/framework/db/simple_module_db/session.py b/framework/db/simple_module_db/session.py index 9aafd500..a3f0c3df 100644 --- a/framework/db/simple_module_db/session.py +++ b/framework/db/simple_module_db/session.py @@ -25,22 +25,38 @@ class DatabaseState: _listeners_registered: bool = field(default=False, repr=False) -def init_db(database_url: str, *, echo: bool = False) -> DatabaseState: +def init_db( + database_url: str, + *, + echo: bool = False, + pool_size: int = 10, + max_overflow: int = 20, + pool_pre_ping: bool = True, + pool_recycle: int = 1800, +) -> DatabaseState: """Create an async engine and session factory. + The pool options only take effect for server-side providers (Postgres). + SQLite uses SQLAlchemy's default pool (single-file, no network), so + passing ``pool_size``/etc. would raise ``TypeError`` — skipped below. + Returns a ``DatabaseState`` that should be stored on ``app.state.db``. """ provider = detect_provider(database_url) connect_args: dict = {} + engine_kwargs: dict = {"echo": echo, "connect_args": connect_args} if provider == DatabaseProvider.SQLITE: connect_args["check_same_thread"] = False + else: + engine_kwargs.update( + pool_size=pool_size, + max_overflow=max_overflow, + pool_pre_ping=pool_pre_ping, + pool_recycle=pool_recycle, + ) - engine = create_async_engine( - database_url, - echo=echo, - connect_args=connect_args, - ) + engine = create_async_engine(database_url, **engine_kwargs) # Scoped Session subclass so event listeners only fire for this engine's sessions scoped_session_class = type("ScopedSession", (Session,), {}) session_factory = async_sessionmaker( diff --git a/framework/db/tests/test_session.py b/framework/db/tests/test_session.py index c1b01930..a85cf112 100644 --- a/framework/db/tests/test_session.py +++ b/framework/db/tests/test_session.py @@ -32,6 +32,20 @@ async def test_separate_init_db_calls_are_independent(self): await db1.engine.dispose() await db2.engine.dispose() + async def test_sqlite_ignores_pool_kwargs(self): + """SQLite must not receive pool_size/max_overflow — they would raise TypeError.""" + db_state = init_db( + "sqlite+aiosqlite:///:memory:", + pool_size=50, + max_overflow=100, + pool_pre_ping=True, + pool_recycle=60, + ) + try: + assert db_state.engine is not None + finally: + await db_state.engine.dispose() + class TestGetDbDependency: async def test_get_db_yields_session(self): diff --git a/framework/hosting/simple_module_hosting/_inertia_shared.py b/framework/hosting/simple_module_hosting/_inertia_shared.py new file mode 100644 index 00000000..fb90abb0 --- /dev/null +++ b/framework/hosting/simple_module_hosting/_inertia_shared.py @@ -0,0 +1,58 @@ +"""Helpers for building Inertia shared-props payloads.""" + +from __future__ import annotations + +import logging + +from starlette.datastructures import Headers +from starlette.requests import Request +from starlette.types import Scope + +logger = logging.getLogger(__name__) + +_I18N_SESSION_LOCALE_KEY = "__i18n_locale" + + +def build_i18n_block(scope: Scope, request: Request) -> dict: + """Assemble the ``i18n`` shared-props block for the current request. + + Rules: + + * No registry / no locale → serve an empty English block and log once. + * Inertia XHR partials (``X-Inertia: true``) reuse the client-side + cached messages; send ``messages: None`` unless the locale differs + from what was last served on this session. + * Full page loads and locale transitions ship the complete dict. + """ + registry = getattr(request.app.state, "i18n_registry", None) + locale = getattr(request.state, "locale", None) + if registry is None or locale is None: + logger.warning( + "InertiaLayoutDataMiddleware: i18n not fully wired " + "(registry_present=%s, locale_present=%s); serving empty messages", + registry is not None, + locale is not None, + ) + return {"locale": "en", "supportedLocales": ["en"], "messages": {}} + + is_inertia = Headers(scope=scope).get("x-inertia") == "true" + session_dict = scope.get("session") + # When the session is absent (pre-session-middleware routes, WebSocket + # upgrades), treat locale as "unchanged" so Inertia XHR requests still + # skip the messages payload. Non-Inertia requests will always ship them + # regardless of the session state. + if session_dict is not None: + last_locale = session_dict.get(_I18N_SESSION_LOCALE_KEY) + locale_changed = last_locale != locale + if locale_changed: + session_dict[_I18N_SESSION_LOCALE_KEY] = locale + else: + locale_changed = False + send_messages = (not is_inertia) or locale_changed + return { + "locale": locale, + "supportedLocales": registry.available_locales(), + # Inertia's JSON encoder doesn't accept the MappingProxyType view + # returned by messages(); copy to a plain dict at the boundary. + "messages": dict(registry.messages(locale)) if send_messages else None, + } diff --git a/framework/hosting/simple_module_hosting/_observability.py b/framework/hosting/simple_module_hosting/_observability.py new file mode 100644 index 00000000..66ba9d6a --- /dev/null +++ b/framework/hosting/simple_module_hosting/_observability.py @@ -0,0 +1,102 @@ +"""ASGI middlewares for correlation IDs and structured request logging.""" + +from __future__ import annotations + +import logging +import time +import uuid + +from starlette.datastructures import Headers, MutableHeaders +from starlette.types import ASGIApp, Message, Receive, Scope, Send + +from simple_module_hosting.logging import correlation_id + +_request_logger = logging.getLogger("simple_module.request") + +# Paths that produce noisy, low-value log entries +_QUIET_PREFIXES = ("/health", "/static/") + + +class CorrelationIdMiddleware: + """Generate or propagate a correlation ID for every request. + + Reads the incoming ``X-Correlation-ID`` header (or generates a UUID4) and + stores it in a :class:`~contextvars.ContextVar` so that every log record + emitted during the request automatically includes the ID. The same value + is echoed back in the response header. + """ + + HEADER = "X-Correlation-ID" + + 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 + + cid = Headers(scope=scope).get(self.HEADER) or uuid.uuid4().hex + + async def send_with_header(message: Message) -> None: + if message["type"] == "http.response.start": + headers = MutableHeaders(scope=message) + headers[self.HEADER] = cid + await send(message) + + token = correlation_id.set(cid) + try: + await self.app(scope, receive, send_with_header) + finally: + correlation_id.reset(token) + + +class RequestLoggingMiddleware: + """Log every request/response pair with timing and status information.""" + + 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"] + if any(path.startswith(p) for p in _QUIET_PREFIXES): + await self.app(scope, receive, send) + return + + method = scope["method"] + client = scope.get("client") + client_ip = client[0] if client else "unknown" + + _request_logger.debug( + "request.started", + extra={"method": method, "path": path, "client_ip": client_ip}, + ) + + status_code: int | None = None + start = time.perf_counter() + + async def send_capture(message: Message) -> None: + nonlocal status_code + if message["type"] == "http.response.start": + status_code = message["status"] + await send(message) + + try: + await self.app(scope, receive, send_capture) + finally: + # Log completion even when the inner app raises, so 500s are observable. + duration_ms = round((time.perf_counter() - start) * 1000, 2) + _request_logger.info( + "request.completed", + extra={ + "method": method, + "path": path, + "status_code": status_code, + "duration_ms": duration_ms, + "client_ip": client_ip, + }, + ) diff --git a/framework/hosting/simple_module_hosting/_phase_helpers.py b/framework/hosting/simple_module_hosting/_phase_helpers.py index b3c67545..4a61d401 100644 --- a/framework/hosting/simple_module_hosting/_phase_helpers.py +++ b/framework/hosting/simple_module_hosting/_phase_helpers.py @@ -27,6 +27,7 @@ not_found_error_handler, unhandled_exception_handler, ) +from simple_module_hosting.csrf import CSRFMiddleware from simple_module_hosting.i18n_middleware import LocaleMiddleware from simple_module_hosting.middleware import ( CorrelationIdMiddleware, @@ -67,7 +68,7 @@ def install_middleware( """Install the full middleware pipeline. Order matters: last added = first executed. Execution order: - CorrelationId → RequestLogging → Security → Session + CorrelationId → RequestLogging → Security → Session → CSRF → [module] → (Tenant, if multi_tenant) → Locale → Inertia. """ app.add_middleware( @@ -85,6 +86,9 @@ def install_middleware( app.add_middleware(TenantMiddleware, header=settings.tenant_header or None) for mod in modules: mod.register_middleware(app) + # CSRF runs immediately after SessionMiddleware loads the session so that + # scope["session"] is populated by the time we validate the token. + app.add_middleware(CSRFMiddleware) app.add_middleware(SessionMiddleware, secret_key=settings.secret_key) app.add_middleware(SecurityHeadersMiddleware) app.add_middleware(RequestLoggingMiddleware) diff --git a/framework/hosting/simple_module_hosting/app_builder.py b/framework/hosting/simple_module_hosting/app_builder.py index 3d4124d4..9dac6656 100644 --- a/framework/hosting/simple_module_hosting/app_builder.py +++ b/framework/hosting/simple_module_hosting/app_builder.py @@ -197,7 +197,14 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: ) # ── Phase 6: Initialize database ─────────────────────── - db_state = init_db(settings.database_url, echo=settings.debug) + db_state = init_db( + settings.database_url, + echo=settings.debug, + pool_size=settings.db_pool_size, + max_overflow=settings.db_max_overflow, + pool_pre_ping=settings.db_pool_pre_ping, + pool_recycle=settings.db_pool_recycle, + ) register_listeners(db_state) app.state.db = db_state diff --git a/framework/hosting/simple_module_hosting/csrf.py b/framework/hosting/simple_module_hosting/csrf.py new file mode 100644 index 00000000..a151d148 --- /dev/null +++ b/framework/hosting/simple_module_hosting/csrf.py @@ -0,0 +1,82 @@ +"""CSRF protection middleware. + +Validates a per-session token on state-changing requests. The token is minted +once per session by :class:`InertiaLayoutDataMiddleware` and embedded in the +Inertia shared props as ``csrf_token``; the frontend echoes it back in the +``X-CSRF-Token`` (or ``X-XSRF-Token``) header on every non-safe request. + +Installed *after* ``SessionMiddleware`` so ``scope["session"]`` is populated. +""" + +from __future__ import annotations + +import logging +import secrets + +from starlette.datastructures import Headers +from starlette.responses import PlainTextResponse +from starlette.types import ASGIApp, Receive, Scope, Send + +logger = logging.getLogger(__name__) + +SESSION_CSRF_TOKEN_KEY = "csrf_token" +"""Session-dict key under which :class:`InertiaLayoutDataMiddleware` mints +the per-session CSRF token that :class:`CSRFMiddleware` validates.""" + +_SAFE_METHODS = frozenset({"GET", "HEAD", "OPTIONS", "TRACE"}) +_HEADER_NAMES = ("x-csrf-token", "x-xsrf-token") + + +class CSRFMiddleware: + """Reject unsafe-method requests without a valid session-bound CSRF token. + + ``exempt_path_prefixes`` opts specific path prefixes out of the check — + use sparingly, only for endpoints that genuinely cannot carry a session + (e.g. webhook receivers authenticated by signature). + """ + + def __init__( + self, + app: ASGIApp, + *, + exempt_path_prefixes: tuple[str, ...] = (), + ) -> None: + self.app = app + self.exempt_path_prefixes = exempt_path_prefixes + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] != "http": + await self.app(scope, receive, send) + return + + method = scope["method"] + if method in _SAFE_METHODS: + await self.app(scope, receive, send) + return + + path = scope["path"] + if any(path.startswith(p) for p in self.exempt_path_prefixes): + await self.app(scope, receive, send) + return + + session = scope.get("session") + expected = session.get(SESSION_CSRF_TOKEN_KEY) if session is not None else None + + headers = Headers(scope=scope) + provided = "" + for name in _HEADER_NAMES: + value = headers.get(name) + if value: + provided = value + break + + if not expected or not provided or not secrets.compare_digest(str(expected), provided): + logger.info( + "csrf.rejected", + extra={"method": method, "path": path, "has_token": bool(expected)}, + ) + response = PlainTextResponse("CSRF validation failed", status_code=403) + await response(scope, receive, send) + return + + await self.app(scope, receive, send) diff --git a/framework/hosting/simple_module_hosting/inertia_utils.py b/framework/hosting/simple_module_hosting/inertia_utils.py index 03648d10..4f8a8881 100644 --- a/framework/hosting/simple_module_hosting/inertia_utils.py +++ b/framework/hosting/simple_module_hosting/inertia_utils.py @@ -6,6 +6,8 @@ from pydantic import ValidationError from starlette.responses import RedirectResponse +from simple_module_hosting.redirects import safe_referer_or_root + SESSION_ERRORS_KEY = "_errors" @@ -20,7 +22,10 @@ def validation_errors_to_dict(exc: ValidationError) -> dict[str, str]: def redirect_back_with_errors(request: Request, errors: dict[str, str]) -> RedirectResponse: - """Store validation errors in the session and redirect to the referring page.""" + """Store validation errors in the session and redirect to the referring page. + + Uses ``safe_referer_or_root`` to reject attacker-controlled Referer values + (cross-origin URLs) — otherwise this becomes a reflected open redirect + accessible to any attacker who can trigger a form validation error.""" request.session[SESSION_ERRORS_KEY] = errors - referer = request.headers.get("referer", "/") - return RedirectResponse(referer, status_code=303) + return RedirectResponse(safe_referer_or_root(request), status_code=303) diff --git a/framework/hosting/simple_module_hosting/middleware.py b/framework/hosting/simple_module_hosting/middleware.py index 79348ec9..2a566923 100644 --- a/framework/hosting/simple_module_hosting/middleware.py +++ b/framework/hosting/simple_module_hosting/middleware.py @@ -1,4 +1,6 @@ -"""Middleware: security headers, tenant isolation, correlation IDs, request logging, layout data. +"""Middleware: security headers, tenant isolation, Inertia shared-props. + +Correlation IDs and request logging live in :mod:`._observability`. All middleware classes use the raw ASGI pattern instead of ``BaseHTTPMiddleware`` to avoid its known issues with streaming responses, extra task creation, @@ -9,8 +11,6 @@ import logging import secrets -import time -import uuid from typing import TYPE_CHECKING from simple_module_db import current_tenant_id @@ -18,110 +18,66 @@ from starlette.requests import Request from starlette.types import ASGIApp, Message, Receive, Scope, Send -from simple_module_hosting.logging import correlation_id +from simple_module_hosting._inertia_shared import build_i18n_block +from simple_module_hosting._observability import ( + CorrelationIdMiddleware, + RequestLoggingMiddleware, +) +from simple_module_hosting.csrf import SESSION_CSRF_TOKEN_KEY from simple_module_hosting.permissions import expand_permissions, resolve_permissions if TYPE_CHECKING: from simple_module_core.menu import MenuRegistry from simple_module_core.permissions import PermissionRegistry -_request_logger = logging.getLogger("simple_module.request") logger = logging.getLogger(__name__) -# Paths that produce noisy, low-value log entries -_QUIET_PREFIXES = ("/health", "/static/") +__all__ = [ + "TENANT_HEADER", + "CorrelationIdMiddleware", + "InertiaLayoutDataMiddleware", + "RequestLoggingMiddleware", + "SecurityHeadersMiddleware", + "TenantMiddleware", +] -class CorrelationIdMiddleware: - """Generate or propagate a correlation ID for every request. +class SecurityHeadersMiddleware: + """Add security headers to every response. - Reads the incoming ``X-Correlation-ID`` header (or generates a UUID4) and - stores it in a :class:`~contextvars.ContextVar` so that every log record - emitted during the request automatically includes the ID. The same value - is echoed back in the response header. + ``content_security_policy`` and ``strict_transport_security`` accept a + string to override the defaults, or ``None`` to suppress that header + (useful in development when Vite's HMR client loads cross-origin scripts, + or behind plain-HTTP loopbacks where HSTS would lock users out). """ - HEADER = "X-Correlation-ID" - - 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 - - cid = Headers(scope=scope).get(self.HEADER) or uuid.uuid4().hex - - async def send_with_header(message: Message) -> None: - if message["type"] == "http.response.start": - headers = MutableHeaders(scope=message) - headers[self.HEADER] = cid - await send(message) - - token = correlation_id.set(cid) - try: - await self.app(scope, receive, send_with_header) - finally: - correlation_id.reset(token) - - -class RequestLoggingMiddleware: - """Log every request/response pair with timing and status information.""" + _DEFAULT_CSP = ( + "default-src 'self'; " + # Inertia embeds the initial page blob inline; Vite injects a React + # Refresh shim at boot. Both require 'unsafe-inline' for scripts. + # Production builds compile to hashed bundles, so this can be + # tightened with a nonce once Vite's preamble is removed in prod. + "script-src 'self' 'unsafe-inline' 'unsafe-eval'; " + "style-src 'self' 'unsafe-inline' https://fonts.googleapis.com; " + "img-src 'self' data: blob:; " + "font-src 'self' https://fonts.gstatic.com data:; " + "connect-src 'self'; " + "frame-ancestors 'self'; " + "base-uri 'self'; " + "form-action 'self'" + ) + _DEFAULT_HSTS = "max-age=31536000; includeSubDomains" - 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"] - if any(path.startswith(p) for p in _QUIET_PREFIXES): - await self.app(scope, receive, send) - return - - method = scope["method"] - client = scope.get("client") - client_ip = client[0] if client else "unknown" - - _request_logger.debug( - "request.started", - extra={"method": method, "path": path, "client_ip": client_ip}, - ) - - status_code: int | None = None - start = time.perf_counter() - - async def send_capture(message: Message) -> None: - nonlocal status_code - if message["type"] == "http.response.start": - status_code = message["status"] - await send(message) - - try: - await self.app(scope, receive, send_capture) - finally: - # Log completion even when the inner app raises, so 500s are observable. - duration_ms = round((time.perf_counter() - start) * 1000, 2) - _request_logger.info( - "request.completed", - extra={ - "method": method, - "path": path, - "status_code": status_code, - "duration_ms": duration_ms, - "client_ip": client_ip, - }, - ) - - -class SecurityHeadersMiddleware: - """Add security headers to every response.""" - - def __init__(self, app: ASGIApp) -> None: + def __init__( + self, + app: ASGIApp, + *, + content_security_policy: str | None = _DEFAULT_CSP, + strict_transport_security: str | None = _DEFAULT_HSTS, + ) -> None: self.app = app + self.csp = content_security_policy + self.hsts = strict_transport_security async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: if scope["type"] != "http": @@ -135,6 +91,10 @@ async def send_with_headers(message: Message) -> None: headers["X-Frame-Options"] = "SAMEORIGIN" headers["X-XSS-Protection"] = "1; mode=block" headers["Referrer-Policy"] = "strict-origin-when-cross-origin" + if self.csp: + headers["Content-Security-Policy"] = self.csp + if self.hsts: + headers["Strict-Transport-Security"] = self.hsts await send(message) await self.app(scope, receive, send_with_headers) @@ -234,29 +194,19 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: all_perms = self.permission_registry.all_permissions frontend_permissions = expand_permissions(resolved, all_perms) if is_authenticated else [] - registry = getattr(request.app.state, "i18n_registry", None) - locale = getattr(request.state, "locale", None) - if registry is not None and locale is not None: - # Use available_locales() (locales with actual loaded messages) — - # NOT the configured supported_locales list. Offering a locale that - # has no JSON files would render a mostly-empty UI when selected. - i18n_block = { - "locale": locale, - "supportedLocales": registry.available_locales(), - "messages": registry.messages(locale), - } - else: - logger.warning( - "InertiaLayoutDataMiddleware: i18n not fully wired " - "(registry_present=%s, locale_present=%s); serving empty messages", - registry is not None, - locale is not None, - ) - i18n_block = { - "locale": "en", - "supportedLocales": ["en"], - "messages": {}, - } + i18n_block = build_i18n_block(scope, request) + + # Session-scoped CSRF token — minted once per session, reused until + # the session is cleared (logout, session-fixation rotate, etc.). + # Embedded in shared props so the frontend can echo it back in the + # ``X-CSRF-Token`` header on every unsafe request (see CSRFMiddleware). + session = scope.get("session") + csrf_token = "" + if session is not None: + csrf_token = session.get(SESSION_CSRF_TOKEN_KEY) or "" + if not csrf_token: + csrf_token = secrets.token_urlsafe(32) + session[SESSION_CSRF_TOKEN_KEY] = csrf_token shared: dict = { "auth": { @@ -277,7 +227,7 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: is_authenticated=is_authenticated, roles=roles, ), - "csrf_token": secrets.token_urlsafe(32) if is_authenticated else "", + SESSION_CSRF_TOKEN_KEY: csrf_token, "i18n": i18n_block, } request.state.inertia_shared = shared diff --git a/framework/hosting/simple_module_hosting/redirects.py b/framework/hosting/simple_module_hosting/redirects.py new file mode 100644 index 00000000..f813a0a9 --- /dev/null +++ b/framework/hosting/simple_module_hosting/redirects.py @@ -0,0 +1,45 @@ +"""Shared helpers for validating redirect targets. + +A raw ``Referer`` header is attacker-controlled — a crafted form on a third +party site can set it to any value. Any endpoint that 303s back to the +referring page must validate the URL is same-origin before trusting it, or +become a reflected open-redirect. +""" + +from __future__ import annotations + +from urllib.parse import urlsplit + +from fastapi import Request + + +def safe_referer_or_root(request: Request) -> str: + """Return the Referer iff it's same-origin; otherwise fall back to ``/``. + + Only honors references that (a) resolve to the same scheme+host as the + current request, or (b) are relative paths that don't try to escape to a + protocol-relative URL (``//evil.example``). + """ + referer = request.headers.get("referer") + if not referer: + return "/" + + # Protocol-relative URLs like "//evil.example/foo" resolve against the + # origin in browsers but leave the site — reject them. + if referer.startswith("//"): + return "/" + + parsed = urlsplit(referer) + # Relative path with no scheme+host → same-origin by construction. + if not parsed.scheme and not parsed.netloc: + return referer if referer.startswith("/") else "/" + + # Absolute URL → must match the current request's origin. + current = request.url + if parsed.scheme == current.scheme and parsed.netloc == current.netloc: + path = parsed.path or "/" + if parsed.query: + path = f"{path}?{parsed.query}" + return path + + return "/" diff --git a/framework/hosting/simple_module_hosting/settings.py b/framework/hosting/simple_module_hosting/settings.py index 51a937ff..c0ee6863 100644 --- a/framework/hosting/simple_module_hosting/settings.py +++ b/framework/hosting/simple_module_hosting/settings.py @@ -4,6 +4,9 @@ from pydantic import model_validator from pydantic_settings import BaseSettings, SettingsConfigDict +from simple_module_core.environments import NON_PROD_ENVIRONMENTS + +_PLACEHOLDER_SECRET_KEY = "change-me-in-production" class Settings(BaseSettings): @@ -17,6 +20,12 @@ class Settings(BaseSettings): # Database database_url: str = "sqlite+aiosqlite:///./app.db" + # Connection pool — server-side providers only (Postgres); SQLite ignores these. + db_pool_size: int = 10 + db_max_overflow: int = 20 + db_pool_pre_ping: bool = True + db_pool_recycle: int = 1800 # 30 minutes — matches common proxy idle-timeouts + # App environment: str = "development" secret_key: str = "change-me-in-production" @@ -74,3 +83,21 @@ def _check_default_locale_supported(self) -> "Settings": f"i18n_supported_locales {self.i18n_supported_locales}" ) return self + + @model_validator(mode="after") + def _forbid_placeholder_secret_in_production(self) -> "Settings": + """Fail boot if production is running with the placeholder secret key. + + Session cookies and any other itsdangerous-signed payloads are signed + with ``secret_key``; a known default would let anyone forge them. + """ + if ( + self.environment not in NON_PROD_ENVIRONMENTS + and self.secret_key == _PLACEHOLDER_SECRET_KEY + ): + raise ValueError( + f"SM_SECRET_KEY must be set to a non-default value when " + f"SM_ENVIRONMENT={self.environment!r}. Generate one with " + "`python -c 'import secrets; print(secrets.token_urlsafe(48))'`." + ) + return self diff --git a/framework/hosting/tests/test_csrf_middleware.py b/framework/hosting/tests/test_csrf_middleware.py new file mode 100644 index 00000000..51646d59 --- /dev/null +++ b/framework/hosting/tests/test_csrf_middleware.py @@ -0,0 +1,119 @@ +"""Tests for CSRFMiddleware: unsafe-method rejection and header validation.""" + +from __future__ import annotations + +import pytest +from simple_module_hosting.csrf import CSRFMiddleware + + +def _scope(method: str, headers: list[tuple[bytes, bytes]] | None = None, session=None) -> dict: + return { + "type": "http", + "method": method, + "path": "/api/anything", + "headers": headers or [], + "state": {}, + "session": session if session is not None else {}, + } + + +async def _noop_receive(): # pragma: no cover + return {"type": "http.request", "body": b"", "more_body": False} + + +async def _unreached_app(scope, receive, send): # pragma: no cover + raise AssertionError("inner app must not be reached when CSRF rejects") + + +class _CapturingSend: + def __init__(self) -> None: + self.start_message: dict | None = None + + async def __call__(self, message): + if message["type"] == "http.response.start": + self.start_message = message + + +class TestCSRFMiddleware: + @pytest.mark.parametrize("method", ["GET", "HEAD", "OPTIONS", "TRACE"]) + async def test_safe_methods_pass_through(self, method: str) -> None: + """GET/HEAD/OPTIONS/TRACE must never be rejected.""" + calls = {"inner": 0} + + async def inner(scope, receive, send): + calls["inner"] += 1 + + mw = CSRFMiddleware(inner) + await mw(_scope(method, session={}), _noop_receive, _CapturingSend()) + assert calls["inner"] == 1 + + async def test_post_without_session_is_rejected(self) -> None: + """POST with no session (should never happen behind SessionMiddleware) fails closed.""" + mw = CSRFMiddleware(_unreached_app) + send = _CapturingSend() + await mw({**_scope("POST"), "session": {}}, _noop_receive, send) + assert send.start_message is not None + assert send.start_message["status"] == 403 + + async def test_post_with_mismatched_token_is_rejected(self) -> None: + send = _CapturingSend() + mw = CSRFMiddleware(_unreached_app) + await mw( + _scope( + "POST", + headers=[(b"x-csrf-token", b"wrong")], + session={"csrf_token": "expected"}, + ), + _noop_receive, + send, + ) + assert send.start_message is not None + assert send.start_message["status"] == 403 + + async def test_post_with_matching_header_passes(self) -> None: + calls = {"inner": 0} + + async def inner(scope, receive, send): + calls["inner"] += 1 + + mw = CSRFMiddleware(inner) + await mw( + _scope( + "POST", + headers=[(b"x-csrf-token", b"expected")], + session={"csrf_token": "expected"}, + ), + _noop_receive, + _CapturingSend(), + ) + assert calls["inner"] == 1 + + async def test_accepts_x_xsrf_token_alias(self) -> None: + """Accept the axios-default header name as well.""" + calls = {"inner": 0} + + async def inner(scope, receive, send): + calls["inner"] += 1 + + mw = CSRFMiddleware(inner) + await mw( + _scope( + "POST", + headers=[(b"x-xsrf-token", b"expected")], + session={"csrf_token": "expected"}, + ), + _noop_receive, + _CapturingSend(), + ) + assert calls["inner"] == 1 + + async def test_exempt_prefix_bypasses_check(self) -> None: + """Exempt prefixes should skip validation entirely.""" + calls = {"inner": 0} + + async def inner(scope, receive, send): + calls["inner"] += 1 + + mw = CSRFMiddleware(inner, exempt_path_prefixes=("/api/anything",)) + await mw(_scope("POST", session={}), _noop_receive, _CapturingSend()) + assert calls["inner"] == 1 diff --git a/framework/hosting/tests/test_settings_secrets.py b/framework/hosting/tests/test_settings_secrets.py new file mode 100644 index 00000000..c29d88f1 --- /dev/null +++ b/framework/hosting/tests/test_settings_secrets.py @@ -0,0 +1,25 @@ +"""Tests for the production secret-key guard on Settings.""" + +from __future__ import annotations + +import pytest +from simple_module_hosting.settings import Settings + + +def test_placeholder_secret_ok_in_development() -> None: + Settings(environment="development", secret_key="change-me-in-production") + + +def test_placeholder_secret_ok_in_testing() -> None: + Settings(environment="testing", secret_key="change-me-in-production") + + +def test_placeholder_secret_rejected_in_production() -> None: + with pytest.raises(ValueError, match="SM_SECRET_KEY"): + Settings(environment="production", secret_key="change-me-in-production") + + +def test_real_secret_accepted_in_production() -> None: + s = Settings(environment="production", secret_key="real-secret-value") + assert s.secret_key == "real-secret-value" + assert s.environment == "production" diff --git a/framework/testing/simple_module_testing/__init__.py b/framework/testing/simple_module_testing/__init__.py index 304eb0e6..90f94a3f 100644 --- a/framework/testing/simple_module_testing/__init__.py +++ b/framework/testing/simple_module_testing/__init__.py @@ -19,9 +19,11 @@ from simple_module_testing.app_factory import build_test_app from simple_module_testing.fake_events import FakeEventBus, RecordedEvent +from simple_module_testing.session_cookie import forge_session_cookie __all__ = [ "FakeEventBus", "RecordedEvent", "build_test_app", + "forge_session_cookie", ] diff --git a/framework/testing/simple_module_testing/session_cookie.py b/framework/testing/simple_module_testing/session_cookie.py new file mode 100644 index 00000000..b36fec71 --- /dev/null +++ b/framework/testing/simple_module_testing/session_cookie.py @@ -0,0 +1,26 @@ +"""Helpers for forging signed Starlette session cookies in tests. + +Starlette's ``SessionMiddleware`` encodes sessions as base64-encoded JSON +signed with an ``itsdangerous`` ``TimestampSigner``. Tests that want to +skip the HTTP flow (login / locale-switcher / etc.) to establish a session +can call :func:`forge_session_cookie` to build the exact cookie value the +middleware would emit, then inject it via ``httpx.AsyncClient(cookies=...)``. +""" + +from __future__ import annotations + +import json +from base64 import b64encode + +from itsdangerous import TimestampSigner + + +def forge_session_cookie(secret_key: str, session_data: dict) -> str: + """Return the signed cookie value Starlette's SessionMiddleware would emit. + + The encoding (``b64(json)`` signed with ``TimestampSigner``) must match + Starlette exactly — otherwise the middleware rejects the cookie and the + test request arrives with an empty session. + """ + data = b64encode(json.dumps(session_data).encode()) + return TimestampSigner(str(secret_key)).sign(data).decode("utf-8") diff --git a/host/client_app/app.tsx b/host/client_app/app.tsx index 7c0a2d8f..cdffc541 100644 --- a/host/client_app/app.tsx +++ b/host/client_app/app.tsx @@ -1,16 +1,45 @@ import { createInertiaApp, router } from '@inertiajs/react'; import { ErrorBoundary } from '@simple-module/ui/components/ErrorBoundary'; +import { getCsrfToken, setCsrfToken } from '@simple-module/ui/lib/csrf'; import { useEffect, useRef } from 'react'; import { createRoot } from 'react-dom/client'; import { bootI18nFromInitialPage, subscribeI18nToNavigation } from './i18n'; import { resolvePage } from './pages'; +function readCsrfToken(pageProps: unknown): string { + if (pageProps && typeof pageProps === 'object' && 'csrf_token' in pageProps) { + const token = (pageProps as { csrf_token?: unknown }).csrf_token; + if (typeof token === 'string') return token; + } + return ''; +} + +// Stamp every non-GET Inertia visit with the server's session-scoped CSRF +// token. Raw fetch() callers use ``fetchWithCsrf`` which reads from the same +// module-level store that ``setCsrfToken`` updates. +router.on('before', (event) => { + const method = event.detail.visit.method?.toLowerCase(); + const token = getCsrfToken(); + if (method && method !== 'get' && token) { + event.detail.visit.headers = { + ...(event.detail.visit.headers ?? {}), + 'X-CSRF-Token': token, + }; + } +}); + +router.on('success', (event) => { + const next = readCsrfToken(event.detail.page.props); + if (next) setCsrfToken(next); +}); + createInertiaApp({ resolve: async (name) => { const page = await resolvePage(name); return page; }, setup({ el, App, props }) { + setCsrfToken(readCsrfToken(props.initialPage.props)); bootI18nFromInitialPage(props.initialPage.props); function Root() { diff --git a/host/client_app/i18n.ts b/host/client_app/i18n.ts index 0a976dff..7551b978 100644 --- a/host/client_app/i18n.ts +++ b/host/client_app/i18n.ts @@ -13,7 +13,9 @@ import { configureI18n, updateI18n } from '@simple-module/i18n'; interface I18nSharedProps { locale: string; supportedLocales: string[]; - messages: Record; + // ``null`` on Inertia XHR visits where the backend skipped the messages + // payload because the client already has them from the initial page load. + messages: Record | null; } export function bootI18nFromInitialPage(props: PageProps): void { @@ -22,7 +24,8 @@ export function bootI18nFromInitialPage(props: PageProps): void { configureI18n({ locale: 'en', messages: {} }); return; } - configureI18n({ locale: i18n.locale, messages: i18n.messages }); + configureI18n({ locale: i18n.locale, messages: i18n.messages ?? {} }); + activeLocale = i18n.locale; } let activeLocale: string | null = null; @@ -31,7 +34,7 @@ export function subscribeI18nToNavigation(): () => void { return router.on('success', (event) => { const i18n = (event.detail.page.props as unknown as { i18n?: I18nSharedProps }).i18n; if (!i18n) return; - if (i18n.locale !== activeLocale) { + if (i18n.locale !== activeLocale && i18n.messages) { updateI18n({ locale: i18n.locale, messages: i18n.messages }); activeLocale = i18n.locale; } diff --git a/host/client_app/package.json b/host/client_app/package.json index 16e59c10..6e3ac100 100644 --- a/host/client_app/package.json +++ b/host/client_app/package.json @@ -44,6 +44,7 @@ "@vitejs/plugin-react-swc": "^4.3.0", "autoprefixer": "^10.4.0", "postcss": "^8.4.0", + "rollup-plugin-visualizer": "~5.12.0", "tailwindcss": "^4.0.0", "typescript": "^5.7.0", "vite": "^8.0.8", diff --git a/host/client_app/vite.config.ts b/host/client_app/vite.config.ts index 9d1cb008..22314b85 100644 --- a/host/client_app/vite.config.ts +++ b/host/client_app/vite.config.ts @@ -2,11 +2,16 @@ import fs from 'node:fs'; import path from 'node:path'; import tailwindcss from '@tailwindcss/vite'; import react from '@vitejs/plugin-react-swc'; +import { visualizer } from 'rollup-plugin-visualizer'; import { defineConfig } from 'vite'; import tsconfigPaths from 'vite-tsconfig-paths'; const projectRoot = path.resolve(__dirname, '../..'); +// ANALYZE=1 npm run build emits host/static/dist/stats.html — a sunburst of +// every chunk and its constituent modules. Open it to chase bundle bloat. +const analyzeBundle = process.env.ANALYZE === '1'; + // Load the module pages manifest written by the Python host at boot. // Each entry points at an absolute pages/ directory — typically inside a // pip-installed module wheel. Vite needs these in server.fs.allow so the @@ -48,6 +53,16 @@ export default defineConfig({ }), react(), tailwindcss(), + ...(analyzeBundle + ? [ + visualizer({ + filename: path.resolve(__dirname, '../static/dist/stats.html'), + template: 'sunburst', + gzipSize: true, + brotliSize: true, + }), + ] + : []), ], root: __dirname, build: { diff --git a/host/migrations/versions/b7e1af4c9d02_add_perf_indexes_and_fix_products.py b/host/migrations/versions/b7e1af4c9d02_add_perf_indexes_and_fix_products.py new file mode 100644 index 00000000..32a2ff8c --- /dev/null +++ b/host/migrations/versions/b7e1af4c9d02_add_perf_indexes_and_fix_products.py @@ -0,0 +1,133 @@ +"""Add perf indexes and drop low-cardinality boolean index. + +Creates: + * ``ix_users_user_email_lower`` — functional index on ``lower(email)`` so the + fastapi-users ``get_by_email`` query (which wraps email in ``lower()``) + can use an index instead of a seq-scan. + * ``ix_users_access_token_user_id`` — Postgres does not auto-index foreign + keys; this covers reverse lookups (e.g. revoking a user's sessions). + * ``ix_users_user_role_role_id`` — same reasoning for "who has role X?" + queries. The composite PK already covers ``user_id``-first lookups. + * ``ix_products_product_deleted`` — supports soft-delete filtering in the + product listing query. + +Drops: + * ``ix_products_product_is_active`` — a plain B-tree over a 2-value boolean + is almost never preferred by the planner over a seq-scan, yet it costs + writes on every insert/update. Replaced implicitly by the combined + filtering the ``products_product_deleted`` index covers. + +PostgreSQL path uses ``CREATE INDEX CONCURRENTLY`` via ``postgresql_concurrently`` ++ ``autocommit_block`` so index builds on large tables do not block writes. +SQLite ignores the flag (its ``CREATE INDEX`` is already fast and non-locking +for this workload). + +Revision ID: b7e1af4c9d02 +Revises: a01185374312 +Create Date: 2026-04-16 12:00:00.000000 +""" + +from __future__ import annotations + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision: str = "b7e1af4c9d02" +down_revision: str | None = "a01185374312" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +# Functional index: ``lower(email)``. Using a literal SQL expression keeps +# the syntax identical on Postgres and SQLite; SQLAlchemy's ``text()`` is +# escaped appropriately by the dialect in both cases. +_EMAIL_LOWER_EXPR = sa.text("lower(email)") + + +def upgrade() -> None: + is_postgres = op.get_context().dialect.name == "postgresql" + + # Each ``CREATE/DROP INDEX CONCURRENTLY`` auto-commits as soon as it finishes. + # ``if_exists`` / ``if_not_exists`` make the migration re-runnable after a + # partial failure mid-block: otherwise we'd have to manually clean up the + # already-committed indexes before retrying. + with op.get_context().autocommit_block(): + op.create_index( + "ix_users_user_email_lower", + "users_user", + [_EMAIL_LOWER_EXPR], + unique=False, + if_not_exists=True, + postgresql_concurrently=is_postgres, + ) + op.create_index( + "ix_users_access_token_user_id", + "users_access_token", + ["user_id"], + unique=False, + if_not_exists=True, + postgresql_concurrently=is_postgres, + ) + op.create_index( + "ix_users_user_role_role_id", + "users_user_role", + ["role_id"], + unique=False, + if_not_exists=True, + postgresql_concurrently=is_postgres, + ) + op.create_index( + "ix_products_product_is_deleted", + "products_product", + ["is_deleted"], + unique=False, + if_not_exists=True, + postgresql_concurrently=is_postgres, + ) + op.drop_index( + "ix_products_product_is_active", + table_name="products_product", + if_exists=True, + postgresql_concurrently=is_postgres, + ) + + +def downgrade() -> None: + is_postgres = op.get_context().dialect.name == "postgresql" + + with op.get_context().autocommit_block(): + op.create_index( + "ix_products_product_is_active", + "products_product", + ["is_active"], + unique=False, + if_not_exists=True, + postgresql_concurrently=is_postgres, + ) + op.drop_index( + "ix_products_product_is_deleted", + table_name="products_product", + if_exists=True, + postgresql_concurrently=is_postgres, + ) + op.drop_index( + "ix_users_user_role_role_id", + table_name="users_user_role", + if_exists=True, + postgresql_concurrently=is_postgres, + ) + op.drop_index( + "ix_users_access_token_user_id", + table_name="users_access_token", + if_exists=True, + postgresql_concurrently=is_postgres, + ) + op.drop_index( + "ix_users_user_email_lower", + table_name="users_user", + if_exists=True, + postgresql_concurrently=is_postgres, + ) diff --git a/host/routes_i18n.py b/host/routes_i18n.py index edcf19bf..f0c98177 100644 --- a/host/routes_i18n.py +++ b/host/routes_i18n.py @@ -7,9 +7,8 @@ from __future__ import annotations -from urllib.parse import urlsplit - from fastapi import APIRouter, Form, HTTPException, Request +from simple_module_hosting.redirects import safe_referer_or_root from starlette.responses import RedirectResponse router = APIRouter() @@ -17,41 +16,6 @@ _ONE_YEAR_SECONDS = 60 * 60 * 24 * 365 -def _safe_redirect_target(request: Request) -> str: - """Return the Referer iff it's same-origin; otherwise fall back to ``/``. - - An attacker can control the ``Referer`` header (e.g. via a crafted form on - a third-party site). We only honor references that (a) resolve to the - same scheme+host as the current request, or (b) are relative paths that - don't try to escape to a protocol-relative URL (``//evil.example``). - """ - referer = request.headers.get("referer") - if not referer: - return "/" - - # Reject protocol-relative URLs like "//evil.example/foo" that browsers - # would resolve against the origin but a crafted Referer could use to - # leave the site. - if referer.startswith("//"): - return "/" - - parsed = urlsplit(referer) - # Relative path with no scheme+host → same-origin by construction. - if not parsed.scheme and not parsed.netloc: - return referer if referer.startswith("/") else "/" - - # Absolute URL → must match the current request's origin. - current = request.url - if parsed.scheme == current.scheme and parsed.netloc == current.netloc: - # Preserve the path + query. - path = parsed.path or "/" - if parsed.query: - path = f"{path}?{parsed.query}" - return path - - return "/" - - @router.post("/i18n/set-locale", response_model=None) async def set_locale(request: Request, locale: str = Form(...)) -> RedirectResponse: """Persist the user's locale choice in a long-lived cookie. @@ -74,14 +38,18 @@ async def set_locale(request: Request, locale: str = Form(...)) -> RedirectRespo detail=f"Unsupported locale '{locale}' (available: {', '.join(supported)})", ) - destination = _safe_redirect_target(request) + destination = safe_referer_or_root(request) response = RedirectResponse(destination, status_code=303) + # ``secure`` lets the cookie travel over HTTP in development (SM_SECRET_KEY + # is the only thing that disambiguates dev vs. prod here); in production + # the reverse proxy strips http traffic, so the flag is safe to always set. response.set_cookie( key=cookie_name, value=locale, max_age=_ONE_YEAR_SECONDS, path="/", samesite="lax", + secure=request.url.scheme == "https", httponly=False, ) return response diff --git a/modules/auth/auth/contracts/schemas.py b/modules/auth/auth/contracts/schemas.py index 3eed8cd2..34b7b0c1 100644 --- a/modules/auth/auth/contracts/schemas.py +++ b/modules/auth/auth/contracts/schemas.py @@ -37,6 +37,37 @@ def from_user(cls, user: User | Any) -> UserContext: tenant_id=user.tenant_id, ) + def to_session_dict(self) -> dict[str, Any]: + """Serialize to a JSON-safe dict for the signed session cookie. + + The inverse of :meth:`from_session_dict`. Adding a new field to + ``UserContext`` requires updating both methods here — not the + AuthMiddleware cache helper, which stays schema-agnostic.""" + return { + "id": self.id, + "email": self.email, + "name": self.name, + "roles": list(self.roles), + "tenant_id": self.tenant_id, + } + + @classmethod + def from_session_dict(cls, payload: Any) -> UserContext | None: + """Rebuild from :meth:`to_session_dict` output. Returns ``None`` on any + shape mismatch so callers can fall through to a fresh DB load.""" + if not isinstance(payload, dict): + return None + try: + return cls( + id=str(payload["id"]), + email=str(payload["email"]), + name=str(payload["name"]), + roles=list(payload.get("roles") or []), + tenant_id=payload.get("tenant_id"), + ) + except (KeyError, TypeError, ValueError): + return None + def has_role(self, role: str) -> bool: return role in self.roles diff --git a/modules/products/products/models.py b/modules/products/products/models.py index a10c434b..aa36cab8 100644 --- a/modules/products/products/models.py +++ b/modules/products/products/models.py @@ -6,7 +6,7 @@ from simple_module_db.base import create_module_base from simple_module_db.mixins import AuditMixin, SoftDeleteMixin -from sqlalchemy import Numeric, String +from sqlalchemy import Index, Numeric, String from sqlalchemy.orm import Mapped, mapped_column # Provider is auto-detected from SM_DATABASE_URL (falls back to SQLite). @@ -24,4 +24,9 @@ class Product(Base, AuditMixin, SoftDeleteMixin): # ty: ignore[unsupported-base name: Mapped[str] = mapped_column(String(200)) description: Mapped[str | None] = mapped_column(String(2000), default=None) price: Mapped[Decimal] = mapped_column(Numeric(10, 2)) - is_active: Mapped[bool] = mapped_column(default=True, index=True) + # A plain B-tree over a 2-value boolean is rarely preferred by the planner + # and costs writes on every insert/update. The listing query filters on + # ``is_deleted`` (indexed) first, so leaving this un-indexed is faster. + is_active: Mapped[bool] = mapped_column(default=True) + + __table_args__ = (Index("ix_products_product_is_deleted", "is_deleted"),) diff --git a/modules/products/products/service.py b/modules/products/products/service.py index d9bc5789..2d461b19 100644 --- a/modules/products/products/service.py +++ b/modules/products/products/service.py @@ -22,9 +22,15 @@ async def get_all( per_page: int = 10, search: str | None = None, ) -> tuple[list[ProductOut], int]: - """Return paginated products and total count.""" - query = select(Product).where(Product.is_active.is_(True)) - count_query = select(func.count()).select_from(Product).where(Product.is_active.is_(True)) + """Return paginated products and total count. + + Excludes soft-deleted rows. Without the ``is_deleted`` filter, a + product deleted while still flagged ``is_active=True`` would leak into + public listings. + """ + base_filter = Product.is_active.is_(True) & Product.is_deleted.is_(False) + query = select(Product).where(base_filter) + count_query = select(func.count()).select_from(Product).where(base_filter) if search: pattern = f"%{search}%" diff --git a/modules/products/tests/test_products.py b/modules/products/tests/test_products.py index 751494f6..76e9d8f7 100644 --- a/modules/products/tests/test_products.py +++ b/modules/products/tests/test_products.py @@ -21,11 +21,23 @@ def _has_index_on(model, column_name: str) -> bool: class TestProductModelIndexes: - def test_is_active_is_indexed(self): + def test_is_deleted_is_indexed(self): from products.models import Product - assert _has_index_on(Product, "is_active"), ( - "Product.is_active must be indexed (used by dashboard product count query)" + assert _has_index_on(Product, "is_deleted"), ( + "Product.is_deleted must be indexed — it's the primary filter on the " + "public listing query and the planner can use the index to skip " + "tombstoned rows cheaply." + ) + + def test_is_active_is_not_indexed(self): + """``is_active`` is intentionally un-indexed — a plain B-tree over a + 2-value boolean is rarely preferred by the planner and wastes writes.""" + from products.models import Product + + assert not _has_index_on(Product, "is_active"), ( + "Product.is_active should not be indexed — drop the index if you " + "reintroduce it by mistake." ) diff --git a/modules/users/tests/conftest.py b/modules/users/tests/conftest.py index cd9b1a98..547d6954 100644 --- a/modules/users/tests/conftest.py +++ b/modules/users/tests/conftest.py @@ -11,16 +11,15 @@ from __future__ import annotations -import json import uuid -from base64 import b64encode from collections.abc import AsyncGenerator import httpx import pytest from fastapi_users.password import PasswordHelper -from itsdangerous import TimestampSigner +from simple_module_hosting.csrf import SESSION_CSRF_TOKEN_KEY from simple_module_hosting.settings import Settings +from simple_module_testing import forge_session_cookie from sqlalchemy.ext.asyncio import AsyncSession from users.constants import ADMIN_ROLE_ID, USER_ROLE_ID @@ -136,19 +135,44 @@ async def users_app_signup(monkeypatch): # --------------------------------------------------------------------------- +_TEST_CSRF_TOKEN = "test-csrf-token" + + @pytest.fixture async def anon_client(users_app) -> AsyncGenerator[httpx.AsyncClient, None]: - """Unauthenticated client against users_app.""" + """Unauthenticated client against users_app. + + Pre-seeded with a signed anonymous session that carries a CSRF token and a + matching ``X-CSRF-Token`` header, so POST flows (login, accept-invite, etc.) + pass validation without first making a GET to mint a token.""" + cookie = forge_session_cookie( + str(users_app.state.settings.secret_key), + {SESSION_CSRF_TOKEN_KEY: _TEST_CSRF_TOKEN}, + ) transport = httpx.ASGITransport(app=users_app) - async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as c: + async with httpx.AsyncClient( + transport=transport, + base_url="http://testserver", + cookies={"session": cookie}, + headers={"X-CSRF-Token": _TEST_CSRF_TOKEN}, + ) as c: yield c @pytest.fixture async def anon_client_signup(users_app_signup) -> AsyncGenerator[httpx.AsyncClient, None]: """Unauthenticated client against users_app_signup (signup enabled).""" + cookie = forge_session_cookie( + str(users_app_signup.state.settings.secret_key), + {SESSION_CSRF_TOKEN_KEY: _TEST_CSRF_TOKEN}, + ) transport = httpx.ASGITransport(app=users_app_signup) - async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as c: + async with httpx.AsyncClient( + transport=transport, + base_url="http://testserver", + cookies={"session": cookie}, + headers={"X-CSRF-Token": _TEST_CSRF_TOKEN}, + ) as c: yield c @@ -168,26 +192,20 @@ async def _make_admin_user(app): return user -def _sign_session(session_data: dict, secret_key: str) -> str: - """Encode and sign a session dict exactly as Starlette's SessionMiddleware does.""" - data = b64encode(json.dumps(session_data).encode()) - signer = TimestampSigner(secret_key) - return signer.sign(data).decode("utf-8") - - @pytest.fixture async def admin_client(users_app) -> AsyncGenerator[httpx.AsyncClient, None]: - """Client with a signed local-user session cookie (admin role).""" + """Client with a signed local-user session cookie (admin role) and CSRF token.""" user = await _make_admin_user(users_app) - cookie = _sign_session( - {"user_id": str(user.id)}, + cookie = forge_session_cookie( str(users_app.state.settings.secret_key), + {"user_id": str(user.id), SESSION_CSRF_TOKEN_KEY: _TEST_CSRF_TOKEN}, ) transport = httpx.ASGITransport(app=users_app) async with httpx.AsyncClient( transport=transport, base_url="http://testserver", cookies={"session": cookie}, + headers={"X-CSRF-Token": _TEST_CSRF_TOKEN}, ) as c: yield c diff --git a/modules/users/tests/test_api_auth.py b/modules/users/tests/test_api_auth.py index 75d154d2..bfe1f69c 100644 --- a/modules/users/tests/test_api_auth.py +++ b/modules/users/tests/test_api_auth.py @@ -185,3 +185,36 @@ async def test_register_mounted_when_signup_enabled(self, anon_client_signup): body = resp.json() assert body["email"] == "newuser@example.com" assert body["is_verified"] is False + + +# --------------------------------------------------------------------------- +# Throughput limit on shared auth side-effect endpoints +# --------------------------------------------------------------------------- + + +class TestAuthThroughputLimit: + @pytest.mark.anyio + async def test_forgot_password_rate_limited_after_threshold( + self, anon_client, users_app, users_db + ): + """After the configured attempt budget, /forgot-password returns 429.""" + from users.rate_limit import ThroughputLimiter + + # Tighten the limit for the test so we don't need to hit 10 real endpoints + users_app.state.auth_throughput_limiter = ThroughputLimiter( + max_attempts=2, window_seconds=60 + ) + + # Use a real email so the first few attempts exercise the real code path + await _make_user(users_db, email="throttle@example.com", password="SecurePass1!") + + payload = {"email": "throttle@example.com"} + # 2 within budget + r1 = await anon_client.post("/api/users/auth/forgot-password", json=payload) + r2 = await anon_client.post("/api/users/auth/forgot-password", json=payload) + assert r1.status_code in (200, 202) + assert r2.status_code in (200, 202) + + # 3rd is throttled + r3 = await anon_client.post("/api/users/auth/forgot-password", json=payload) + assert r3.status_code == 429 diff --git a/modules/users/tests/test_rate_limit.py b/modules/users/tests/test_rate_limit.py index d6f5a999..46caee5a 100644 --- a/modules/users/tests/test_rate_limit.py +++ b/modules/users/tests/test_rate_limit.py @@ -1,9 +1,9 @@ -"""Tests for LoginRateLimiter.""" +"""Tests for LoginRateLimiter and ThroughputLimiter.""" from __future__ import annotations import pytest -from users.rate_limit import LoginRateLimiter +from users.rate_limit import LoginRateLimiter, ThroughputLimiter @pytest.fixture @@ -107,3 +107,26 @@ def test_same_email_different_ip_independent(self, limiter): limiter.record_failure(key1) assert limiter.is_locked(key1) is True assert limiter.is_locked(key2) is False + + +class TestThroughputLimiter: + def test_under_budget_passes(self): + limiter = ThroughputLimiter(max_attempts=3, window_seconds=60) + assert limiter.check_and_record("k") is True + assert limiter.check_and_record("k") is True + assert limiter.check_and_record("k") is True + + def test_over_budget_rejected(self): + limiter = ThroughputLimiter(max_attempts=3, window_seconds=60) + for _ in range(3): + limiter.check_and_record("k") + assert limiter.check_and_record("k") is False + # Still rejected after further attempts within window + assert limiter.check_and_record("k") is False + + def test_keys_isolated(self): + limiter = ThroughputLimiter(max_attempts=2, window_seconds=60) + limiter.check_and_record("a") + limiter.check_and_record("a") + assert limiter.check_and_record("a") is False + assert limiter.check_and_record("b") is True diff --git a/modules/users/tests/test_settings.py b/modules/users/tests/test_settings.py index b9dbc15d..53c1c19a 100644 --- a/modules/users/tests/test_settings.py +++ b/modules/users/tests/test_settings.py @@ -95,3 +95,47 @@ def test_mailer_pattern_accepts_smtp(self): s = UsersSettings(mailer="smtp") assert s.mailer == "smtp" + + +class TestTokenSecretProductionGuard: + def test_placeholder_secrets_ok_in_development(self, monkeypatch): + monkeypatch.setenv("SM_ENVIRONMENT", "development") + from users.settings import UsersSettings + + UsersSettings() # must not raise + + def test_placeholder_secrets_ok_in_testing(self, monkeypatch): + monkeypatch.setenv("SM_ENVIRONMENT", "testing") + from users.settings import UsersSettings + + UsersSettings() # must not raise + + def test_placeholder_reset_secret_rejected_in_production(self, monkeypatch): + monkeypatch.setenv("SM_ENVIRONMENT", "production") + monkeypatch.setenv( + "SM_USERS_VERIFICATION_TOKEN_SECRET", "not-a-placeholder-value-just-for-test" + ) + from users.settings import UsersSettings + + with pytest.raises(ValidationError, match="RESET_PASSWORD_TOKEN_SECRET"): + UsersSettings() + + def test_placeholder_verify_secret_rejected_in_production(self, monkeypatch): + monkeypatch.setenv("SM_ENVIRONMENT", "production") + monkeypatch.setenv( + "SM_USERS_RESET_PASSWORD_TOKEN_SECRET", "not-a-placeholder-value-just-for-test" + ) + from users.settings import UsersSettings + + with pytest.raises(ValidationError, match="VERIFICATION_TOKEN_SECRET"): + UsersSettings() + + def test_real_secrets_accepted_in_production(self, monkeypatch): + monkeypatch.setenv("SM_ENVIRONMENT", "production") + monkeypatch.setenv("SM_USERS_RESET_PASSWORD_TOKEN_SECRET", "real-reset-secret") + monkeypatch.setenv("SM_USERS_VERIFICATION_TOKEN_SECRET", "real-verify-secret") + from users.settings import UsersSettings + + s = UsersSettings() + assert s.reset_password_token_secret == "real-reset-secret" + assert s.verification_token_secret == "real-verify-secret" diff --git a/modules/users/users/endpoints/api.py b/modules/users/users/endpoints/api.py index 6dde8b8f..1a65f4a0 100644 --- a/modules/users/users/endpoints/api.py +++ b/modules/users/users/endpoints/api.py @@ -11,37 +11,26 @@ from __future__ import annotations import logging -import uuid from fastapi import APIRouter, Depends, HTTPException, Request, Response, status from fastapi.security import OAuth2PasswordRequestForm from fastapi_users import exceptions as fu_exceptions -from simple_module_core.events import EventBus -from simple_module_hosting.permissions import RequiresPermission -from users.contracts.events import RoleAssigned, UserDisabled, UserInvited from users.contracts.schemas import ( AcceptInviteRequest, - PasswordResetLink, - RoleAssignment, SelfProfileUpdate, UserCreate, - UserInvite, - UserListItem, UserRead, UserUpdate, ) from users.deps import ( auth_backend, fastapi_users, - get_event_bus, - get_mailer, get_user_manager, - get_user_service, ) +from users.endpoints.api_admin import admin_router from users.manager import UserManager -from users.rate_limit import LoginRateLimiter -from users.service import UserService +from users.rate_limit import LoginRateLimiter, ThroughputLimiter logger = logging.getLogger(__name__) router = APIRouter() @@ -55,6 +44,26 @@ def get_rate_limiter(request: Request) -> LoginRateLimiter: return request.app.state.rate_limiter +def _client_ip(request: Request) -> str: + return request.client.host if request.client else "unknown" + + +async def enforce_auth_throughput_limit(request: Request) -> None: + """FastAPI dependency that rejects the request with 429 when this IP has + exhausted its attempts budget on shared auth side-effect endpoints. + + Applied to forgot-password / register / accept-invite / request-verify-token, + which otherwise allow unlimited email or account-creation spam. + """ + limiter: ThroughputLimiter = request.app.state.auth_throughput_limiter + key = f"{request.url.path}::{_client_ip(request)}" + if not limiter.check_and_record(key): + raise HTTPException( + status_code=status.HTTP_429_TOO_MANY_REQUESTS, + detail="Too many attempts — try again later", + ) + + # ── Wrapper login ──────────────────────────────────────────────────────────── @@ -105,30 +114,45 @@ async def login( def register_auth_routes(api_router: APIRouter, settings) -> None: - """Mount all auth routes, conditionally adding register if allowed.""" + """Mount all auth routes, conditionally adding register if allowed. + + The stock fastapi-users routers (reset/verify/register) ship POST endpoints + that trigger email side-effects or account creation. We wrap them with the + throughput limiter so an attacker can't spam password-reset emails or mint + accounts indefinitely. ``router`` itself is left unwrapped because its + rate-limited endpoints apply the dep themselves (login via LoginRateLimiter, + accept-invite via ``enforce_auth_throughput_limit``). + """ api_router.include_router(router) api_router.include_router( fastapi_users.get_reset_password_router(), prefix="/auth", tags=["users-auth"], + dependencies=[Depends(enforce_auth_throughput_limit)], ) api_router.include_router( fastapi_users.get_verify_router(UserRead), prefix="/auth", tags=["users-auth"], + dependencies=[Depends(enforce_auth_throughput_limit)], ) if settings.allow_signup: api_router.include_router( fastapi_users.get_register_router(UserRead, UserCreate), prefix="/auth", tags=["users-auth"], + dependencies=[Depends(enforce_auth_throughput_limit)], ) # ── Accept-invite (verify + set password + login, one shot) ───────────────── -@router.post("/auth/accept-invite", status_code=204) +@router.post( + "/auth/accept-invite", + status_code=204, + dependencies=[Depends(enforce_auth_throughput_limit)], +) async def accept_invite( body: AcceptInviteRequest, request: Request, @@ -181,120 +205,6 @@ async def update_me( ) -# ── Admin REST ─────────────────────────────────────────────────────────────── - - -@router.get( - "/admin", - response_model=list[UserListItem], - dependencies=[Depends(RequiresPermission("users.manage"))], -) -async def admin_list_users( - page: int = 1, - per_page: int = 20, - q: str | None = None, - service: UserService = Depends(get_user_service), -): - """List all users (paginated, optional search).""" - items, _ = await service.list_users(page=page, per_page=per_page, search=q) - return items - - -@router.post( - "/admin/invite", - response_model=UserListItem, - status_code=status.HTTP_201_CREATED, - dependencies=[Depends(RequiresPermission("users.manage"))], -) -async def admin_invite_user( - data: UserInvite, - request: Request, - bus: EventBus = Depends(get_event_bus), - service: UserService = Depends(get_user_service), - mailer=Depends(get_mailer), -): - """Invite a new user by email, optionally assigning roles.""" - invited_by = getattr(request.state, "user", None) - invited_by_name = invited_by.name if invited_by else "Administrator" - user, token = await service.invite( - data.email, data.full_name, data.role_names, invited_by=invited_by - ) - await mailer.send_invite(user.email, token, invited_by_name) - await bus.publish( - UserInvited( - user_id=user.id, - email=user.email, - invited_by=(str(invited_by.id) if invited_by else None), - ) - ) - return await service.to_list_item(user) - - -@router.patch( - "/admin/{user_id}/disable", - response_model=UserListItem, - dependencies=[Depends(RequiresPermission("users.manage"))], -) -async def admin_disable_user( - user_id: uuid.UUID, - bus: EventBus = Depends(get_event_bus), - service: UserService = Depends(get_user_service), -): - """Disable a user account (sets is_active=False and disabled_at).""" - user = await service.disable(user_id) - await bus.publish(UserDisabled(user_id=user.id)) - return await service.to_list_item(user) - - -@router.patch( - "/admin/{user_id}/enable", - response_model=UserListItem, - dependencies=[Depends(RequiresPermission("users.manage"))], -) -async def admin_enable_user( - user_id: uuid.UUID, - service: UserService = Depends(get_user_service), -): - """Re-enable a previously disabled user account.""" - user = await service.enable(user_id) - return await service.to_list_item(user) - - -@router.put( - "/admin/{user_id}/roles", - response_model=UserListItem, - dependencies=[Depends(RequiresPermission("users.manage"))], -) -async def admin_set_roles( - user_id: uuid.UUID, - data: RoleAssignment, - request: Request, - bus: EventBus = Depends(get_event_bus), - service: UserService = Depends(get_user_service), -): - """Replace a user's role assignments.""" - assigned_by = getattr(request.state, "user", None) - user = await service.set_roles( - user_id, - data.role_names, - assigned_by=str(assigned_by.id) if assigned_by else None, - ) - for role in data.role_names: - await bus.publish(RoleAssigned(user_id=user.id, role_name=role)) - return await service.to_list_item(user) - +# ── Admin REST — defined in api_admin.py, mounted here ────────────────────── -@router.post( - "/admin/{user_id}/reset-password-link", - response_model=PasswordResetLink, - dependencies=[Depends(RequiresPermission("users.manage"))], -) -async def admin_reset_password_link( - user_id: uuid.UUID, - request: Request, - service: UserService = Depends(get_user_service), -): - """Generate a password-reset link for the given user (admin copy).""" - base_url = request.app.state.users_settings.base_url - link = await service.generate_reset_link(user_id, base_url) - return PasswordResetLink(link=link) +router.include_router(admin_router) diff --git a/modules/users/users/endpoints/api_admin.py b/modules/users/users/endpoints/api_admin.py new file mode 100644 index 00000000..8c8588e9 --- /dev/null +++ b/modules/users/users/endpoints/api_admin.py @@ -0,0 +1,124 @@ +"""Admin REST endpoints for the users module. + +Split out of :mod:`.api` to keep per-file complexity manageable. Mounted +into the main ``router`` via ``include_router`` at the bottom of ``api.py``. +""" + +from __future__ import annotations + +import uuid + +from fastapi import APIRouter, Depends, Request, status +from simple_module_core.events import EventBus +from simple_module_hosting.permissions import RequiresPermission + +from users.contracts.events import RoleAssigned, UserDisabled, UserInvited +from users.contracts.schemas import ( + PasswordResetLink, + RoleAssignment, + UserInvite, + UserListItem, +) +from users.deps import get_event_bus, get_mailer, get_user_service +from users.service import UserService + +admin_router = APIRouter( + prefix="/admin", + dependencies=[Depends(RequiresPermission("users.manage"))], + tags=["users-admin"], +) + + +@admin_router.get("", response_model=list[UserListItem]) +async def admin_list_users( + page: int = 1, + per_page: int = 20, + q: str | None = None, + service: UserService = Depends(get_user_service), +): + """List all users (paginated, optional search).""" + items, _ = await service.list_users(page=page, per_page=per_page, search=q) + return items + + +@admin_router.post( + "/invite", + response_model=UserListItem, + status_code=status.HTTP_201_CREATED, +) +async def admin_invite_user( + data: UserInvite, + request: Request, + bus: EventBus = Depends(get_event_bus), + service: UserService = Depends(get_user_service), + mailer=Depends(get_mailer), +): + """Invite a new user by email, optionally assigning roles.""" + invited_by = getattr(request.state, "user", None) + invited_by_name = invited_by.name if invited_by else "Administrator" + user, token = await service.invite( + data.email, data.full_name, data.role_names, invited_by=invited_by + ) + await mailer.send_invite(user.email, token, invited_by_name) + await bus.publish( + UserInvited( + user_id=user.id, + email=user.email, + invited_by=(str(invited_by.id) if invited_by else None), + ) + ) + return await service.to_list_item(user) + + +@admin_router.patch("/{user_id}/disable", response_model=UserListItem) +async def admin_disable_user( + user_id: uuid.UUID, + bus: EventBus = Depends(get_event_bus), + service: UserService = Depends(get_user_service), +): + """Disable a user account (sets is_active=False and disabled_at).""" + user = await service.disable(user_id) + await bus.publish(UserDisabled(user_id=user.id)) + return await service.to_list_item(user) + + +@admin_router.patch("/{user_id}/enable", response_model=UserListItem) +async def admin_enable_user( + user_id: uuid.UUID, + service: UserService = Depends(get_user_service), +): + """Re-enable a previously disabled user account.""" + user = await service.enable(user_id) + return await service.to_list_item(user) + + +@admin_router.put("/{user_id}/roles", response_model=UserListItem) +async def admin_set_roles( + user_id: uuid.UUID, + data: RoleAssignment, + request: Request, + bus: EventBus = Depends(get_event_bus), + service: UserService = Depends(get_user_service), +): + """Replace a user's role assignments.""" + assigned_by = getattr(request.state, "user", None) + user = await service.set_roles( + user_id, + data.role_names, + assigned_by=str(assigned_by.id) if assigned_by else None, + ) + for role in data.role_names: + await bus.publish(RoleAssigned(user_id=user.id, role_name=role)) + return await service.to_list_item(user) + + +@admin_router.post("/{user_id}/reset-password-link", response_model=PasswordResetLink) +async def admin_reset_password_link( + user_id: uuid.UUID, + request: Request, + service: UserService = Depends(get_user_service), +): + """Generate a password-reset link for the given user (admin copy).""" + base_url = request.app.state.users_settings.base_url + link = await service.generate_reset_link(user_id, base_url) + return PasswordResetLink(link=link) diff --git a/modules/users/users/endpoints/views.py b/modules/users/users/endpoints/views.py index 7c6fc932..8f7e5143 100644 --- a/modules/users/users/endpoints/views.py +++ b/modules/users/users/endpoints/views.py @@ -6,15 +6,12 @@ from fastapi import APIRouter, Depends, HTTPException, Request from inertia import InertiaResponse -from simple_module_db.deps import get_db from simple_module_hosting.inertia_deps import InertiaDep from simple_module_hosting.permissions import RequiresPermission -from sqlalchemy import select -from sqlalchemy.ext.asyncio import AsyncSession from starlette.responses import RedirectResponse from users.deps import get_user_service -from users.models import Role +from users.roles_cache import get_roles_cache from users.service import UserService router = APIRouter() @@ -26,41 +23,21 @@ @router.get("/login", response_model=None) async def login_page(request: Request, inertia: InertiaDep) -> InertiaResponse: users_settings = request.app.state.users_settings - allow_signup = users_settings.allow_signup - dev_accounts: list[dict[str, str]] = [] - # Only expose seeded credentials in dev, and only when the env vars that - # actually seeded them are still set (so production configs that happen to - # boot with SM_ENVIRONMENT=development never leak real passwords). - host_settings = getattr(request.app.state, "settings", None) - if host_settings is not None and host_settings.is_development: - if users_settings.bootstrap_email and users_settings.bootstrap_password: - dev_accounts.append( - { - "label": "Admin", - "email": users_settings.bootstrap_email, - "password": users_settings.bootstrap_password, - } - ) - if users_settings.bootstrap_user_email and users_settings.bootstrap_user_password: - dev_accounts.append( - { - "label": "User", - "email": users_settings.bootstrap_user_email, - "password": users_settings.bootstrap_user_password, - } - ) return await inertia.render( "Users/Login", - {"allow_signup": allow_signup, "dev_accounts": dev_accounts}, + {"allow_signup": users_settings.allow_signup}, ) -@router.get("/logout", response_model=None) +@router.post("/logout", response_model=None) async def logout(request: Request) -> RedirectResponse: - """GET-able logout for menu links — clears the session + auth cookie.""" + """Clear the session + auth cookie. POST-only to resist cross-site `` + logout attacks — the menu's logout link submits this as an Inertia form.""" request.session.clear() cookie_name = request.app.state.users_settings.cookie_name - response = RedirectResponse("/", status_code=302) + # 303 forces the follow-up to GET — Inertia treats the redirect as a full + # navigation rather than replaying the POST. + response = RedirectResponse("/", status_code=303) response.delete_cookie(cookie_name, path="/") return response @@ -109,22 +86,21 @@ async def profile_page(inertia: InertiaDep) -> InertiaResponse: dependencies=[Depends(RequiresPermission("users.manage"))], ) async def admin_index( + request: Request, inertia: InertiaDep, service: UserService = Depends(get_user_service), page: int = 1, per_page: int = 20, q: str | None = None, - db: AsyncSession = Depends(get_db), ) -> InertiaResponse: users, total = await service.list_users(page=page, per_page=per_page, search=q) - roles_list = (await db.execute(select(Role).order_by(Role.name))).scalars().all() return await inertia.render( "Users/Users/Index", { "users": [u.model_dump(mode="json") for u in users], "pagination": {"page": page, "per_page": per_page, "total": total}, "query": q or "", - "roles": [{"id": str(r.id), "name": r.name} for r in roles_list], + "roles": [{"id": r.id, "name": r.name} for r in await get_roles_cache(request.app)], }, ) @@ -135,14 +111,13 @@ async def admin_index( dependencies=[Depends(RequiresPermission("users.manage"))], ) async def admin_invite_page( + request: Request, inertia: InertiaDep, - db: AsyncSession = Depends(get_db), ) -> InertiaResponse: - roles_list = (await db.execute(select(Role).order_by(Role.name))).scalars().all() return await inertia.render( "Users/Users/Invite", { - "roles": [{"id": str(r.id), "name": r.name} for r in roles_list], + "roles": [{"id": r.id, "name": r.name} for r in await get_roles_cache(request.app)], }, ) @@ -154,9 +129,9 @@ async def admin_invite_page( ) async def admin_edit_page( user_id: str, + request: Request, inertia: InertiaDep, service: UserService = Depends(get_user_service), - db: AsyncSession = Depends(get_db), ) -> InertiaResponse: try: uid = uuid.UUID(user_id) @@ -165,11 +140,10 @@ async def admin_edit_page( user_item = await service.get_list_item(uid) if user_item is None: raise HTTPException(status_code=404) - roles_list = (await db.execute(select(Role).order_by(Role.name))).scalars().all() return await inertia.render( "Users/Users/Edit", { "user": user_item.model_dump(mode="json"), - "roles": [{"id": str(r.id), "name": r.name} for r in roles_list], + "roles": [{"id": r.id, "name": r.name} for r in await get_roles_cache(request.app)], }, ) diff --git a/modules/users/users/manager.py b/modules/users/users/manager.py index 532322d3..002293a6 100644 --- a/modules/users/users/manager.py +++ b/modules/users/users/manager.py @@ -2,6 +2,7 @@ from __future__ import annotations +import asyncio import uuid from datetime import UTC, datetime from typing import TYPE_CHECKING @@ -91,18 +92,25 @@ async def generate_verification_token(self, user: User) -> str: self.verification_token_lifetime_seconds, ) - def generate_reset_password_token(self, user: User) -> str: + async def generate_reset_password_token(self, user: User) -> str: """Mint a reset-audience JWT without firing on_after_forgot_password. Shape matches ``BaseUserManager.forgot_password``: the payload includes - ``password_fgpt`` (a hash of the current hashed_password) so the token + ``password_fgpt`` (a bcrypt of the current hashed_password) so the token invalidates when the password changes. Used by the admin reset-password-link endpoint, where the admin copies the link instead of triggering an email. + + The fingerprint must be bcrypt-style because the stock + ``reset_password`` path in fastapi-users verifies it with + ``password_helper.verify_and_update``. Bcrypt is CPU-bound (~100ms at + default rounds) so we offload it to a worker thread — otherwise a + single admin action would stall the event loop for other requests. """ + fingerprint = await asyncio.to_thread(self.password_helper.hash, user.hashed_password) token_data = { "sub": str(user.id), - "password_fgpt": self.password_helper.hash(user.hashed_password), + "password_fgpt": fingerprint, "aud": self.reset_password_token_audience, } return generate_jwt( diff --git a/modules/users/users/middleware.py b/modules/users/users/middleware.py index 5fb05e77..c6cd10c8 100644 --- a/modules/users/users/middleware.py +++ b/modules/users/users/middleware.py @@ -4,6 +4,13 @@ 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``. """ @@ -24,10 +31,11 @@ logger = logging.getLogger(__name__) +SESSION_USER_CTX_KEY = "user_ctx" + # Paths that don't require authentication. PUBLIC_PATHS = ( "/users/login", - "/users/logout", "/users/register", "/users/forgot-password", "/users/reset-password", @@ -48,9 +56,9 @@ class AuthMiddleware: """Redirect unauthenticated users to /users/login. - Loads the authenticated user from DB on every request. Sets - ``request.state.user`` and the ``current_user_id`` ContextVar so audit - listeners stamp created_by / updated_by correctly. + 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: @@ -69,16 +77,25 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: user_ctx: UserContext | None = None if raw_user_id: - try: - user_uuid = uuid.UUID(raw_user_id) - except (ValueError, TypeError): - logger.warning("Invalid user_id in session: %r", raw_user_id) - session.pop("user_id", None) - else: - user_ctx = await self._load_user(scope, user_uuid) - if user_ctx is None: - # User was deleted / disabled since session creation. + 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("user_id", 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("user_id", None) + session.pop(SESSION_USER_CTX_KEY, None) + else: + session[SESSION_USER_CTX_KEY] = user_ctx.to_session_dict() if user_ctx is None and not is_public: request = Request(scope) diff --git a/modules/users/users/models.py b/modules/users/users/models.py index ea9f917d..2795956a 100644 --- a/modules/users/users/models.py +++ b/modules/users/users/models.py @@ -16,7 +16,7 @@ from fastapi_users_db_sqlalchemy.generics import GUID from simple_module_db.base import create_module_base from simple_module_db.mixins import AuditMixin -from sqlalchemy import DateTime, ForeignKey, String, func +from sqlalchemy import DateTime, ForeignKey, Index, String, func, text from sqlalchemy.orm import Mapped, declared_attr, mapped_column, relationship Base = create_module_base("users") @@ -43,6 +43,10 @@ class User(SQLAlchemyBaseUserTableUUID, Base, AuditMixin): # ty: ignore[unsuppo back_populates="users", ) + # Functional index so the ``lower(email)`` predicate used by + # ``UserDatabaseWithRoles.get_by_email`` can be served from an index. + __table_args__ = (Index("ix_users_user_email_lower", text("lower(email)")),) + class Role(Base, AuditMixin): # ty: ignore[unsupported-base] """A named role that can hold a set of permission strings.""" @@ -80,6 +84,11 @@ class UserRole(Base): # ty: ignore[unsupported-base] ) assigned_by: Mapped[str | None] = mapped_column(String(255), default=None) + # The composite PK covers ``user_id``-first lookups; add a standalone + # index on ``role_id`` for reverse lookups — PostgreSQL does not + # auto-index FKs. + __table_args__ = (Index("ix_users_user_role_role_id", "role_id"),) + class UserAccessToken(SQLAlchemyBaseAccessTokenTable[uuid.UUID], Base): # ty: ignore[unsupported-base] """fastapi-users DatabaseStrategy-backed access tokens.""" @@ -99,6 +108,7 @@ def user_id(self) -> Mapped[uuid.UUID]: GUID, ForeignKey("users_user.id", ondelete="CASCADE"), nullable=False, + index=True, # PostgreSQL does not auto-index FKs ) diff --git a/modules/users/users/module.py b/modules/users/users/module.py index ca7f6428..2a57dabd 100644 --- a/modules/users/users/module.py +++ b/modules/users/users/module.py @@ -62,6 +62,7 @@ def register_menu_items(self, registry: MenuRegistry) -> None: icon="log-out", order=999, section=MenuSection.USER_DROPDOWN, + method="post", ) ) @@ -82,11 +83,14 @@ def register_middleware(self, app: FastAPI) -> None: async def on_startup(self, app: FastAPI) -> None: """Build the mailer, rate limiter, and apply production cookie params.""" + import asyncio + from users.backend import reconfigure_cookie_transport from users.bootstrap import bootstrap_admin_from_env from users.deps import auth_backend from users.mailer import build_mailer - from users.rate_limit import LoginRateLimiter + from users.rate_limit import LoginRateLimiter, ThroughputLimiter + from users.roles_cache import refresh_roles_cache s = app.state.users_settings app.state.mailer = build_mailer(s) @@ -95,8 +99,16 @@ async def on_startup(self, app: FastAPI) -> None: window_seconds=s.login_rate_limit_window_seconds, cooldown_seconds=s.login_rate_limit_cooldown_seconds, ) + app.state.auth_throughput_limiter = ThroughputLimiter( + max_attempts=s.auth_rate_limit_attempts, + window_seconds=s.auth_rate_limit_window_seconds, + ) reconfigure_cookie_transport(auth_backend, s) - # Auto-create admin iff users table empty and both env vars set. - # Runs LAST so the mailer + cookie state are already built. - await bootstrap_admin_from_env(app) + # Bootstrap + roles-cache hit different tables and have no data + # dependency on each other — run them concurrently to shave a DB + # round-trip off startup. + await asyncio.gather( + bootstrap_admin_from_env(app), + refresh_roles_cache(app), + ) diff --git a/modules/users/users/pages/AcceptInvite.tsx b/modules/users/users/pages/AcceptInvite.tsx index ff00993b..c0482007 100644 --- a/modules/users/users/pages/AcceptInvite.tsx +++ b/modules/users/users/pages/AcceptInvite.tsx @@ -10,6 +10,7 @@ import { import { Input } from '@simple-module/ui/components/ui/input'; import { Label } from '@simple-module/ui/components/ui/label'; import { AuthCardShell } from '@simple-module/ui/layouts/AuthCardShell'; +import { fetchWithCsrf } from '@simple-module/ui/lib/csrf'; import { useState } from 'react'; interface Props { @@ -37,7 +38,7 @@ function AcceptInvite() { return; } setLoading(true); - fetch('/api/users/auth/accept-invite', { + fetchWithCsrf('/api/users/auth/accept-invite', { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ token, password }), diff --git a/modules/users/users/pages/ForgotPassword.tsx b/modules/users/users/pages/ForgotPassword.tsx index 2d90fbdc..9e44d31b 100644 --- a/modules/users/users/pages/ForgotPassword.tsx +++ b/modules/users/users/pages/ForgotPassword.tsx @@ -9,6 +9,7 @@ import { import { Input } from '@simple-module/ui/components/ui/input'; import { Label } from '@simple-module/ui/components/ui/label'; import { AuthCardShell } from '@simple-module/ui/layouts/AuthCardShell'; +import { fetchWithCsrf } from '@simple-module/ui/lib/csrf'; import { useState } from 'react'; function ForgotPassword() { @@ -19,7 +20,7 @@ function ForgotPassword() { const handleSubmit = (e: React.FormEvent) => { e.preventDefault(); setLoading(true); - fetch('/api/users/auth/forgot-password', { + fetchWithCsrf('/api/users/auth/forgot-password', { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ email }), diff --git a/modules/users/users/pages/Login.tsx b/modules/users/users/pages/Login.tsx index 171f94e0..ccf85186 100644 --- a/modules/users/users/pages/Login.tsx +++ b/modules/users/users/pages/Login.tsx @@ -10,21 +10,15 @@ import { import { Input } from '@simple-module/ui/components/ui/input'; import { Label } from '@simple-module/ui/components/ui/label'; import { AuthCardShell } from '@simple-module/ui/layouts/AuthCardShell'; +import { fetchWithCsrf } from '@simple-module/ui/lib/csrf'; import { useState } from 'react'; -interface DevAccount { - label: string; - email: string; - password: string; -} - interface Props { allow_signup: boolean; - dev_accounts: DevAccount[]; } function Login() { - const { allow_signup, dev_accounts } = usePage<{ props: Props }>().props as unknown as Props; + const { allow_signup } = usePage<{ props: Props }>().props as unknown as Props; const [email, setEmail] = useState(''); const [password, setPassword] = useState(''); @@ -42,7 +36,7 @@ function Login() { setNeedsVerification(false); setLoading(true); const body = new URLSearchParams({ username, password: pwd }); - fetch('/api/users/auth/login', { + fetchWithCsrf('/api/users/auth/login', { method: 'POST', body, headers: { 'Content-Type': 'application/x-www-form-urlencoded' }, @@ -71,14 +65,8 @@ function Login() { submitLogin(email, password); }; - const handleDevLogin = (account: DevAccount) => { - setEmail(account.email); - setPassword(account.password); - submitLogin(account.email, account.password); - }; - const handleResendVerification = () => { - fetch('/api/users/auth/request-verify-token', { + fetchWithCsrf('/api/users/auth/request-verify-token', { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ email }), @@ -146,27 +134,6 @@ function Login() { - {dev_accounts && dev_accounts.length > 0 && ( -
-

Dev: log in as seeded user

-
- {dev_accounts.map((account) => ( - - ))} -
-
- )} - {allow_signup && (

Don't have an account?{' '} diff --git a/modules/users/users/pages/Profile.tsx b/modules/users/users/pages/Profile.tsx index 0737d7b1..9b876384 100644 --- a/modules/users/users/pages/Profile.tsx +++ b/modules/users/users/pages/Profile.tsx @@ -6,6 +6,7 @@ import { Card, CardContent } from '@simple-module/ui/components/ui/card'; import { Input } from '@simple-module/ui/components/ui/input'; import { Label } from '@simple-module/ui/components/ui/label'; import { AuthenticatedLayout } from '@simple-module/ui/layouts/AuthenticatedLayout'; +import { fetchWithCsrf } from '@simple-module/ui/lib/csrf'; import { useState } from 'react'; import { toast } from 'sonner'; @@ -33,7 +34,7 @@ function Profile() { const handleSubmit = (e: React.FormEvent) => { e.preventDefault(); setSaving(true); - fetch('/api/users/me', { + fetchWithCsrf('/api/users/me', { method: 'PATCH', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ full_name: fullName }), diff --git a/modules/users/users/pages/Register.tsx b/modules/users/users/pages/Register.tsx index 156e9971..3d8f04d2 100644 --- a/modules/users/users/pages/Register.tsx +++ b/modules/users/users/pages/Register.tsx @@ -9,6 +9,7 @@ import { import { Input } from '@simple-module/ui/components/ui/input'; import { Label } from '@simple-module/ui/components/ui/label'; import { AuthCardShell } from '@simple-module/ui/layouts/AuthCardShell'; +import { fetchWithCsrf } from '@simple-module/ui/lib/csrf'; import { useState } from 'react'; function Register() { @@ -28,7 +29,7 @@ function Register() { return; } setLoading(true); - fetch('/api/users/auth/register', { + fetchWithCsrf('/api/users/auth/register', { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ email, password, full_name: fullName }), diff --git a/modules/users/users/pages/ResetPassword.tsx b/modules/users/users/pages/ResetPassword.tsx index 99fe85d5..a6627d04 100644 --- a/modules/users/users/pages/ResetPassword.tsx +++ b/modules/users/users/pages/ResetPassword.tsx @@ -10,6 +10,7 @@ import { import { Input } from '@simple-module/ui/components/ui/input'; import { Label } from '@simple-module/ui/components/ui/label'; import { AuthCardShell } from '@simple-module/ui/layouts/AuthCardShell'; +import { fetchWithCsrf } from '@simple-module/ui/lib/csrf'; import { useState } from 'react'; interface Props { @@ -39,7 +40,7 @@ function ResetPassword() { return; } setLoading(true); - fetch('/api/users/auth/reset-password', { + fetchWithCsrf('/api/users/auth/reset-password', { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ token, password }), diff --git a/modules/users/users/pages/Users/Edit.tsx b/modules/users/users/pages/Users/Edit.tsx index 8bddce44..fcda114b 100644 --- a/modules/users/users/pages/Users/Edit.tsx +++ b/modules/users/users/pages/Users/Edit.tsx @@ -6,6 +6,7 @@ import { Card, CardContent, CardHeader, CardTitle } from '@simple-module/ui/comp import { Checkbox } from '@simple-module/ui/components/ui/checkbox'; import { Label } from '@simple-module/ui/components/ui/label'; import { AuthenticatedLayout } from '@simple-module/ui/layouts/AuthenticatedLayout'; +import { fetchWithCsrf } from '@simple-module/ui/lib/csrf'; import { useState } from 'react'; import { toast } from 'sonner'; @@ -48,7 +49,7 @@ function Edit() { const endpoint = isActive ? `/api/users/admin/${user.id}/disable` : `/api/users/admin/${user.id}/enable`; - fetch(endpoint, { method: 'PATCH' }) + fetchWithCsrf(endpoint, { method: 'PATCH' }) .then(async (res) => { if (res.ok) { const newActive = !isActive; @@ -65,7 +66,7 @@ function Edit() { const handleSaveRoles = () => { setSavingRoles(true); - fetch(`/api/users/admin/${user.id}/roles`, { + fetchWithCsrf(`/api/users/admin/${user.id}/roles`, { method: 'PUT', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ role_names: selectedRoles }), @@ -83,7 +84,7 @@ function Edit() { }; const handleCopyResetLink = () => { - fetch(`/api/users/admin/${user.id}/reset-password-link`, { method: 'POST' }) + fetchWithCsrf(`/api/users/admin/${user.id}/reset-password-link`, { method: 'POST' }) .then(async (res) => { if (res.ok) { const data = await res.json(); diff --git a/modules/users/users/pages/Users/Invite.tsx b/modules/users/users/pages/Users/Invite.tsx index 6e878f3b..d0d7513b 100644 --- a/modules/users/users/pages/Users/Invite.tsx +++ b/modules/users/users/pages/Users/Invite.tsx @@ -6,6 +6,7 @@ import { Checkbox } from '@simple-module/ui/components/ui/checkbox'; import { Input } from '@simple-module/ui/components/ui/input'; import { Label } from '@simple-module/ui/components/ui/label'; import { AuthenticatedLayout } from '@simple-module/ui/layouts/AuthenticatedLayout'; +import { fetchWithCsrf } from '@simple-module/ui/lib/csrf'; import { useState } from 'react'; import { toast } from 'sonner'; @@ -37,7 +38,7 @@ function Invite() { e.preventDefault(); setError(null); setLoading(true); - fetch('/api/users/admin/invite', { + fetchWithCsrf('/api/users/admin/invite', { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ email, full_name: fullName || null, role_names: selectedRoles }), diff --git a/modules/users/users/pages/VerifyEmail.tsx b/modules/users/users/pages/VerifyEmail.tsx index e94722b8..2598b723 100644 --- a/modules/users/users/pages/VerifyEmail.tsx +++ b/modules/users/users/pages/VerifyEmail.tsx @@ -8,6 +8,7 @@ import { CardTitle, } from '@simple-module/ui/components/ui/card'; import { AuthCardShell } from '@simple-module/ui/layouts/AuthCardShell'; +import { fetchWithCsrf } from '@simple-module/ui/lib/csrf'; import { useEffect, useState } from 'react'; interface Props { @@ -34,7 +35,7 @@ function VerifyEmail() { return; } - fetch('/api/users/auth/verify', { + fetchWithCsrf('/api/users/auth/verify', { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ token }), diff --git a/modules/users/users/rate_limit.py b/modules/users/users/rate_limit.py index 8f71a147..5f96f5e1 100644 --- a/modules/users/users/rate_limit.py +++ b/modules/users/users/rate_limit.py @@ -1,4 +1,10 @@ -"""In-process login rate limiter — TTL caches, no Redis.""" +"""In-process rate limiters — TTL caches, no Redis. + +Note: counters live in the worker process, so a multi-worker deployment (e.g. +``uvicorn --workers 4``) has independent counters per worker. Effective +thresholds scale with worker count. Swap for a Redis-backed store when you +deploy behind more than a single worker. +""" from __future__ import annotations @@ -31,3 +37,23 @@ def record_failure(self, key: str) -> None: def reset(self, key: str) -> None: self._fails.pop(key, None) self._locks.pop(key, None) + + +class ThroughputLimiter: + """N requests per rolling window per key — counts every attempt. + + Used to dampen enumeration and email-spam on endpoints like + ``/forgot-password`` and ``/register`` where a failure-based lockout + (``LoginRateLimiter``) isn't the right shape — the attacker is after the + side-effect itself, not a correct credential. + """ + + def __init__(self, max_attempts: int = 10, window_seconds: int = 300) -> None: + self._hits: TTLCache = TTLCache(maxsize=10_000, ttl=window_seconds) + self._max = max_attempts + + def check_and_record(self, key: str) -> bool: + """Return True if this attempt is within budget; False if throttled.""" + count = self._hits.get(key, 0) + 1 + self._hits[key] = count + return count <= self._max diff --git a/modules/users/users/roles_cache.py b/modules/users/users/roles_cache.py new file mode 100644 index 00000000..3c4d76c0 --- /dev/null +++ b/modules/users/users/roles_cache.py @@ -0,0 +1,62 @@ +"""Cached list of roles for admin-page rendering. + +Roles are seed data — created by the ``e3ce9754e6dc_seed_users_roles`` +migration and only mutated by hand. Caching on ``app.state`` lets admin views +build Inertia payloads without a per-request ``SELECT * FROM users_role``. + +The cache stores detached :class:`RoleSummary` values (not ORM objects, which +would blow up on attribute access after their session closes). Each view is +responsible for shaping them into whatever the page needs. + +Refresh entry points: + +* ``UsersModule.on_startup`` — initial population. +* ``refresh_roles_cache(app)`` — callable after any future role-management + code paths that add/remove roles. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING + +from sqlalchemy import select + +from users.models import Role + +if TYPE_CHECKING: + from fastapi import FastAPI + +_ROLES_CACHE_KEY = "users_roles_cache" + + +@dataclass(frozen=True, slots=True) +class RoleSummary: + """Minimal role projection safe to hold across request boundaries.""" + + id: str + name: str + + +async def refresh_roles_cache(app: FastAPI) -> list[RoleSummary]: + """Reload the roles list from the DB into ``app.state.users_roles_cache``.""" + async with app.state.db.session_factory() as db: + result = await db.execute(select(Role).order_by(Role.name)) + cached = [RoleSummary(id=str(r.id), name=r.name) for r in result.scalars().all()] + setattr(app.state, _ROLES_CACHE_KEY, cached) + return cached + + +async def get_roles_cache(app: FastAPI) -> list[RoleSummary]: + """Return the cached roles list, populating it from the DB on first miss. + + The cache is pre-warmed in ``UsersModule.on_startup``. The lazy fallback + covers scenarios where startup ran before the ``users_role`` table had any + rows (e.g. tests that use ``metadata.create_all`` instead of running the + seed migration, or staged deployments that seed roles separately). Once + populated, subsequent calls are O(1) attribute reads. + """ + cached = getattr(app.state, _ROLES_CACHE_KEY, None) + if cached: + return cached + return await refresh_roles_cache(app) diff --git a/modules/users/users/service.py b/modules/users/users/service.py index 6cb6d5e0..6d905012 100644 --- a/modules/users/users/service.py +++ b/modules/users/users/service.py @@ -191,7 +191,7 @@ async def generate_reset_link(self, user_id: uuid.UUID, base_url: str) -> str: raise HTTPException(status_code=404, detail="User not found") - token = self._manager.generate_reset_password_token(user) + token = await self._manager.generate_reset_password_token(user) return f"{base_url.rstrip('/')}/users/reset-password?token={token}" async def get_with_roles(self, user_id: uuid.UUID) -> User | None: diff --git a/modules/users/users/settings.py b/modules/users/users/settings.py index 42ffec63..a2a26964 100644 --- a/modules/users/users/settings.py +++ b/modules/users/users/settings.py @@ -2,8 +2,14 @@ from __future__ import annotations -from pydantic import Field +import os + +from pydantic import Field, model_validator from pydantic_settings import BaseSettings, SettingsConfigDict +from simple_module_core.environments import NON_PROD_ENVIRONMENTS + +_PLACEHOLDER_RESET_SECRET = "dev-reset-token-secret-change-me" +_PLACEHOLDER_VERIFY_SECRET = "dev-verify-token-secret-change-me" class UsersSettings(BaseSettings): @@ -38,11 +44,16 @@ class UsersSettings(BaseSettings): smtp_from: str = "no-reply@localhost" smtp_tls: bool = True - # Rate limit (login) + # Rate limit (login — failure-based lockout) login_rate_limit_failures: int = 5 login_rate_limit_window_seconds: int = 300 login_rate_limit_cooldown_seconds: int = 900 + # Rate limit (auth side-effects: forgot-password, register, accept-invite, + # request-verify-token). Counts every attempt per IP per window. + auth_rate_limit_attempts: int = 10 + auth_rate_limit_window_seconds: int = 300 + # Bootstrap (env-var auto-create users on first boot) bootstrap_email: str = "" bootstrap_password: str = "" @@ -50,3 +61,29 @@ class UsersSettings(BaseSettings): # testing non-admin flows without logging out/in repeatedly. bootstrap_user_email: str = "" bootstrap_user_password: str = "" + + @model_validator(mode="after") + def _forbid_placeholder_token_secrets_in_production(self) -> UsersSettings: + """Fail boot if the reset/verify token secrets are still placeholders. + + Both are HMAC keys for fastapi-users JWTs. A well-known default lets + an attacker mint password-reset or email-verification tokens for any + user. The environment is read from ``SM_ENVIRONMENT`` (host setting) + so this check has no runtime coupling to the hosting package. + """ + env = os.environ.get("SM_ENVIRONMENT", "development") + if env in NON_PROD_ENVIRONMENTS: + return self + bad = [] + if self.reset_password_token_secret == _PLACEHOLDER_RESET_SECRET: + bad.append("SM_USERS_RESET_PASSWORD_TOKEN_SECRET") + if self.verification_token_secret == _PLACEHOLDER_VERIFY_SECRET: + bad.append("SM_USERS_VERIFICATION_TOKEN_SECRET") + if bad: + names = ", ".join(bad) + raise ValueError( + f"{names} must be set to non-default value(s) when " + f"SM_ENVIRONMENT={env!r}. Generate with " + "`python -c 'import secrets; print(secrets.token_urlsafe(48))'`." + ) + return self diff --git a/package-lock.json b/package-lock.json index d17e42eb..cdfddd3a 100644 --- a/package-lock.json +++ b/package-lock.json @@ -60,6 +60,7 @@ "@vitejs/plugin-react-swc": "^4.3.0", "autoprefixer": "^10.4.0", "postcss": "^8.4.0", + "rollup-plugin-visualizer": "~5.12.0", "tailwindcss": "^4.0.0", "typescript": "^5.7.0", "vite": "^8.0.8", @@ -5074,7 +5075,6 @@ "integrity": "sha512-quJQXlTSUGL2LH9SUXo8VwsY4soanhgo6LNSm84E1LBcE8s3O0wpdiRzyR9z/ZZJMlMWv37qOOb9pdJlMUEKFQ==", "dev": true, "license": "MIT", - "peer": true, "engines": { "node": ">=8" } @@ -5325,6 +5325,21 @@ "url": "https://polar.sh/cva" } }, + "node_modules/cliui": { + "version": "8.0.1", + "resolved": "https://registry.npmjs.org/cliui/-/cliui-8.0.1.tgz", + "integrity": "sha512-BSeNnyus75C4//NQ9gQt1/csTXyo/8Sb+afLAkzAptFuMsod9HFokGNudZpi/oQV73hnVK+sR+5PVRMd+Dr7YQ==", + "dev": true, + "license": "ISC", + "dependencies": { + "string-width": "^4.2.0", + "strip-ansi": "^6.0.1", + "wrap-ansi": "^7.0.0" + }, + "engines": { + "node": ">=12" + } + }, "node_modules/clsx": { "version": "2.1.1", "resolved": "https://registry.npmjs.org/clsx/-/clsx-2.1.1.tgz", @@ -5350,6 +5365,26 @@ "react-dom": "^18 || ^19 || ^19.0.0-rc" } }, + "node_modules/color-convert": { + "version": "2.0.1", + "resolved": "https://registry.npmjs.org/color-convert/-/color-convert-2.0.1.tgz", + "integrity": "sha512-RRECPsj7iu/xb5oKYcsFHSppFNnsj/52OVTRKb4zP5onXwVF3zVmmToNcOfGC+CRDpfK/U584fMg38ZHCaElKQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "color-name": "~1.1.4" + }, + "engines": { + "node": ">=7.0.0" + } + }, + "node_modules/color-name": { + "version": "1.1.4", + "resolved": "https://registry.npmjs.org/color-name/-/color-name-1.1.4.tgz", + "integrity": "sha512-dOy+3AuW3a2wNbZHIuMZpTcgjGuLU/uBL/ubcZF9OXbDo8ff4O8yVp5Bf0efS8uEoYo5q4Fx7dY9OgQGXgAsQA==", + "dev": true, + "license": "MIT" + }, "node_modules/combined-stream": { "version": "1.0.8", "resolved": "https://registry.npmjs.org/combined-stream/-/combined-stream-1.0.8.tgz", @@ -5589,6 +5624,16 @@ "node": ">=6" } }, + "node_modules/define-lazy-prop": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/define-lazy-prop/-/define-lazy-prop-2.0.0.tgz", + "integrity": "sha512-Ds09qNh8yw3khSjiJjiUInaGX9xlqZDY7JVryGxdxV7NPeuqQfplOpQ66yJFZut3jLa5zOwkXw1g9EI2uKh4Og==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=8" + } + }, "node_modules/delayed-stream": { "version": "1.0.0", "resolved": "https://registry.npmjs.org/delayed-stream/-/delayed-stream-1.0.0.tgz", @@ -5681,6 +5726,13 @@ "embla-carousel": "8.6.0" } }, + "node_modules/emoji-regex": { + "version": "8.0.0", + "resolved": "https://registry.npmjs.org/emoji-regex/-/emoji-regex-8.0.0.tgz", + "integrity": "sha512-MSjYzcWNOA0ewAHpz0MxpYFvwg6yjy1NG3xteoqz644VCo/RPgnr1/GGt+ic3iJTzQ8Eu3TdM14SawnVUmGE6A==", + "dev": true, + "license": "MIT" + }, "node_modules/enhanced-resolve": { "version": "5.20.1", "resolved": "https://registry.npmjs.org/enhanced-resolve/-/enhanced-resolve-5.20.1.tgz", @@ -5898,6 +5950,16 @@ "url": "https://github.com/sponsors/ljharb" } }, + "node_modules/get-caller-file": { + "version": "2.0.5", + "resolved": "https://registry.npmjs.org/get-caller-file/-/get-caller-file-2.0.5.tgz", + "integrity": "sha512-DyFP3BM/3YHTQOCUL/w0OZHR0lpKeGrxotcHWcqNEdnltqFwXVfhEBQ94eIo34AfQpo0rGki4cyIiftY06h2Fg==", + "dev": true, + "license": "ISC", + "engines": { + "node": "6.* || 8.* || >= 10.*" + } + }, "node_modules/get-intrinsic": { "version": "1.3.0", "resolved": "https://registry.npmjs.org/get-intrinsic/-/get-intrinsic-1.3.0.tgz", @@ -6134,6 +6196,32 @@ "node": ">=12" } }, + "node_modules/is-docker": { + "version": "2.2.1", + "resolved": "https://registry.npmjs.org/is-docker/-/is-docker-2.2.1.tgz", + "integrity": "sha512-F+i2BKsFrH66iaUFc0woD8sLy8getkwTwtOBjvs56Cx4CgJDeKQeqfz8wAYiSb8JOprWhHH5p77PbmYCvvUuXQ==", + "dev": true, + "license": "MIT", + "bin": { + "is-docker": "cli.js" + }, + "engines": { + "node": ">=8" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/is-fullwidth-code-point": { + "version": "3.0.0", + "resolved": "https://registry.npmjs.org/is-fullwidth-code-point/-/is-fullwidth-code-point-3.0.0.tgz", + "integrity": "sha512-zymm5+u+sCsSWyD9qNaejV3DFvhCKclKdizYaJUuHA83RLjb7nSuGnddCHGv0hk+KY7BMAlsWeK4Ueg6EV6XQg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=8" + } + }, "node_modules/is-potential-custom-element-name": { "version": "1.0.1", "resolved": "https://registry.npmjs.org/is-potential-custom-element-name/-/is-potential-custom-element-name-1.0.1.tgz", @@ -6141,6 +6229,19 @@ "dev": true, "license": "MIT" }, + "node_modules/is-wsl": { + "version": "2.2.0", + "resolved": "https://registry.npmjs.org/is-wsl/-/is-wsl-2.2.0.tgz", + "integrity": "sha512-fKzAra0rGJUUBwGBgNkHZuToZcn+TtXHpeCgmkMJMMYx1sQDYaCSyjJBSCa2nH1DGm7s3n1oBnohoVTBaN7Lww==", + "dev": true, + "license": "MIT", + "dependencies": { + "is-docker": "^2.0.0" + }, + "engines": { + "node": ">=8" + } + }, "node_modules/jiti": { "version": "2.6.1", "resolved": "https://registry.npmjs.org/jiti/-/jiti-2.6.1.tgz", @@ -6623,6 +6724,24 @@ "url": "https://github.com/sponsors/ljharb" } }, + "node_modules/open": { + "version": "8.4.2", + "resolved": "https://registry.npmjs.org/open/-/open-8.4.2.tgz", + "integrity": "sha512-7x81NCL719oNbsq/3mh+hVrAWmFuEYUqrq/Iw3kUzH8ReypT9QQ0BLoJS7/G9k6N81XjW4qHWtjWwe/9eLy1EQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "define-lazy-prop": "^2.0.0", + "is-docker": "^2.1.1", + "is-wsl": "^2.2.0" + }, + "engines": { + "node": ">=12" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, "node_modules/parse5": { "version": "7.3.0", "resolved": "https://registry.npmjs.org/parse5/-/parse5-7.3.0.tgz", @@ -7138,6 +7257,16 @@ "redux": "^5.0.0" } }, + "node_modules/require-directory": { + "version": "2.1.1", + "resolved": "https://registry.npmjs.org/require-directory/-/require-directory-2.1.1.tgz", + "integrity": "sha512-fGxEI7+wsG9xrvdjsrlmL22OMTTiHRwAMroiEeMgq8gzoLC/PQr7RsRDSTLUg/bZAZtF+TVIkHc6/4RIKrui+Q==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=0.10.0" + } + }, "node_modules/reselect": { "version": "5.1.1", "resolved": "https://registry.npmjs.org/reselect/-/reselect-5.1.1.tgz", @@ -7178,6 +7307,46 @@ "@rolldown/binding-win32-x64-msvc": "1.0.0-rc.15" } }, + "node_modules/rollup-plugin-visualizer": { + "version": "5.12.0", + "resolved": "https://registry.npmjs.org/rollup-plugin-visualizer/-/rollup-plugin-visualizer-5.12.0.tgz", + "integrity": "sha512-8/NU9jXcHRs7Nnj07PF2o4gjxmm9lXIrZ8r175bT9dK8qoLlvKTwRMArRCMgpMGlq8CTLugRvEmyMeMXIU2pNQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "open": "^8.4.0", + "picomatch": "^2.3.1", + "source-map": "^0.7.4", + "yargs": "^17.5.1" + }, + "bin": { + "rollup-plugin-visualizer": "dist/bin/cli.js" + }, + "engines": { + "node": ">=14" + }, + "peerDependencies": { + "rollup": "2.x || 3.x || 4.x" + }, + "peerDependenciesMeta": { + "rollup": { + "optional": true + } + } + }, + "node_modules/rollup-plugin-visualizer/node_modules/picomatch": { + "version": "2.3.2", + "resolved": "https://registry.npmjs.org/picomatch/-/picomatch-2.3.2.tgz", + "integrity": "sha512-V7+vQEJ06Z+c5tSye8S+nHUfI51xoXIXjHQ99cQtKUkQqqO1kO/KCJUfZXuB47h/YBlDhah2H3hdUGXn8ie0oA==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=8.6" + }, + "funding": { + "url": "https://github.com/sponsors/jonschlinkert" + } + }, "node_modules/rrweb-cssom": { "version": "0.7.1", "resolved": "https://registry.npmjs.org/rrweb-cssom/-/rrweb-cssom-0.7.1.tgz", @@ -7304,6 +7473,16 @@ "react-dom": "^18.0.0 || ^19.0.0 || ^19.0.0-rc" } }, + "node_modules/source-map": { + "version": "0.7.6", + "resolved": "https://registry.npmjs.org/source-map/-/source-map-0.7.6.tgz", + "integrity": "sha512-i5uvt8C3ikiWeNZSVZNWcfZPItFQOsYTUAOkcUPGd8DqDy1uOUikjt5dG+uRlwyvR108Fb9DOd4GvXfT0N2/uQ==", + "dev": true, + "license": "BSD-3-Clause", + "engines": { + "node": ">= 12" + } + }, "node_modules/source-map-js": { "version": "1.2.1", "resolved": "https://registry.npmjs.org/source-map-js/-/source-map-js-1.2.1.tgz", @@ -7328,6 +7507,34 @@ "dev": true, "license": "MIT" }, + "node_modules/string-width": { + "version": "4.2.3", + "resolved": "https://registry.npmjs.org/string-width/-/string-width-4.2.3.tgz", + "integrity": "sha512-wKyQRQpjJ0sIp62ErSZdGsjMJWsap5oRNihHhu6G7JVO/9jIB6UyevL+tXuOqrng8j/cxKTWyWUwvSTriiZz/g==", + "dev": true, + "license": "MIT", + "dependencies": { + "emoji-regex": "^8.0.0", + "is-fullwidth-code-point": "^3.0.0", + "strip-ansi": "^6.0.1" + }, + "engines": { + "node": ">=8" + } + }, + "node_modules/strip-ansi": { + "version": "6.0.1", + "resolved": "https://registry.npmjs.org/strip-ansi/-/strip-ansi-6.0.1.tgz", + "integrity": "sha512-Y38VPSHcqkFrCpFnQ9vuSXmquuv5oXOKpGeT6aGrr3o3Gc9AlVa6JBfUSOCnbxGGZF+/0ooI7KrPuUSztUdU5A==", + "dev": true, + "license": "MIT", + "dependencies": { + "ansi-regex": "^5.0.1" + }, + "engines": { + "node": ">=8" + } + }, "node_modules/strip-indent": { "version": "3.0.0", "resolved": "https://registry.npmjs.org/strip-indent/-/strip-indent-3.0.0.tgz", @@ -7954,6 +8161,40 @@ "node": ">=8" } }, + "node_modules/wrap-ansi": { + "version": "7.0.0", + "resolved": "https://registry.npmjs.org/wrap-ansi/-/wrap-ansi-7.0.0.tgz", + "integrity": "sha512-YVGIj2kamLSTxw6NsZjoBxfSwsn0ycdesmc4p+Q21c5zPuZ1pl+NfxVdxPtdHvmNVOQ6XSYG4AUtyt/Fi7D16Q==", + "dev": true, + "license": "MIT", + "dependencies": { + "ansi-styles": "^4.0.0", + "string-width": "^4.1.0", + "strip-ansi": "^6.0.0" + }, + "engines": { + "node": ">=10" + }, + "funding": { + "url": "https://github.com/chalk/wrap-ansi?sponsor=1" + } + }, + "node_modules/wrap-ansi/node_modules/ansi-styles": { + "version": "4.3.0", + "resolved": "https://registry.npmjs.org/ansi-styles/-/ansi-styles-4.3.0.tgz", + "integrity": "sha512-zbB9rCJAT1rbjiVDb2hqKFHNYLxgtk8NURxZ3IZwD3F6NtxbXZQCnnSi1Lkx+IDohdPlFp222wVALIheZJQSEg==", + "dev": true, + "license": "MIT", + "dependencies": { + "color-convert": "^2.0.1" + }, + "engines": { + "node": ">=8" + }, + "funding": { + "url": "https://github.com/chalk/ansi-styles?sponsor=1" + } + }, "node_modules/ws": { "version": "8.20.0", "resolved": "https://registry.npmjs.org/ws/-/ws-8.20.0.tgz", @@ -7993,6 +8234,45 @@ "dev": true, "license": "MIT" }, + "node_modules/y18n": { + "version": "5.0.8", + "resolved": "https://registry.npmjs.org/y18n/-/y18n-5.0.8.tgz", + "integrity": "sha512-0pfFzegeDWJHJIAmTLRP2DwHjdF5s7jo9tuztdQxAhINCdvS+3nGINqPd00AphqJR/0LhANUS6/+7SCb98YOfA==", + "dev": true, + "license": "ISC", + "engines": { + "node": ">=10" + } + }, + "node_modules/yargs": { + "version": "17.7.2", + "resolved": "https://registry.npmjs.org/yargs/-/yargs-17.7.2.tgz", + "integrity": "sha512-7dSzzRQ++CKnNI/krKnYRV7JKKPUXMEh61soaHKg9mrWEhzFWhFnxPxGl+69cD1Ou63C13NUPCnmIcrvqCuM6w==", + "dev": true, + "license": "MIT", + "dependencies": { + "cliui": "^8.0.1", + "escalade": "^3.1.1", + "get-caller-file": "^2.0.5", + "require-directory": "^2.1.1", + "string-width": "^4.2.3", + "y18n": "^5.0.5", + "yargs-parser": "^21.1.1" + }, + "engines": { + "node": ">=12" + } + }, + "node_modules/yargs-parser": { + "version": "21.1.1", + "resolved": "https://registry.npmjs.org/yargs-parser/-/yargs-parser-21.1.1.tgz", + "integrity": "sha512-tVpsJW7DdjecAiFpbIB1e3qxIQsE6NoPc5/eTdrbbIC4h0LVsWhnoa3g+m2HclBIujHzsxZ4VJVA+GUuc2/LBw==", + "dev": true, + "license": "ISC", + "engines": { + "node": ">=12" + } + }, "node_modules/zod": { "version": "4.3.6", "resolved": "https://registry.npmjs.org/zod/-/zod-4.3.6.tgz", diff --git a/packages/ui/src/index.ts b/packages/ui/src/index.ts index e4be644b..ba1350ba 100644 --- a/packages/ui/src/index.ts +++ b/packages/ui/src/index.ts @@ -7,4 +7,5 @@ export { AppLayout } from './layouts/AppLayout'; export { AuthenticatedLayout } from './layouts/AuthenticatedLayout'; export { PublicLayout } from './layouts/PublicLayout'; export { SidebarLayout } from './layouts/SidebarLayout'; +export { fetchWithCsrf, getCsrfToken, setCsrfToken } from './lib/csrf'; export type { MenuItem, SharedProps } from './types'; diff --git a/packages/ui/src/layouts/SidebarLayout.tsx b/packages/ui/src/layouts/SidebarLayout.tsx index 3e11026f..45a1cfb0 100644 --- a/packages/ui/src/layouts/SidebarLayout.tsx +++ b/packages/ui/src/layouts/SidebarLayout.tsx @@ -229,10 +229,15 @@ export function SidebarLayout({ )} {menus?.userDropdown?.map((item) => ( - + {item.label} - + ))} diff --git a/packages/ui/src/lib/csrf.ts b/packages/ui/src/lib/csrf.ts new file mode 100644 index 00000000..0ada1da5 --- /dev/null +++ b/packages/ui/src/lib/csrf.ts @@ -0,0 +1,38 @@ +/** + * CSRF token helpers for non-Inertia fetch() calls. + * + * Inertia visits route through the ``router.on('before', ...)`` hook in + * ``host/client_app/app.tsx`` and are handled there. Anything that uses raw + * ``fetch()`` (e.g. login, forgot-password, admin mutations) must use + * ``fetchWithCsrf`` instead so the server-side CSRFMiddleware accepts the + * request. + * + * The token is session-scoped on the server (see ``CSRFMiddleware``) and + * arrives in the Inertia shared props as ``csrf_token`` on every response. + * ``app.tsx`` calls ``setCsrfToken`` on boot and after each navigation to + * keep this module's copy fresh. + */ + +let csrfToken = ''; + +export function setCsrfToken(token: string): void { + csrfToken = token; +} + +export function getCsrfToken(): string { + return csrfToken; +} + +const SAFE_METHODS = new Set(['GET', 'HEAD', 'OPTIONS']); + +export function fetchWithCsrf(input: RequestInfo | URL, init?: RequestInit): Promise { + const method = (init?.method ?? 'GET').toUpperCase(); + if (SAFE_METHODS.has(method)) { + return fetch(input, init); + } + const headers = new Headers(init?.headers); + if (csrfToken && !headers.has('X-CSRF-Token')) { + headers.set('X-CSRF-Token', csrfToken); + } + return fetch(input, { ...init, headers }); +} diff --git a/packages/ui/src/types.ts b/packages/ui/src/types.ts index 2720b6ce..9a13f712 100644 --- a/packages/ui/src/types.ts +++ b/packages/ui/src/types.ts @@ -2,6 +2,7 @@ export interface MenuItem { label: string; url: string; icon: string; + method?: 'get' | 'post'; } export interface SharedProps { @@ -16,4 +17,5 @@ export interface SharedProps { navbar: MenuItem[]; userDropdown: MenuItem[]; }; + csrf_token?: string; } diff --git a/pyproject.toml b/pyproject.toml index ddf98d74..37feec55 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -73,4 +73,6 @@ exclude = ["framework/hosting/simple_module_hosting/templates/**"] asyncio_mode = "auto" testpaths = ["framework/core/tests", "framework/db/tests", "framework/hosting/tests", "framework/testing/tests", "host/tests", "modules/auth/tests", "modules/dashboard/tests", "modules/products/tests", "modules/users/tests", "tests/integration", "tests/e2e"] markers = ["e2e: end-to-end tests requiring a live browser"] -addopts = "-m 'not e2e'" +# --durations=20 prints the 20 slowest tests after every run so slow fixtures +# (e.g. function-scoped ``app`` rebuilds) surface without needing an opt-in flag. +addopts = "-m 'not e2e' --durations=20" diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 70b8ec93..47bc1601 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -12,22 +12,17 @@ from __future__ import annotations -import json import uuid as _uuid -from base64 import b64encode from collections.abc import AsyncGenerator import httpx import pytest from fastapi_users.password import PasswordHelper -from itsdangerous import TimestampSigner +from simple_module_hosting.csrf import SESSION_CSRF_TOKEN_KEY +from simple_module_testing import forge_session_cookie from sqlalchemy import select - -def _sign_session(session_data: dict, secret: str) -> str: - """Build a signed ``session`` cookie value matching SessionMiddleware.""" - data = b64encode(json.dumps(session_data).encode()) - return TimestampSigner(secret).sign(data).decode("utf-8") +_TEST_CSRF_TOKEN = "test-csrf-token" async def _seed_user_with_roles(app, email: str, role_names: list[str]): @@ -82,12 +77,16 @@ async def _seed_user_with_roles(app, email: str, role_names: list[str]): def _make_client( app, user_id: str, *, extra_headers: dict[str, str] | None = None ) -> httpx.AsyncClient: - cookie = _sign_session({"user_id": user_id}, str(app.state.settings.secret_key)) + cookie = forge_session_cookie( + str(app.state.settings.secret_key), + {"user_id": user_id, SESSION_CSRF_TOKEN_KEY: _TEST_CSRF_TOKEN}, + ) + headers = {"X-CSRF-Token": _TEST_CSRF_TOKEN, **(extra_headers or {})} return httpx.AsyncClient( transport=httpx.ASGITransport(app=app), base_url="http://testserver", cookies={"session": cookie}, - headers=extra_headers or {}, + headers=headers, ) diff --git a/tests/integration/test_i18n_flow.py b/tests/integration/test_i18n_flow.py index 68796e70..436fdcd0 100644 --- a/tests/integration/test_i18n_flow.py +++ b/tests/integration/test_i18n_flow.py @@ -49,9 +49,26 @@ async def app_with_host_routes(app): # type: ignore[no-untyped-def] @pytest.fixture async def host_client(app_with_host_routes) -> AsyncGenerator[httpx.AsyncClient, None]: - """Unauthenticated client against the host-routes-enabled app.""" + """Unauthenticated client against the host-routes-enabled app. + + Pre-seeded with a signed anonymous session carrying a CSRF token so + POST flows (e.g. /i18n/set-locale) clear CSRFMiddleware. + """ + from simple_module_hosting.csrf import SESSION_CSRF_TOKEN_KEY + from simple_module_testing import forge_session_cookie + + csrf = "test-csrf-token" + signed = forge_session_cookie( + str(app_with_host_routes.state.settings.secret_key), + {SESSION_CSRF_TOKEN_KEY: csrf}, + ) transport = httpx.ASGITransport(app=app_with_host_routes) - async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as c: + async with httpx.AsyncClient( + transport=transport, + base_url="http://testserver", + cookies={"session": signed}, + headers={"X-CSRF-Token": csrf}, + ) as c: yield c