diff --git a/app/auth.py b/app/auth.py index 2d5c8c53..3f8a43fa 100644 --- a/app/auth.py +++ b/app/auth.py @@ -80,6 +80,7 @@ async def login(request: Request): "show_oauth": OAUTH_CONFIGURED, "oauth_provider_name": OAUTH_PROVIDER_NAME, "app_version": settings.version, # Changed from app_version to version + "csrf_token": getattr(request.state, "csrf_token", ""), }, ) diff --git a/app/main.py b/app/main.py index 6a0151a2..59d61908 100644 --- a/app/main.py +++ b/app/main.py @@ -19,6 +19,7 @@ from app.auth import router as auth_router from app.config import settings from app.database import init_db from app.middleware.audit_log import AuditLogMiddleware +from app.middleware.csrf import CSRFMiddleware from app.middleware.rate_limit import create_limiter, get_rate_limit_exceeded_handler from app.middleware.request_size_limit import RequestSizeLimitMiddleware from app.middleware.security_headers import SecurityHeadersMiddleware @@ -123,6 +124,11 @@ app.add_middleware(SecurityHeadersMiddleware, config=settings) # See SECURITY_AUDIT.md – Code Security section app.add_middleware(RequestSizeLimitMiddleware, config=settings) +# 3) CSRF Protection Middleware - validates CSRF tokens for state-changing operations +# Only active when AUTH_ENABLED=True. Exempts OAuth callback endpoints. +# Tokens are stored in the session and validated via X-CSRF-Token header or form field. +app.add_middleware(CSRFMiddleware, config=settings) + # 2) Audit Logging Middleware - logs all requests with sensitive data masking # Configure via AUDIT_LOGGING_ENABLED environment variable # See SECURITY_AUDIT.md – Infrastructure Security section diff --git a/app/middleware/csrf.py b/app/middleware/csrf.py new file mode 100644 index 00000000..4494e98b --- /dev/null +++ b/app/middleware/csrf.py @@ -0,0 +1,162 @@ +#!/usr/bin/env python3 + +""" +CSRF Protection Middleware for DocuElevate. + +This middleware implements Cross-Site Request Forgery (CSRF) protection for all +state-changing HTTP operations (POST, PUT, DELETE, PATCH). + +How it works: +- A cryptographically secure token is generated per session and stored in the session. +- The token is attached to ``request.state.csrf_token`` so Jinja2 templates can render it. +- For every state-changing request the middleware validates the submitted token by + checking (in order): + 1. The ``X-CSRF-Token`` HTTP request header (used by AJAX / fetch calls). + 2. The ``csrf_token`` field in ``application/x-www-form-urlencoded`` bodies + (used by traditional HTML forms such as the login form). + Multipart file-upload requests must always supply the token via the header. +- Validation is only enforced when ``AUTH_ENABLED=True``. When authentication is + disabled (development / single-user mode) the middleware is a no-op. + +Exempt paths (CSRF is not checked even for state-changing methods): +- ``/oauth-callback`` – OAuth 2.0 callback; protected by the ``state`` parameter. +""" + +import logging +import secrets +from typing import Callable + +from fastapi import Request, Response +from starlette.middleware.base import BaseHTTPMiddleware +from starlette.responses import JSONResponse, RedirectResponse + +logger = logging.getLogger(__name__) + +# HTTP methods that change server state and therefore require a valid CSRF token. +CSRF_PROTECTED_METHODS = {"POST", "PUT", "DELETE", "PATCH"} + +# Paths that must never be CSRF-checked (e.g. OAuth flow endpoints that carry +# their own replay-protection mechanism). +CSRF_EXEMPT_PATHS = { + "/oauth-callback", +} + + +class CSRFMiddleware(BaseHTTPMiddleware): + """ + Middleware that generates and validates CSRF tokens for state-changing requests. + + Token lifecycle + --------------- + 1. On the first request for a session a 64-character hex token is created with + ``secrets.token_hex(32)`` and stored in ``request.session["csrf_token"]``. + 2. On every subsequent request the existing token is read from the session. + 3. The token is always attached to ``request.state.csrf_token`` so that + Jinja2 templates (and response processors) can embed it. + + Validation + ---------- + For ``POST``, ``PUT``, ``DELETE``, and ``PATCH`` requests the middleware + checks whether the submitted token matches the session token using a + constant-time comparison (``secrets.compare_digest``) to prevent timing + attacks. + + Failure response + ---------------- + - API routes (``/api/*``): HTTP 403 JSON response. + - All other routes: HTTP 302 redirect to ``/login?error=…``. + """ + + def __init__(self, app, config): + """ + Initialise the middleware. + + Args: + app: The ASGI application. + config: Application settings object (``app.config.Settings``). + ``config.auth_enabled`` controls whether CSRF enforcement is active. + """ + super().__init__(app) + self.config = config + self.enabled = config.auth_enabled + if self.enabled: + logger.info("CSRF protection middleware enabled") + else: + logger.info("CSRF protection middleware disabled (AUTH_ENABLED=False)") + + async def dispatch(self, request: Request, call_next: Callable) -> Response: + """ + Process the request: generate/attach the token and validate it when required. + + Args: + request: Incoming HTTP request. + call_next: Next middleware or route handler in the ASGI chain. + + Returns: + HTTP response, or an error response when CSRF validation fails. + """ + if not self.enabled: + return await call_next(request) + + # Generate or retrieve the per-session CSRF token. + csrf_token = request.session.get("csrf_token") + if not csrf_token: + csrf_token = secrets.token_hex(32) + request.session["csrf_token"] = csrf_token + + # Attach token to request state so templates and route handlers can access it. + request.state.csrf_token = csrf_token + + # Validate for state-changing methods on non-exempt paths. + if request.method in CSRF_PROTECTED_METHODS and request.url.path not in CSRF_EXEMPT_PATHS: + submitted_token = await self._get_submitted_token(request) + if not submitted_token or not secrets.compare_digest(csrf_token, submitted_token): + logger.warning( + f"[SECURITY] CSRF_VALIDATION_FAILED method={request.method} path={request.url.path}" + ) + if request.url.path.startswith("/api/"): + return JSONResponse( + status_code=403, + content={"detail": "CSRF token missing or invalid"}, + ) + return RedirectResponse(url="/login?error=Invalid+request", status_code=302) + + return await call_next(request) + + @staticmethod + async def _get_submitted_token(request: Request) -> str | None: + """ + Extract the CSRF token submitted by the client. + + Checks (in priority order): + 1. ``X-CSRF-Token`` request header – preferred for AJAX / fetch requests. + 2. ``csrf_token`` form field in ``application/x-www-form-urlencoded`` bodies – + used by plain HTML forms (e.g. the login form). + + Multipart bodies (file uploads) are intentionally not parsed here to avoid + buffering large uploads in middleware; those endpoints must send the token + via the header instead. + + Args: + request: The incoming HTTP request. + + Returns: + The submitted CSRF token string, or ``None`` if not found. + """ + # 1. Check the request header (AJAX / fetch). + token = request.headers.get("X-CSRF-Token") + if token: + return token + + # 2. For URL-encoded form bodies only (plain HTML form submissions). + content_type = request.headers.get("content-type", "") + if "application/x-www-form-urlencoded" in content_type: + try: + form = await request.form() + token = form.get("csrf_token") + if token: + return str(token) + except Exception as exc: + logger.debug(f"CSRF: could not parse form body: {exc}") + + return None diff --git a/app/views/base.py b/app/views/base.py index 45ffa7b9..c3b338a7 100644 --- a/app/views/base.py +++ b/app/views/base.py @@ -26,12 +26,19 @@ original_template_response = templates.TemplateResponse def template_response_with_version(*args, **kwargs): - """Wrapper for TemplateResponse to include version in all templates""" + """Wrapper for TemplateResponse to include version and CSRF token in all templates""" # If context dict is provided, add version to it if len(args) >= 2 and isinstance(args[1], dict): args[1].setdefault("version", settings.version) + # Inject CSRF token from request state when available + req = args[1].get("request") + if req is not None and hasattr(req.state, "csrf_token"): + args[1].setdefault("csrf_token", req.state.csrf_token) elif "context" in kwargs and isinstance(kwargs["context"], dict): kwargs["context"].setdefault("version", settings.version) + req = kwargs["context"].get("request") + if req is not None and hasattr(req.state, "csrf_token"): + kwargs["context"].setdefault("csrf_token", req.state.csrf_token) return original_template_response(*args, **kwargs) diff --git a/frontend/static/js/common.js b/frontend/static/js/common.js index 0d906055..07c15e09 100644 --- a/frontend/static/js/common.js +++ b/frontend/static/js/common.js @@ -1,5 +1,39 @@ // frontend/static/js/common.js +// --------------------------------------------------------------------------- +// CSRF token helper +// --------------------------------------------------------------------------- +// Read the CSRF token from the tag injected by the +// server into base.html for every authenticated page. +function getCsrfToken() { + const meta = document.querySelector('meta[name="csrf-token"]'); + return meta ? meta.getAttribute('content') : ''; +} + +// Wrap the native fetch() so that every state-changing request automatically +// includes the X-CSRF-Token header without requiring callers to remember it. +(function patchFetch() { + const _CSRF_METHODS = new Set(['POST', 'PUT', 'DELETE', 'PATCH']); + const _originalFetch = window.fetch; + window.fetch = function (input, init) { + init = init || {}; + const method = (init.method || 'GET').toUpperCase(); + if (_CSRF_METHODS.has(method)) { + const token = getCsrfToken(); + if (token) { + // Merge headers so a caller-supplied X-CSRF-Token is not overwritten, + // but add the token when no override is present. + const headers = Object.assign({}, init.headers || {}); + if (!headers['X-CSRF-Token']) { + headers['X-CSRF-Token'] = token; + } + init.headers = headers; + } + } + return _originalFetch.call(this, input, init); + }; +})(); + // Check authentication status and update the auth section (async function() { console.log('Checking authentication status...'); diff --git a/frontend/templates/base.html b/frontend/templates/base.html index 4a8715f9..aa0c196f 100644 --- a/frontend/templates/base.html +++ b/frontend/templates/base.html @@ -18,6 +18,8 @@ crossorigin="anonymous" referrerpolicy="no-referrer" /> {% endblock %} {% block head_extra %}{% endblock %} + +
diff --git a/frontend/templates/login.html b/frontend/templates/login.html index c1c6060d..7087120a 100644 --- a/frontend/templates/login.html +++ b/frontend/templates/login.html @@ -33,6 +33,7 @@