mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: validate reasoning field before chat dispatch
This commit is contained in:
parent
624f2f97b6
commit
115f4fd89b
4 changed files with 76 additions and 18 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue