fix(auth): cache request body in CSRF middleware to prevent login failures
Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
This commit is contained in:
@@ -151,6 +151,16 @@ class CSRFMiddleware(BaseHTTPMiddleware):
|
|||||||
logger.debug("CSRF: content_type=%r method=%s path=%s", content_type, request.method, request.url.path)
|
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:
|
if "application/x-www-form-urlencoded" in content_type:
|
||||||
try:
|
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()
|
form = await request.form()
|
||||||
token = form.get("csrf_token")
|
token = form.get("csrf_token")
|
||||||
logger.debug(
|
logger.debug(
|
||||||
|
|||||||
+58
-1
@@ -33,10 +33,13 @@ class TestGetSubmittedToken:
|
|||||||
"""Token is read from the form body for URL-encoded POST data."""
|
"""Token is read from the form body for URL-encoded POST data."""
|
||||||
mock_request = MagicMock(spec=Request)
|
mock_request = MagicMock(spec=Request)
|
||||||
mock_request.headers = {"content-type": "application/x-www-form-urlencoded"}
|
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"})
|
mock_request.form = AsyncMock(return_value={"csrf_token": "form_token_xyz"})
|
||||||
|
|
||||||
token = await CSRFMiddleware._get_submitted_token(mock_request)
|
token = await CSRFMiddleware._get_submitted_token(mock_request)
|
||||||
assert token == "form_token_xyz"
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_header_takes_priority_over_form(self):
|
async def test_header_takes_priority_over_form(self):
|
||||||
@@ -71,18 +74,72 @@ class TestGetSubmittedToken:
|
|||||||
token = await CSRFMiddleware._get_submitted_token(mock_request)
|
token = await CSRFMiddleware._get_submitted_token(mock_request)
|
||||||
assert token is None
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_handles_form_parse_exception_gracefully(self):
|
async def test_handles_form_parse_exception_gracefully(self):
|
||||||
"""A broken form body does not crash the middleware."""
|
"""A broken form body does not crash the middleware."""
|
||||||
mock_request = MagicMock(spec=Request)
|
mock_request = MagicMock(spec=Request)
|
||||||
mock_request.headers = {"content-type": "application/x-www-form-urlencoded"}
|
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"))
|
mock_request.form = AsyncMock(side_effect=Exception("parse error"))
|
||||||
|
|
||||||
token = await CSRFMiddleware._get_submitted_token(mock_request)
|
token = await CSRFMiddleware._get_submitted_token(mock_request)
|
||||||
assert token is None
|
assert token is None
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Unit tests – CSRFMiddleware.dispatch
|
# Unit tests – CSRFMiddleware.dispatch
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user