From dcde79b6b0aa7a8bb51a8da005268186e6e44c09 Mon Sep 17 00:00:00 2001 From: Charan Rathore Date: Mon, 28 Sep 2026 10:19:32 +0530 Subject: [PATCH] fix: require JSON object for adapter chat passthrough --- .../pass_through_endpoints.py | 11 ++-- .../test_pass_through_endpoints.py | 60 +++++++++++++++++++ 2 files changed, 66 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e0a4184291e..a7a6ea6863a 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -1,4 +1,3 @@ -import ast import asyncio import copy import json @@ -220,11 +219,13 @@ async def chat_completion_pass_through_endpoint( data = {"litellm_call_id": litellm_call_id} try: body: Final = await request.body() - body_str: Final = body.decode() try: - data = ast.literal_eval(body_str) | data - except Exception: - data = json.loads(body_str) | data + parsed_body: Final = json.loads(body) + except (UnicodeDecodeError, json.JSONDecodeError) as exc: + raise HTTPException(status_code=400, detail="Request body must be valid JSON") from exc + if not isinstance(parsed_body, dict): + raise HTTPException(status_code=400, detail="Request body must be a JSON object") + data = parsed_body | data data["adapter_id"] = adapter_id diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 3469df082e0..ade94276604 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -7481,6 +7481,66 @@ def test_passthrough_sees_the_public_list_rebound_after_import(monkeypatch: pyte assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]} +@pytest.mark.asyncio +@pytest.mark.parametrize( + "body", + ( + b"{'model': 'gpt-4o-mini'}", + b'(1, 2, {"role": "user"})', + b"[1, 2]", + b'{"model":', + ), +) +async def test_adapter_chat_rejects_non_object_or_non_json_body(monkeypatch: pytest.MonkeyPatch, body: bytes): + proxy_logging = MagicMock() + proxy_logging.post_call_failure_hook = AsyncMock() + add_request_data = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging) + monkeypatch.setattr("litellm.proxy.proxy_server.add_litellm_data_to_request", add_request_data) + request = MagicMock(spec=Request) + request.body = AsyncMock(return_value=body) + + with pytest.raises(ProxyException) as raised: + await chat_completion_pass_through_endpoint( + fastapi_response=Response(), + request=request, + adapter_id="anthropic", + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + + assert (raised.value.type, raised.value.code) == ("invalid_request_error", "400") + add_request_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_adapter_chat_accepts_json_object(monkeypatch: pytest.MonkeyPatch): + proxy_logging = MagicMock() + proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["data"]) + proxy_logging.post_call_failure_hook = AsyncMock() + + async def add_request_data(**kwargs: object) -> object: + return kwargs["data"] + + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging) + monkeypatch.setattr("litellm.proxy.proxy_server.add_litellm_data_to_request", add_request_data) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + request = MagicMock(spec=Request) + request.body = AsyncMock(return_value=b'{"model":"unknown-model","messages":[]}') + + with pytest.raises(ProxyException) as raised: + await chat_completion_pass_through_endpoint( + fastapi_response=Response(), + request=request, + adapter_id="anthropic", + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + + assert (raised.value.type, raised.value.code) == ("invalid_request_error", "400") + assert proxy_logging.pre_call_hook.await_args.kwargs["data"]["model"] == "unknown-model" + + @pytest.mark.asyncio async def test_chat_completion_pass_through_endpoint_answers_an_openai_typed_error_for_an_unknown_model( monkeypatch: pytest.MonkeyPatch,