mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-23 00:41:40 +00:00
Test pass_through_request state_raw_body uses httpx content=
Cover non-streaming (async_client.request) and streaming (build_request) paths so SigV4 bytes on request.state are not replaced by json= of a hook-mutated dict. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
c2ad1fc2e9
commit
9d6a19b4a0
1 changed files with 138 additions and 0 deletions
|
|
@ -20,6 +20,9 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
|||
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
|
||||
pass_through_request,
|
||||
)
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
||||
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.success_handler import (
|
||||
PassThroughEndpointLogging,
|
||||
)
|
||||
|
|
@ -2226,6 +2229,141 @@ async def test_create_pass_through_route_custom_body_url_target():
|
|||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_non_streaming_uses_content_for_state_raw_body():
|
||||
"""
|
||||
Bedrock SigV4 path: exact signed bytes live on request.state; upstream must receive
|
||||
content=... even if pre_call_hook mutates the parsed dict (would change json=).
|
||||
"""
|
||||
# Bytes that were signed (simulated); parsed body + hook will diverge on purpose.
|
||||
raw_signed = b'{"retrievalQuery":{"text":"signed"},"sig":"intact"}'
|
||||
parsed_from_wire = {"retrievalQuery": {"text": "signed"}, "sig": "intact"}
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.query_params = QueryParams({})
|
||||
mock_request.headers = Headers({"Content-Type": "application/json"})
|
||||
mock_request.state = SimpleNamespace()
|
||||
setattr(mock_request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, raw_signed)
|
||||
mock_request.body = AsyncMock(
|
||||
return_value=json.dumps(parsed_from_wire).encode("utf-8")
|
||||
)
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.api_key = "sk-test"
|
||||
|
||||
upstream = httpx.Response(
|
||||
status_code=200,
|
||||
headers={"content-type": "application/json"},
|
||||
content=b'{"ok": true}',
|
||||
request=httpx.Request(
|
||||
"POST",
|
||||
"https://bedrock-agent-runtime.us-east-1.amazonaws.com/knowledgebases/KB/retrieve",
|
||||
),
|
||||
)
|
||||
|
||||
mock_async_client = AsyncMock()
|
||||
mock_async_client.request = AsyncMock(return_value=upstream)
|
||||
mock_client_obj = MagicMock()
|
||||
mock_client_obj.client = mock_async_client
|
||||
|
||||
async def _hook_mutates_body(**kwargs):
|
||||
data = kwargs["data"]
|
||||
if isinstance(data, dict):
|
||||
data["hook_mutated"] = True
|
||||
return data
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client",
|
||||
return_value=mock_client_obj,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj.pre_call_hook",
|
||||
new=AsyncMock(side_effect=_hook_mutates_body),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler",
|
||||
new=AsyncMock(),
|
||||
),
|
||||
):
|
||||
await pass_through_request(
|
||||
request=mock_request,
|
||||
target="https://bedrock-agent-runtime.us-east-1.amazonaws.com/knowledgebases/KB/retrieve",
|
||||
custom_headers={"content-type": "application/json"},
|
||||
user_api_key_dict=mock_user,
|
||||
stream=False,
|
||||
)
|
||||
|
||||
mock_async_client.request.assert_called_once()
|
||||
req_kw = mock_async_client.request.call_args[1]
|
||||
assert req_kw.get("content") == raw_signed
|
||||
assert "json" not in req_kw
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_streaming_uses_content_for_state_raw_body():
|
||||
"""Streaming pass-through with state raw body must use build_request(..., content=...)."""
|
||||
raw_signed = b'{"model":"m","stream":true}'
|
||||
parsed_from_wire = {"model": "m", "stream": True}
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.query_params = QueryParams({})
|
||||
mock_request.headers = Headers({"Content-Type": "application/json"})
|
||||
mock_request.state = SimpleNamespace()
|
||||
setattr(mock_request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, raw_signed)
|
||||
mock_request.body = AsyncMock(
|
||||
return_value=json.dumps(parsed_from_wire).encode("utf-8")
|
||||
)
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.api_key = "sk-test"
|
||||
|
||||
mock_built = MagicMock()
|
||||
mock_async_client = AsyncMock()
|
||||
mock_async_client.build_request = MagicMock(return_value=mock_built)
|
||||
stream_resp = httpx.Response(
|
||||
status_code=200,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
content=b"data: {}\n\n",
|
||||
request=httpx.Request("POST", "https://example.com/v1/messages"),
|
||||
)
|
||||
mock_async_client.send = AsyncMock(return_value=stream_resp)
|
||||
mock_client_obj = MagicMock()
|
||||
mock_client_obj.client = mock_async_client
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client",
|
||||
return_value=mock_client_obj,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj.pre_call_hook",
|
||||
new=AsyncMock(side_effect=lambda **kw: kw["data"]),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler",
|
||||
new=AsyncMock(),
|
||||
),
|
||||
):
|
||||
response = await pass_through_request(
|
||||
request=mock_request,
|
||||
target="https://example.com/v1/messages",
|
||||
custom_headers={"Authorization": "Bearer x"},
|
||||
user_api_key_dict=mock_user,
|
||||
stream=None,
|
||||
)
|
||||
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
mock_async_client.build_request.assert_called_once()
|
||||
br_kw = mock_async_client.build_request.call_args[1]
|
||||
assert br_kw.get("content") == raw_signed
|
||||
assert "json" not in br_kw
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_pass_through_route_no_custom_body_falls_back():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue