diff --git a/app/middleware/csrf.py b/app/middleware/csrf.py index 555c3aa8..dce8dda1 100644 --- a/app/middleware/csrf.py +++ b/app/middleware/csrf.py @@ -151,6 +151,16 @@ class CSRFMiddleware(BaseHTTPMiddleware): logger.debug("CSRF: content_type=%r method=%s path=%s", content_type, request.method, request.url.path) if "application/x-www-form-urlencoded" in content_type: try: + # Cache the raw body bytes before parsing the form. Starlette's + # BaseHTTPMiddleware uses _CachedRequest.wrapped_receive to relay + # the body to downstream handlers. When form() is called it + # internally uses stream() which sets _stream_consumed=True but + # does NOT populate _body. wrapped_receive then sees a consumed + # stream and forwards an empty body, so the auth endpoint gets + # form_keys=[]. Calling body() first stores the bytes in _body; + # wrapped_receive detects this and replays the real body to any + # downstream handler (e.g. the /auth endpoint). + await request.body() form = await request.form() token = form.get("csrf_token") logger.debug( diff --git a/tests/test_csrf.py b/tests/test_csrf.py index 31745f9b..a6041038 100644 --- a/tests/test_csrf.py +++ b/tests/test_csrf.py @@ -33,10 +33,13 @@ class TestGetSubmittedToken: """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.body = AsyncMock(return_value=b"csrf_token=form_token_xyz") mock_request.form = AsyncMock(return_value={"csrf_token": "form_token_xyz"}) token = await CSRFMiddleware._get_submitted_token(mock_request) assert token == "form_token_xyz" + # body() must have been awaited so the body is cached for downstream re-reads. + mock_request.body.assert_awaited_once() @pytest.mark.asyncio async def test_header_takes_priority_over_form(self): @@ -71,18 +74,72 @@ class TestGetSubmittedToken: token = await CSRFMiddleware._get_submitted_token(mock_request) assert token is None + @pytest.mark.asyncio + async def test_body_is_cached_before_form_parse(self): + """body() is called before form() so downstream handlers can re-read the body. + + This covers the Starlette BaseHTTPMiddleware body-replay bug: if form() + is called without first calling body(), _stream_consumed is set to True + but _body remains unset. _CachedRequest.wrapped_receive then forwards + an empty body to downstream apps (e.g. the /auth endpoint) causing + form_keys=[] and login failures. Calling body() first caches _body so + wrapped_receive correctly replays the full body. + """ + mock_request = MagicMock(spec=Request) + mock_request.headers = {"content-type": "application/x-www-form-urlencoded"} + call_order: list[str] = [] + + async def _body(): + call_order.append("body") + return b"csrf_token=tok&username=alice&password=test_password" + + async def _form(): + call_order.append("form") + return {"csrf_token": "tok", "username": "alice", "password": "test_password"} + + mock_request.body = _body + mock_request.form = _form + + token = await CSRFMiddleware._get_submitted_token(mock_request) + + assert token == "tok" + # body() must be called BEFORE form() to ensure body caching. + assert call_order == ["body", "form"] + + @pytest.mark.asyncio + async def test_body_is_not_called_when_header_present(self): + """When the CSRF token is in the X-CSRF-Token header, body() is not called. + + For header-based token submission (AJAX / fetch requests) we skip body + parsing entirely, so the body stream remains unconsumed and downstream + handlers can read it normally. + """ + mock_request = MagicMock(spec=Request) + mock_request.headers = { + "X-CSRF-Token": "header_tok", + "content-type": "application/x-www-form-urlencoded", + } + mock_request.body = AsyncMock(return_value=b"username=alice") + mock_request.form = AsyncMock(return_value={"csrf_token": "header_tok"}) + + token = await CSRFMiddleware._get_submitted_token(mock_request) + + assert token == "header_tok" + # body() must NOT be called – the header path returns early. + mock_request.body.assert_not_awaited() + @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.body = AsyncMock(return_value=b"") 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 # ---------------------------------------------------------------------------