From 115f4fd89bc7555874ef81e625ebb5be3aecafc7 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Tue, 15 Sep 2026 10:37:07 +0200 Subject: [PATCH] fix: validate reasoning field before chat dispatch --- .../reasoning_content_utils.py | 13 ++++++ .../llms/hosted_vllm/chat/transformation.py | 25 ++++++----- .../llms/openai/chat/gpt_transformation.py | 13 ++++-- .../test_hosted_vllm_chat_transformation.py | 43 ++++++++++++++++--- 4 files changed, 76 insertions(+), 18 deletions(-) diff --git a/litellm/litellm_core_utils/reasoning_content_utils.py b/litellm/litellm_core_utils/reasoning_content_utils.py index b307a85804e..63ad70f5b27 100644 --- a/litellm/litellm_core_utils/reasoning_content_utils.py +++ b/litellm/litellm_core_utils/reasoning_content_utils.py @@ -6,9 +6,22 @@ from typing import ( cast, # noqa: TID251 # Preserves arbitrary provider fields without lossy TypedDict validation. ) +from litellm.exceptions import BadRequestError from litellm.types.llms.openai import AllMessageValues +def should_normalize_reasoning_content(field: object, *, model: str, provider: str) -> bool: + if field is None or field == "reasoning_content": + return False + if field == "reasoning": + return True + raise BadRequestError( + message="reasoning_content_field must be reasoning_content or reasoning", + model=model, + llm_provider=provider, + ) + + def normalize_reasoning_content( messages: Sequence[AllMessageValues], *, forward: bool = True, normalize: bool = True ) -> list[AllMessageValues]: # mutable-ok: provider request contract diff --git a/litellm/llms/hosted_vllm/chat/transformation.py b/litellm/llms/hosted_vllm/chat/transformation.py index be3c4206f5c..3f084ca6517 100644 --- a/litellm/llms/hosted_vllm/chat/transformation.py +++ b/litellm/llms/hosted_vllm/chat/transformation.py @@ -11,7 +11,10 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( _get_image_mime_type_from_url, ) from litellm.litellm_core_utils.prompt_templates.factory import _parse_mime_type -from litellm.litellm_core_utils.reasoning_content_utils import normalize_reasoning_content +from litellm.litellm_core_utils.reasoning_content_utils import ( + normalize_reasoning_content, + should_normalize_reasoning_content, +) from litellm.litellm_core_utils.reasoning_effort_utils import ( reasoning_effort_from_thinking_budget, ) @@ -151,14 +154,16 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): self, model: str, messages: list[AllMessageValues], # mutable-ok: provider request contract - optional_params: dict, # mutable-ok: provider request contract - litellm_params: dict, # mutable-ok: provider request contract - headers: dict, # mutable-ok: provider request contract - ) -> dict: # mutable-ok: provider request contract + optional_params: dict[str, object], # mutable-ok: provider request contract + litellm_params: dict[str, object], # mutable-ok: provider request contract + headers: dict[str, str], # mutable-ok: provider request contract + ) -> dict[str, object]: # mutable-ok: provider request contract request_messages: Final = normalize_reasoning_content( messages, forward=litellm_params.get("forward_reasoning_content") is True, - normalize=litellm_params.get("reasoning_content_field") == "reasoning", + normalize=should_normalize_reasoning_content( + litellm_params.get("reasoning_content_field"), model=model, provider="hosted_vllm" + ), ) return super().transform_request(model, request_messages, optional_params, litellm_params, headers) @@ -166,10 +171,10 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): self, model: str, messages: list[AllMessageValues], # mutable-ok: provider request contract - optional_params: dict, # mutable-ok: provider request contract - litellm_params: dict, # mutable-ok: provider request contract - headers: dict, # mutable-ok: provider request contract - ) -> dict: # mutable-ok: provider request contract + optional_params: dict[str, object], # mutable-ok: provider request contract + litellm_params: dict[str, object], # mutable-ok: provider request contract + headers: dict[str, str], # mutable-ok: provider request contract + ) -> dict[str, object]: # mutable-ok: provider request contract return await super().async_transform_request( model, deepcopy(messages), optional_params, litellm_params, headers ) diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index b3824ef77a7..fe5586daf7b 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -30,7 +30,10 @@ from litellm.litellm_core_utils.prompt_templates.image_handling import ( async_convert_url_to_base64, convert_url_to_base64, ) -from litellm.litellm_core_utils.reasoning_content_utils import normalize_reasoning_content +from litellm.litellm_core_utils.reasoning_content_utils import ( + normalize_reasoning_content, + should_normalize_reasoning_content, +) from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.llms.base_llm.base_utils import BaseLLMModelInfo from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException @@ -481,7 +484,9 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): request_messages: Final = ( normalize_reasoning_content(messages) if litellm_params.get("custom_llm_provider") == "openai" - and litellm_params.get("reasoning_content_field") == "reasoning" + and should_normalize_reasoning_content( + litellm_params.get("reasoning_content_field"), model=model, provider="openai" + ) else messages ) messages = self._transform_messages(messages=request_messages, model=model) @@ -516,7 +521,9 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): request_messages: Final = ( normalize_reasoning_content(messages) if litellm_params.get("custom_llm_provider") == "openai" - and litellm_params.get("reasoning_content_field") == "reasoning" + and should_normalize_reasoning_content( + litellm_params.get("reasoning_content_field"), model=model, provider="openai" + ) else messages ) transformed_messages = await self._transform_messages(messages=request_messages, model=model, is_async=True) diff --git a/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py b/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py index e17e490b47f..c133f23487c 100644 --- a/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py +++ b/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py @@ -586,7 +586,6 @@ async def test_reasoning_field_sdk_router_final_wire( ("legacy", {}), ("normalized", {"reasoning_content_field": "reasoning"}), ("explicit-default", {"reasoning_content_field": "reasoning_content"}), - ("unknown", {"reasoning_content_field": "unknown"}), ) ], num_retries=0, @@ -604,7 +603,7 @@ async def test_reasoning_field_sdk_router_final_wire( }, ) for alias, field in ( - ("normalized", "reasoning"), ("legacy", None), ("explicit-default", "reasoning_content"), ("unknown", "unknown") + ("normalized", "reasoning"), ("legacy", None), ("explicit-default", "reasoning_content") ): kwargs: Final = ( {"model": alias, "messages": messages} @@ -640,13 +639,14 @@ async def test_reasoning_field_sdk_router_final_wire( assert "reasoning_content_field" not in payload assert "forward_reasoning_content" not in payload assert messages == original - assert route.call_count == 4 + assert route.call_count == 3 @pytest.mark.asyncio @pytest.mark.parametrize("is_async", [False, True]) @pytest.mark.parametrize("provider", ["deepinfra", "together_ai", None]) -async def test_reasoning_field_does_not_apply_to_inherited_provider(provider: str | None, is_async: bool): +@pytest.mark.parametrize("field", ["reasoning", "invalid-selector"]) +async def test_reasoning_field_does_not_apply_to_inherited_provider(provider: str | None, is_async: bool, field: str): from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig config: Final = OpenAIGPTConfig() @@ -657,8 +657,41 @@ async def test_reasoning_field_does_not_apply_to_inherited_provider(provider: st "messages": messages, "optional_params": {}, "headers": {}, - "litellm_params": {"custom_llm_provider": provider, "reasoning_content_field": "reasoning"}, + "litellm_params": {"custom_llm_provider": provider, "reasoning_content_field": field}, } result: Final = await config.async_transform_request(**kwargs) if is_async else config.transform_request(**kwargs) assert result["messages"] == original assert messages == original + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider", ["hosted_vllm", "openai"]) +@pytest.mark.parametrize("is_async", [False, True]) +@pytest.mark.parametrize("via_router", [False, True]) +@pytest.mark.parametrize("bridge", [False, True]) +@pytest.mark.parametrize("field", ["reasonig", ""]) +@pytest.mark.parametrize("forward", [False, True]) +async def test_invalid_reasoning_field_fails_before_http( + provider, is_async, via_router, bridge, field, forward, monkeypatch +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + messages = [{"role": "user", "content": "Hello"}] + original = deepcopy(messages) + params = { + "model": f"{provider}/reasoning-test", "api_key": "test-key", + "api_base": "https://invalid-reasoning-field.invalid/v1", + "reasoning_content_field": field, "forward_reasoning_content": forward, + **({"use_chat_completions_api": True} if bridge else {}), + } + router = litellm.Router(model_list=[{"model_name": "invalid-field", "litellm_params": params}], num_retries=0) + client = router if via_router else litellm + kwargs = {**({"model": "invalid-field"} if via_router else params), "input" if bridge else "messages": messages} + method = (client.aresponses if is_async else client.responses) if bridge else (client.acompletion if is_async else client.completion) + with respx.mock(assert_all_called=False) as mock: + with pytest.raises(litellm.BadRequestError) as error: + await method(**kwargs) if is_async else method(**kwargs) + assert error.value.status_code == 400 + assert "reasoning_content_field must be reasoning_content or reasoning" in str(error.value) + assert "reasonig" not in str(error.value) + assert len(mock.calls) == 0 + assert messages == original