"""CSRF middleware — Origin header + CSRF token validation for state-changing requests.""" from __future__ import annotations import logging from fastapi import Request, status from starlette.middleware.base import BaseHTTPMiddleware from starlette.responses import JSONResponse from app.config import get_settings logger = logging.getLogger(__name__) class SecurityHeadersMiddleware(BaseHTTPMiddleware): """Add security headers to all responses.""" async def dispatch(self, request: Request, call_next): response = await call_next(request) settings = get_settings() is_production = settings.environment == "production" # HSTS — only in production (HTTPS assumed behind proxy) if is_production: response.headers["Strict-Transport-Security"] = ( "max-age=63072000; includeSubDomains; preload" ) # Prevent MIME type sniffing response.headers["X-Content-Type-Options"] = "nosniff" # Prevent clickjacking response.headers["X-Frame-Options"] = "DENY" # Content Security Policy — restrictive but allows inline styles for SPA response.headers["Content-Security-Policy"] = ( "default-src 'self'; " "script-src 'self'; " "style-src 'self' 'unsafe-inline' https://fonts.googleapis.com; " "img-src 'self' data: blob:; " "font-src 'self' https://fonts.gstatic.com; " "connect-src 'self' wss: ws:; " "frame-ancestors 'none'; " "base-uri 'self'; " "form-action 'self'" ) # Referrer policy response.headers["Referrer-Policy"] = "strict-origin-when-cross-origin" # Permissions policy response.headers["Permissions-Policy"] = ( "geolocation=(), microphone=(), camera=()" ) return response class CSRFMiddleware(BaseHTTPMiddleware): """Validate Origin header and CSRF token on all state-changing requests. SameSite=Strict cookie + Origin validation + double-submit CSRF token. The CSRF token is generated at login and stored in the Redis session. The client must send it via the X-CSRF-Token header on unsafe methods. """ UNSAFE_METHODS = {"POST", "PATCH", "PUT", "DELETE"} async def dispatch(self, request: Request, call_next): # Skip WebSocket upgrade requests — they use GET and are handled separately if request.headers.get("upgrade", "").lower() == "websocket": return await call_next(request) if request.method in self.UNSAFE_METHODS: # 1. Origin header check origin = request.headers.get("origin") if not origin: return JSONResponse( status_code=status.HTTP_403_FORBIDDEN, content={"detail": "Missing Origin header", "code": "csrf_missing_origin"}, ) settings = get_settings() allowed = settings.cors_origin_list if origin not in allowed: return JSONResponse( status_code=status.HTTP_403_FORBIDDEN, content={"detail": "Invalid Origin", "code": "csrf_invalid_origin"}, ) # 2. CSRF token validation (double-submit pattern) # Skip CSRF token check for auth endpoints (login/password-reset) path = request.url.path if path.endswith("/auth/login") or path.endswith("/auth/logout") or path.endswith("/guest/login") or path.endswith("/guest/logout") or path.endswith("/password-reset/request") or path.endswith("/password-reset/confirm") or path.endswith("/api/v1/errors") or path == "/api/v1/errors": return await call_next(request) csrf_header = request.headers.get("x-csrf-token") if not csrf_header: return JSONResponse( status_code=status.HTTP_403_FORBIDDEN, content={"detail": "Missing X-CSRF-Token header", "code": "csrf_missing_token"}, ) # Get session ID from cookie to look up stored CSRF token session_id = request.cookies.get(settings.session_cookie_name) if not session_id: return JSONResponse( status_code=status.HTTP_403_FORBIDDEN, content={"detail": "No session for CSRF validation", "code": "csrf_no_session"}, ) # Look up CSRF token from session (Redis with DB fallback) from app.core.auth import get_redis, get_session_data redis = get_redis() session_data = await get_session_data(redis, session_id) if session_data is None: return JSONResponse( status_code=status.HTTP_403_FORBIDDEN, content={"detail": "Session expired for CSRF validation", "code": "csrf_session_expired"}, ) stored_token = session_data.get("csrf_token") if not stored_token or stored_token != csrf_header: return JSONResponse( status_code=status.HTTP_403_FORBIDDEN, content={"detail": "CSRF token mismatch", "code": "csrf_token_mismatch"}, ) # Sliding session: also extend TTL on CSRF-validated unsafe requests # (best-effort — ignore Redis errors during outage) try: await redis.expire(f"session:{session_id}", settings.session_ttl_seconds) except Exception: pass return await call_next(request)