fix: add .copy() to create_request_copy and tests for header caching

Protect cached headers from mutation in create_request_copy sites and
add unit tests for _safe_get_request_headers caching behavior.
This commit is contained in:
Ryan Crabbe 2026-02-18 09:44:35 -08:00 committed by Sameer Kankute
parent 97d3942c78
commit 071a7bbad1
3 changed files with 48 additions and 2 deletions

View file

@ -61,7 +61,7 @@ def create_request_copy(request: Request):
return {
"method": request.method,
"url": str(request.url),
"headers": _safe_get_request_headers(request),
"headers": _safe_get_request_headers(request).copy(),
"cookies": request.cookies,
"query_params": dict(request.query_params),
}

View file

@ -33,7 +33,7 @@ def create_request_copy(request: Request):
return {
"method": request.method,
"url": str(request.url),
"headers": _safe_get_request_headers(request),
"headers": _safe_get_request_headers(request).copy(),
"cookies": request.cookies,
"query_params": dict(request.query_params),
}

View file

@ -761,3 +761,49 @@ async def test_request_body_with_html_script_tags():
f"Message content with HTML was modified during parsing: "
f"expected={msg['content']!r}, got={result['messages'][2]['content']!r}"
)
def test_safe_get_request_headers_caches_on_request_state():
"""
Test that _safe_get_request_headers caches the result on request.state
and returns the same object on subsequent calls.
"""
mock_request = MagicMock()
mock_request.headers = {"content-type": "application/json", "authorization": "Bearer sk-123"}
mock_request.state = MagicMock(spec=[]) # empty spec so getattr returns default
# First call — should create and cache
result1 = _safe_get_request_headers(mock_request)
assert result1 == {"content-type": "application/json", "authorization": "Bearer sk-123"}
assert mock_request.state._cached_headers is result1
# Second call — should return the cached object (same identity)
result2 = _safe_get_request_headers(mock_request)
assert result2 is result1
def test_safe_get_request_headers_none_request():
"""
Test that _safe_get_request_headers returns empty dict for None request.
"""
result = _safe_get_request_headers(None)
assert result == {}
def test_safe_get_request_headers_copy_protects_cache():
"""
Test that callers using .copy() before mutation do not corrupt the cache.
"""
mock_request = MagicMock()
mock_request.headers = {"authorization": "Bearer sk-123", "host": "localhost"}
mock_request.state = MagicMock(spec=[])
original = _safe_get_request_headers(mock_request)
# Simulate what mutation call sites do: copy then pop
mutable = _safe_get_request_headers(mock_request).copy()
mutable.pop("authorization", None)
# Cache must be unaffected
assert "authorization" in _safe_get_request_headers(mock_request)
assert _safe_get_request_headers(mock_request) is original