fix: validate reasoning field before chat dispatch

This commit is contained in:
jibanez-staticduo 2026-09-15 10:37:07 +02:00
parent 624f2f97b6
commit 115f4fd89b
No known key found for this signature in database
4 changed files with 76 additions and 18 deletions

View file

@ -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

View file

@ -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
)

View file

@ -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)

View file

@ -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