feat(auth): implement CSRF token protection for state-changing operations

Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
This commit is contained in:
copilot-swe-agent[bot]
2026-02-22 16:27:37 +00:00
parent bb94eccd86
commit c54cbcbdd6
8 changed files with 545 additions and 1 deletions
+1
View File
@@ -80,6 +80,7 @@ async def login(request: Request):
"show_oauth": OAUTH_CONFIGURED, "show_oauth": OAUTH_CONFIGURED,
"oauth_provider_name": OAUTH_PROVIDER_NAME, "oauth_provider_name": OAUTH_PROVIDER_NAME,
"app_version": settings.version, # Changed from app_version to version "app_version": settings.version, # Changed from app_version to version
"csrf_token": getattr(request.state, "csrf_token", ""),
}, },
) )
+6
View File
@@ -19,6 +19,7 @@ from app.auth import router as auth_router
from app.config import settings from app.config import settings
from app.database import init_db from app.database import init_db
from app.middleware.audit_log import AuditLogMiddleware 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.rate_limit import create_limiter, get_rate_limit_exceeded_handler
from app.middleware.request_size_limit import RequestSizeLimitMiddleware from app.middleware.request_size_limit import RequestSizeLimitMiddleware
from app.middleware.security_headers import SecurityHeadersMiddleware from app.middleware.security_headers import SecurityHeadersMiddleware
@@ -123,6 +124,11 @@ app.add_middleware(SecurityHeadersMiddleware, config=settings)
# See SECURITY_AUDIT.md Code Security section # See SECURITY_AUDIT.md Code Security section
app.add_middleware(RequestSizeLimitMiddleware, config=settings) 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 # 2) Audit Logging Middleware - logs all requests with sensitive data masking
# Configure via AUDIT_LOGGING_ENABLED environment variable # Configure via AUDIT_LOGGING_ENABLED environment variable
# See SECURITY_AUDIT.md Infrastructure Security section # See SECURITY_AUDIT.md Infrastructure Security section
+162
View File
@@ -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
+8 -1
View File
@@ -26,12 +26,19 @@ original_template_response = templates.TemplateResponse
def template_response_with_version(*args, **kwargs): 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 context dict is provided, add version to it
if len(args) >= 2 and isinstance(args[1], dict): if len(args) >= 2 and isinstance(args[1], dict):
args[1].setdefault("version", settings.version) 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): elif "context" in kwargs and isinstance(kwargs["context"], dict):
kwargs["context"].setdefault("version", settings.version) 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) return original_template_response(*args, **kwargs)
+34
View File
@@ -1,5 +1,39 @@
// frontend/static/js/common.js // frontend/static/js/common.js
// ---------------------------------------------------------------------------
// CSRF token helper
// ---------------------------------------------------------------------------
// Read the CSRF token from the <meta name="csrf-token"> 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 // Check authentication status and update the auth section
(async function() { (async function() {
console.log('Checking authentication status...'); console.log('Checking authentication status...');
+2
View File
@@ -18,6 +18,8 @@
crossorigin="anonymous" referrerpolicy="no-referrer" /> crossorigin="anonymous" referrerpolicy="no-referrer" />
{% endblock %} {% endblock %}
{% block head_extra %}{% endblock %} {% block head_extra %}{% endblock %}
<!-- CSRF token for AJAX/fetch requests -->
<meta name="csrf-token" content="{{ csrf_token | default('', true) }}">
</head> </head>
<body class="bg-gray-50 min-h-screen flex flex-col"> <body class="bg-gray-50 min-h-screen flex flex-col">
+1
View File
@@ -33,6 +33,7 @@
<div class="mb-8" id="local-auth"> <div class="mb-8" id="local-auth">
<h2 class="text-lg font-semibold mb-4 text-gray-700">Sign in with username</h2> <h2 class="text-lg font-semibold mb-4 text-gray-700">Sign in with username</h2>
<form method="POST" action="/auth" class="space-y-4"> <form method="POST" action="/auth" class="space-y-4">
<input type="hidden" name="csrf_token" value="{{ csrf_token | default('', true) }}">
<div> <div>
<label for="username" class="block text-sm font-medium text-gray-700">Username</label> <label for="username" class="block text-sm font-medium text-gray-700">Username</label>
<input type="text" id="username" name="username" required <input type="text" id="username" name="username" required
+331
View File
@@ -0,0 +1,331 @@
"""Tests for the CSRF protection middleware (app/middleware/csrf.py)."""
import secrets
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import Request
from starlette.responses import JSONResponse, RedirectResponse
from app.middleware.csrf import CSRF_EXEMPT_PATHS, CSRF_PROTECTED_METHODS, CSRFMiddleware
# ---------------------------------------------------------------------------
# Unit tests CSRFMiddleware._get_submitted_token
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestGetSubmittedToken:
"""Unit tests for the CSRF token extraction helper."""
@pytest.mark.asyncio
async def test_returns_header_token(self):
"""Token is read from the X-CSRF-Token request header."""
mock_request = MagicMock(spec=Request)
mock_request.headers = {"X-CSRF-Token": "abc123"}
mock_request.form = AsyncMock(return_value={})
token = await CSRFMiddleware._get_submitted_token(mock_request)
assert token == "abc123"
@pytest.mark.asyncio
async def test_returns_form_token_for_urlencoded(self):
"""Token is read from the form body for URL-encoded POST data."""
mock_request = MagicMock(spec=Request)
mock_request.headers = {"content-type": "application/x-www-form-urlencoded"}
mock_request.form = AsyncMock(return_value={"csrf_token": "form_token_xyz"})
token = await CSRFMiddleware._get_submitted_token(mock_request)
assert token == "form_token_xyz"
@pytest.mark.asyncio
async def test_header_takes_priority_over_form(self):
"""Header token takes priority over form body token."""
mock_request = MagicMock(spec=Request)
mock_request.headers = {
"X-CSRF-Token": "header_token",
"content-type": "application/x-www-form-urlencoded",
}
mock_request.form = AsyncMock(return_value={"csrf_token": "form_token"})
token = await CSRFMiddleware._get_submitted_token(mock_request)
assert token == "header_token"
@pytest.mark.asyncio
async def test_returns_none_when_no_token(self):
"""Returns None when no token is present in header or body."""
mock_request = MagicMock(spec=Request)
mock_request.headers = {"content-type": "application/json"}
mock_request.form = AsyncMock(return_value={})
token = await CSRFMiddleware._get_submitted_token(mock_request)
assert token is None
@pytest.mark.asyncio
async def test_returns_none_for_multipart_without_header(self):
"""Multipart bodies without a header should return None (not parsed)."""
mock_request = MagicMock(spec=Request)
mock_request.headers = {"content-type": "multipart/form-data; boundary=----boundary"}
mock_request.form = AsyncMock(return_value={"csrf_token": "should_not_be_read"})
token = await CSRFMiddleware._get_submitted_token(mock_request)
assert token is None
@pytest.mark.asyncio
async def test_handles_form_parse_exception_gracefully(self):
"""A broken form body does not crash the middleware."""
mock_request = MagicMock(spec=Request)
mock_request.headers = {"content-type": "application/x-www-form-urlencoded"}
mock_request.form = AsyncMock(side_effect=Exception("parse error"))
token = await CSRFMiddleware._get_submitted_token(mock_request)
assert token is None
# ---------------------------------------------------------------------------
# Unit tests CSRFMiddleware.dispatch
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestCSRFMiddlewareDispatch:
"""Unit tests for the CSRFMiddleware.dispatch method."""
def _make_middleware(self, auth_enabled: bool = True):
mock_app = AsyncMock()
mock_config = MagicMock()
mock_config.auth_enabled = auth_enabled
return CSRFMiddleware(mock_app, mock_config)
def _make_request(self, method="GET", path="/", session=None, headers=None, state=None):
mock_request = MagicMock(spec=Request)
mock_request.method = method
mock_request.url = MagicMock()
mock_request.url.path = path
mock_request.session = session if session is not None else {}
mock_request.headers = headers or {}
mock_request.state = MagicMock()
mock_request.state.csrf_token = None
return mock_request
@pytest.mark.asyncio
async def test_noop_when_auth_disabled(self):
"""Middleware is a no-op when AUTH_ENABLED is False."""
middleware = self._make_middleware(auth_enabled=False)
request = self._make_request(method="POST", path="/api/test")
next_response = MagicMock()
call_next = AsyncMock(return_value=next_response)
result = await middleware.dispatch(request, call_next)
call_next.assert_called_once_with(request)
assert result is next_response
@pytest.mark.asyncio
async def test_generates_token_when_not_in_session(self):
"""A new CSRF token is generated and stored in the session when absent."""
middleware = self._make_middleware()
request = self._make_request(method="GET", session={})
call_next = AsyncMock(return_value=MagicMock())
await middleware.dispatch(request, call_next)
assert "csrf_token" in request.session
token = request.session["csrf_token"]
assert len(token) == 64 # secrets.token_hex(32) -> 64 hex chars
@pytest.mark.asyncio
async def test_reuses_existing_token_from_session(self):
"""An existing session token is reused instead of regenerating."""
existing_token = secrets.token_hex(32)
middleware = self._make_middleware()
request = self._make_request(method="GET", session={"csrf_token": existing_token})
call_next = AsyncMock(return_value=MagicMock())
await middleware.dispatch(request, call_next)
assert request.session["csrf_token"] == existing_token
assert request.state.csrf_token == existing_token
@pytest.mark.asyncio
async def test_attaches_token_to_request_state(self):
"""Token is always attached to request.state.csrf_token."""
middleware = self._make_middleware()
request = self._make_request(method="GET", session={})
call_next = AsyncMock(return_value=MagicMock())
await middleware.dispatch(request, call_next)
assert request.state.csrf_token is not None
assert len(request.state.csrf_token) == 64
@pytest.mark.asyncio
async def test_safe_methods_pass_without_token(self):
"""GET/HEAD/OPTIONS requests pass through without CSRF validation."""
middleware = self._make_middleware()
for method in ("GET", "HEAD", "OPTIONS"):
request = self._make_request(method=method, session={})
call_next = AsyncMock(return_value=MagicMock())
result = await middleware.dispatch(request, call_next)
call_next.assert_called_once_with(request)
call_next.reset_mock()
@pytest.mark.asyncio
async def test_post_with_valid_header_token_passes(self):
"""POST with a matching X-CSRF-Token header passes validation."""
token = secrets.token_hex(32)
middleware = self._make_middleware()
request = self._make_request(
method="POST",
path="/api/process/",
session={"csrf_token": token},
headers={"X-CSRF-Token": token},
)
call_next = AsyncMock(return_value=MagicMock())
with patch.object(CSRFMiddleware, "_get_submitted_token", new=AsyncMock(return_value=token)):
result = await middleware.dispatch(request, call_next)
call_next.assert_called_once_with(request)
@pytest.mark.asyncio
async def test_post_with_invalid_token_returns_403_for_api(self):
"""POST with a wrong token on an API route returns HTTP 403."""
token = secrets.token_hex(32)
middleware = self._make_middleware()
request = self._make_request(
method="POST",
path="/api/process/",
session={"csrf_token": token},
)
call_next = AsyncMock(return_value=MagicMock())
with patch.object(CSRFMiddleware, "_get_submitted_token", new=AsyncMock(return_value="wrong_token")):
result = await middleware.dispatch(request, call_next)
assert isinstance(result, JSONResponse)
assert result.status_code == 403
call_next.assert_not_called()
@pytest.mark.asyncio
async def test_post_with_missing_token_returns_403_for_api(self):
"""POST with no CSRF token on an API route returns HTTP 403."""
token = secrets.token_hex(32)
middleware = self._make_middleware()
request = self._make_request(
method="POST",
path="/api/settings/bulk-update",
session={"csrf_token": token},
)
call_next = AsyncMock(return_value=MagicMock())
with patch.object(CSRFMiddleware, "_get_submitted_token", new=AsyncMock(return_value=None)):
result = await middleware.dispatch(request, call_next)
assert isinstance(result, JSONResponse)
assert result.status_code == 403
@pytest.mark.asyncio
async def test_post_with_invalid_token_redirects_for_frontend(self):
"""POST with a wrong token on a frontend route redirects to /login."""
token = secrets.token_hex(32)
middleware = self._make_middleware()
request = self._make_request(
method="POST",
path="/auth",
session={"csrf_token": token},
)
call_next = AsyncMock(return_value=MagicMock())
with patch.object(CSRFMiddleware, "_get_submitted_token", new=AsyncMock(return_value="bad_token")):
result = await middleware.dispatch(request, call_next)
assert isinstance(result, RedirectResponse)
assert result.status_code == 302
assert "/login?error=Invalid+request" in result.headers["location"]
call_next.assert_not_called()
@pytest.mark.asyncio
async def test_delete_with_valid_token_passes(self):
"""DELETE with a matching token passes through."""
token = secrets.token_hex(32)
middleware = self._make_middleware()
request = self._make_request(
method="DELETE",
path="/api/files/1",
session={"csrf_token": token},
)
call_next = AsyncMock(return_value=MagicMock())
with patch.object(CSRFMiddleware, "_get_submitted_token", new=AsyncMock(return_value=token)):
result = await middleware.dispatch(request, call_next)
call_next.assert_called_once_with(request)
@pytest.mark.asyncio
async def test_oauth_callback_is_exempt(self):
"""OAuth callback path is exempt from CSRF validation even on POST."""
middleware = self._make_middleware()
request = self._make_request(
method="POST",
path="/oauth-callback",
session={"csrf_token": secrets.token_hex(32)},
)
call_next = AsyncMock(return_value=MagicMock())
with patch.object(CSRFMiddleware, "_get_submitted_token", new=AsyncMock(return_value=None)):
result = await middleware.dispatch(request, call_next)
call_next.assert_called_once_with(request)
# ---------------------------------------------------------------------------
# Integration tests via TestClient
# ---------------------------------------------------------------------------
@pytest.mark.integration
@pytest.mark.security
class TestCSRFIntegration:
"""Integration tests for CSRF protection using the FastAPI TestClient.
The shared ``client`` fixture runs with ``AUTH_ENABLED=False`` (see
``conftest.py``), so the CSRF middleware is disabled by design. The tests
below verify constants and confirm that the middleware is a no-op in that
configuration.
"""
def test_csrf_constants(self):
"""Verify the constant sets have the expected members."""
assert "POST" in CSRF_PROTECTED_METHODS
assert "PUT" in CSRF_PROTECTED_METHODS
assert "DELETE" in CSRF_PROTECTED_METHODS
assert "PATCH" in CSRF_PROTECTED_METHODS
assert "GET" not in CSRF_PROTECTED_METHODS
assert "/oauth-callback" in CSRF_EXEMPT_PATHS
def test_csrf_middleware_noop_when_auth_disabled(self):
"""When AUTH_ENABLED=False the middleware dispatch is a no-op (no validation)."""
# Build a middleware instance with auth disabled.
mock_app = AsyncMock()
mock_config = MagicMock()
mock_config.auth_enabled = False
middleware = CSRFMiddleware(mock_app, mock_config)
import asyncio
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = MagicMock()
mock_request.url.path = "/api/process/"
mock_request.session = {}
call_next = AsyncMock(return_value=MagicMock())
asyncio.run(middleware.dispatch(mock_request, call_next))
# call_next must have been called (request was not blocked).
call_next.assert_called_once_with(mock_request)
# Session should remain untouched (no token generated).
assert "csrf_token" not in mock_request.session