feat: configure historical assistant reasoning field

This commit is contained in:
jibanez-staticduo 2026-09-15 08:02:35 +02:00
parent 8a82ea72ac
commit 9e47c58e1d
No known key found for this signature in database
13 changed files with 338 additions and 21 deletions

View file

@ -372,6 +372,11 @@ class Cache:
)
if forward_reasoning_content is True:
cache_key += "forward_reasoning_content: True"
reasoning_content_field: Final = kwargs.get(
"reasoning_content_field", nested_litellm_params.get("reasoning_content_field")
)
if reasoning_content_field == "reasoning":
cache_key += "reasoning_content_field: reasoning"
if is_semantic_cache:
cache_key += self._get_semantic_cache_tenant_scope(kwargs)

View file

@ -26,7 +26,9 @@ AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset(
# Keys `completion()` forwards from its own kwargs into `get_litellm_params`,
# which are otherwise invisible to it because that call site passes explicit
# named arguments rather than `**kwargs`.
FORWARDED_KWARGS_KEYS: Final = AWS_CREDENTIAL_KWARGS_KEYS | frozenset({"forward_reasoning_content"})
FORWARDED_KWARGS_KEYS: Final = AWS_CREDENTIAL_KWARGS_KEYS | frozenset(
{"forward_reasoning_content", "reasoning_content_field"}
)
# Pre-define optional kwargs keys as frozenset for O(1) lookups
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls

View file

@ -0,0 +1,37 @@
from collections.abc import Mapping, Sequence
from copy import deepcopy
from types import MappingProxyType
from typing import (
Final,
cast, # noqa: TID251 # Preserves arbitrary provider fields without lossy TypedDict validation.
)
from litellm.types.llms.openai import AllMessageValues
def normalize_reasoning_content(
messages: Sequence[AllMessageValues], *, forward: bool = True
) -> list[AllMessageValues]: # mutable-ok: provider request contract
def normalize_message(message: AllMessageValues) -> AllMessageValues:
if message["role"] != "assistant":
return message
history: Final[Mapping[str, object]] = message
reasoning: Final = (
history.get("reasoning") if history.get("reasoning") is not None else history.get("reasoning_content")
)
normalized: Final[Mapping[str, object]] = MappingProxyType(
{
**MappingProxyType(
{key: value for key, value in history.items() if key not in ("reasoning", "reasoning_content")}
),
**(
MappingProxyType({"reasoning": reasoning})
if forward and reasoning is not None
else MappingProxyType({})
),
}
)
result: Final = dict(normalized) # mutable-ok: provider request contract
return cast(AllMessageValues, result) # cast-ok: only optional reasoning keys change
return [normalize_message(message) for message in deepcopy(messages)] # mutable-ok: provider request contract

View file

@ -11,6 +11,7 @@ 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_effort_utils import (
reasoning_effort_from_thinking_budget,
)
@ -154,7 +155,11 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
litellm_params: dict, # mutable-ok: provider request contract
headers: dict, # mutable-ok: provider request contract
) -> dict: # mutable-ok: provider request contract
request_messages: Final = deepcopy(messages)
request_messages: Final = (
normalize_reasoning_content(messages, forward=litellm_params.get("forward_reasoning_content") is True)
if litellm_params.get("reasoning_content_field") == "reasoning"
else deepcopy(messages)
)
if litellm_params.get("forward_reasoning_content") is not True:
for message in request_messages:
if message["role"] == "assistant":

View file

@ -30,6 +30,7 @@ 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.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
@ -477,7 +478,13 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
Returns:
dict: The transformed request. Sent as the body of the API call.
"""
messages = self._transform_messages(messages=messages, model=model)
request_messages: Final = (
normalize_reasoning_content(messages)
if litellm_params.get("custom_llm_provider") == "openai"
and litellm_params.get("reasoning_content_field") == "reasoning"
else messages
)
messages = self._transform_messages(messages=request_messages, model=model)
if not self._should_preserve_cache_control_for_endpoint(
litellm_params.get("custom_llm_provider"), litellm_params.get("api_base")
):
@ -506,7 +513,13 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
litellm_params: dict,
headers: dict,
) -> dict:
transformed_messages = await self._transform_messages(messages=messages, model=model, is_async=True)
request_messages: Final = (
normalize_reasoning_content(messages)
if litellm_params.get("custom_llm_provider") == "openai"
and litellm_params.get("reasoning_content_field") == "reasoning"
else messages
)
transformed_messages = await self._transform_messages(messages=request_messages, model=model, is_async=True)
if not self._should_preserve_cache_control_for_endpoint(
litellm_params.get("custom_llm_provider"), litellm_params.get("api_base")
):

View file

@ -18968,7 +18968,7 @@
}
}
},
"description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n "
"description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n"
},
"500": {
"content": {

View file

@ -358,6 +358,7 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True)
merge_reasoning_content_in_choices: bool | None = False
forward_reasoning_content: bool | None = False
reasoning_content_field: Literal["reasoning_content", "reasoning"] = "reasoning_content"
model_info: dict | None = None
mock_response: str | ModelResponse | Exception | Any | None = None

View file

@ -3827,6 +3827,7 @@ all_litellm_params = (
"use_in_pass_through",
"merge_reasoning_content_in_choices",
"forward_reasoning_content",
"reasoning_content_field",
"litellm_credential_name",
"allowed_openai_params",
"litellm_session_id",

View file

@ -124,10 +124,14 @@ async def test_router_aliases_isolate_reasoning_flag_on_same_backend(async_mode:
@pytest.mark.asyncio
@pytest.mark.parametrize("async_mode", [False, True], ids=["sync", "async"])
@pytest.mark.parametrize("surface", ["sdk", "router", "responses"])
@pytest.mark.parametrize("surface", ["sdk", "router", "responses", "router-responses"])
@pytest.mark.parametrize("provider", ["hosted_vllm", "openai"])
@pytest.mark.parametrize("normalize", [False, True])
async def test_local_cache_separates_forwarded_reasoning_history(
async_mode: bool, surface: str, monkeypatch: pytest.MonkeyPatch
async_mode: bool, surface: str, provider: str, normalize: bool, monkeypatch: pytest.MonkeyPatch
):
if provider == "openai" and not normalize:
pytest.skip("OpenAI does not apply hosted_vllm forwarding policy")
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setattr(litellm, "cache", Cache(type="local", namespace="reasoning-cache-test"))
messages: Final = _messages()
@ -151,10 +155,22 @@ async def test_local_cache_separates_forwarded_reasoning_history(
{
"model_name": alias,
"litellm_params": {
"model": MODEL,
"model": f"{provider}/reasoning-test",
"api_base": URL.removesuffix("/chat/completions"),
"api_key": "test-key",
**({} if enabled is None else {"forward_reasoning_content": enabled}),
"use_chat_completions_api": True,
**(
{
"forward_reasoning_content": True,
**(
{}
if enabled is None
else {"reasoning_content_field": "reasoning" if enabled else "reasoning_content"}
),
}
if normalize
else ({} if enabled is None else {"forward_reasoning_content": enabled})
),
},
}
for alias, enabled in (("default", None), ("disabled", False), ("enabled", True))
@ -167,7 +183,8 @@ async def test_local_cache_separates_forwarded_reasoning_history(
def backend(request: httpx.Request) -> httpx.Response:
body: Final = json.loads(request.content)
assert "forward_reasoning_content" not in body
enabled: Final = body["messages"][1].get("reasoning_content") == REASONING
assert "reasoning_content_field" not in body
enabled: Final = body["messages"][1].get("reasoning" if normalize else "reasoning_content") == REASONING
return httpx.Response(
200,
json={
@ -196,13 +213,30 @@ async def test_local_cache_separates_forwarded_reasoning_history(
("enabled", True, 2),
):
kwargs: Final = {
"model": MODEL,
"model": f"{provider}/reasoning-test",
"api_base": URL.removesuffix("/chat/completions"),
"api_key": "test-key",
"caching": True,
**({} if forward is None else {"forward_reasoning_content": forward}),
**(
{
"forward_reasoning_content": True,
**(
{}
if forward is None
else {"reasoning_content_field": "reasoning" if forward else "reasoning_content"}
),
}
if normalize
else ({} if forward is None else {"forward_reasoning_content": forward})
),
}
if surface == "router":
if surface == "router-responses":
response = (
await router.aresponses(model=alias, input=input_items)
if async_mode
else router.responses(model=alias, input=input_items)
)
elif surface == "router":
response = (
await router.acompletion(model=alias, messages=messages)
if async_mode
@ -222,7 +256,9 @@ async def test_local_cache_separates_forwarded_reasoning_history(
)
await asyncio.gather(*tuple(_PENDING_CACHE_WRITES))
content: Final = (
response.output[0].content[0].text if surface == "responses" else response.choices[0].message.content
response.output[0].content[0].text
if surface in ("responses", "router-responses")
else response.choices[0].message.content
)
assert content == ("forwarded" if forward is True else "omitted")
assert route.call_count == expected_calls

View file

@ -298,3 +298,42 @@ def test_exact_cache_key_includes_anthropic_messages_params(anthropic_param):
assert baseline != cache.get_cache_key(
model="claude-sonnet-4-5", messages=messages, **anthropic_param
)
@pytest.mark.parametrize("semantic", [False, True])
@pytest.mark.parametrize("provider", ["hosted_vllm", "openai"])
def test_reasoning_field_cache_identity(semantic: bool, provider: str):
cache = _semantic_cache(namespace="history-field") if semantic else Cache(type="local", namespace="history-field")
request = {
"model": f"{provider}/reasoning-test",
"messages": [{"role": "user", "content": "hi"}],
"metadata": {"model_group": "first", "caching_groups": [("first", "second")], "user_api_key": "tenant-a"},
}
legacy = cache.get_cache_key(**request)
assert cache.get_cache_key(**request, reasoning_content_field="reasoning_content") == legacy
assert cache.get_cache_key(**request, litellm_params={"reasoning_content_field": "reasoning_content"}) == legacy
normalized = cache.get_cache_key(**request, reasoning_content_field="reasoning")
assert normalized != legacy
assert normalized == cache.get_cache_key(**request, litellm_params={"reasoning_content_field": "reasoning"})
assert normalized != cache.get_cache_key(
**request, reasoning_content_field="reasoning", forward_reasoning_content=True
)
assert (
cache.get_cache_key(
**request,
reasoning_content_field="reasoning_content",
litellm_params={"reasoning_content_field": "reasoning"},
)
== legacy
)
assert normalized == cache.get_cache_key(
**{**request, "metadata": {**request["metadata"], "model_group": "second"}}, reasoning_content_field="reasoning"
)
assert normalized != cache.get_cache_key(
**request, reasoning_content_field="reasoning", cache={"namespace": "other"}
)
if semantic:
assert normalized != cache.get_cache_key(
**{**request, "metadata": {**request["metadata"], "user_api_key": "tenant-b"}},
reasoning_content_field="reasoning",
)

View file

@ -1,5 +1,11 @@
import json
from copy import deepcopy
from typing import Final
import httpx
import respx
import litellm
from unittest.mock import MagicMock, patch
import pytest
@ -517,3 +523,139 @@ def test_hosted_vllm_custom_tools_use_top_level_input_schema():
assert tools[0]["function"]["name"] == "search"
assert tools[0]["function"]["description"] == "Search docs"
assert tools[0]["function"]["parameters"] == input_schema
@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("forward", [None, False, True])
@pytest.mark.parametrize(
"history, expected",
[
({"reasoning_content": "source"}, "source"),
({"reasoning": "target"}, "target"),
({"reasoning_content": "same", "reasoning": "same"}, "same"),
({"reasoning_content": "source", "reasoning": "target"}, "target"),
({"reasoning_content": "source", "reasoning": None}, "source"),
({"reasoning_content": "source", "reasoning": ""}, ""),
({"reasoning_content": "", "reasoning": None}, ""),
({"reasoning_content": None, "reasoning": None}, None),
({}, None),
],
)
async def test_reasoning_field_sdk_router_final_wire(
provider: str,
is_async: bool,
via_router: bool,
forward: bool | None,
history: dict[str, str | None],
expected: str | None,
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
messages: Final = [
{"role": "user", "content": "Check both records"},
*[
message
for index in (1, 2)
for message in (
{
"role": "assistant",
"content": f"Checking {index}",
**history,
"tool_calls": [
{"id": f"call_{index}", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}
],
},
{"role": "tool", "tool_call_id": f"call_{index}", "content": f"record {index}"},
)
],
]
original: Final = deepcopy(messages)
base_params: Final = {
"model": f"{provider}/reasoning-test",
"api_base": "https://history-field.invalid/v1",
"api_key": "test-key",
**({} if forward is None else {"forward_reasoning_content": forward}),
}
router: Final = litellm.Router(
model_list=[
{"model_name": alias, "litellm_params": {**base_params, **params}}
for alias, params in (
("legacy", {}),
("normalized", {"reasoning_content_field": "reasoning"}),
("explicit-default", {"reasoning_content_field": "reasoning_content"}),
)
],
num_retries=0,
)
with respx.mock(assert_all_called=True) as mock:
route: Final = mock.post("https://history-field.invalid/v1/chat/completions").respond(
200,
json={
"id": "chatcmpl-history",
"object": "chat.completion",
"created": 1,
"model": "reasoning-test",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Done"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 1, "total_tokens": 11},
},
)
for alias, field in (("normalized", "reasoning"), ("legacy", None), ("explicit-default", "reasoning_content")):
kwargs: Final = (
{"model": alias, "messages": messages}
if via_router
else {
**base_params,
"messages": messages,
**({} if field is None else {"reasoning_content_field": field}),
}
)
client: Final = router if via_router else litellm
response: Final = await client.acompletion(**kwargs) if is_async else client.completion(**kwargs)
assert response.choices[0].message.content == "Done"
payload: Final = json.loads(route.calls[-1].request.content)
forwarded: Final = provider == "openai" or forward is True
expected_history: Final = (
({"reasoning": expected} if expected is not None and forwarded else {})
if field == "reasoning"
else {
key: value
for key, value in history.items()
if value is not None and (forwarded or key != "reasoning_content")
}
)
assert payload["messages"] == [
(
{**{key: value for key, value in message.items() if key not in history}, **expected_history}
if message["role"] == "assistant"
else message
)
for message in original
]
assert "reasoning_content_field" not in payload
assert "forward_reasoning_content" not in payload
assert messages == original
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):
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
config: Final = OpenAIGPTConfig()
messages: Final = [{"role": "assistant", "content": "Done", "reasoning_content": "source", "reasoning": "target"}]
original: Final = deepcopy(messages)
kwargs: Final = {
"model": "reasoning-test",
"messages": messages,
"optional_params": {},
"headers": {},
"litellm_params": {"custom_llm_provider": provider, "reasoning_content_field": "reasoning"},
}
result: Final = await config.async_transform_request(**kwargs) if is_async else config.transform_request(**kwargs)
assert result["messages"] == original
assert messages == original

View file

@ -36,8 +36,17 @@ from litellm.types.utils import (
@pytest.mark.parametrize("async_mode", [False, True], ids=["responses", "aresponses"])
@pytest.mark.parametrize("forward", [None, False, True], ids=["absent", "false", "true"])
@pytest.mark.parametrize("sequential", [False, True], ids=["parallel-tools", "sequential-tools"])
@pytest.mark.parametrize("provider", ["hosted_vllm", "openai"])
@pytest.mark.parametrize("field", [None, "reasoning"])
@pytest.mark.parametrize("via_router", [False, True])
async def test_hosted_vllm_responses_reasoning_and_parallel_tools_final_wire(
async_mode: bool, forward: bool | None, sequential: bool, monkeypatch: pytest.MonkeyPatch
async_mode: bool,
forward: bool | None,
sequential: bool,
provider: str,
field: str | None,
via_router: bool,
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
reasoning: Final = "Inspect both tool results before answering."
@ -70,13 +79,27 @@ async def test_hosted_vllm_responses_reasoning_and_parallel_tools_final_wire(
]
original: Final = deepcopy(input_items)
kwargs: Final = {
"model": "hosted_vllm/reasoning-test",
"model": f"{provider}/reasoning-test",
"input": input_items,
"api_base": "https://responses-reasoning-test.invalid/v1",
"api_key": "test-key",
"use_chat_completions_api": True,
**({} if forward is None else {"forward_reasoning_content": forward}),
**({} if field is None else {"reasoning_content_field": field}),
}
router: Final = litellm.Router(
model_list=[
{
"model_name": "history-alias",
"litellm_params": {key: value for key, value in kwargs.items() if key != "input"},
}
],
num_retries=0,
)
client: Final = router if via_router else litellm
request: Final = {"model": "history-alias", "input": input_items} if via_router else kwargs
forwarded: Final = provider == "openai" or forward is True
output_field: Final = field or "reasoning_content"
with respx.mock(assert_all_called=True) as mock:
route: Final = mock.post("https://responses-reasoning-test.invalid/v1/chat/completions").mock(
return_value=httpx.Response(
@ -93,13 +116,14 @@ async def test_hosted_vllm_responses_reasoning_and_parallel_tools_final_wire(
},
)
)
response: Final = await litellm.aresponses(**kwargs) if async_mode else litellm.responses(**kwargs)
response: Final = await client.aresponses(**request) if async_mode else client.responses(**request)
assert response.output[0].content[0].text == "Compared"
assert route.call_count == 1
payload: Final = json.loads(route.calls[0].request.content)
assert payload["model"] == "reasoning-test"
assert "forward_reasoning_content" not in route.calls[0].request.content.decode()
assert "use_chat_completions_api" not in payload
assert "reasoning_content_field" not in payload
messages: Final = payload["messages"]
assert [message["role"] for message in messages] == (
["user", "assistant", "tool", "assistant", "tool"] if sequential else ["user", "assistant", "tool", "tool"]
@ -108,11 +132,11 @@ async def test_hosted_vllm_responses_reasoning_and_parallel_tools_final_wire(
results: Final = [message for message in messages if message["role"] == "tool"]
assert [message["tool_call_id"] for message in results] == ["call_1", "call_2"]
assert [message["content"] for message in results] == ["first record", "second record"]
assert messages[1].get("reasoning_content") == (reasoning if forward is True else None)
assert route.calls[0].request.content.decode().count(reasoning) == int(forward is True)
assert messages[1].get(output_field) == (reasoning if forwarded else None)
assert route.calls[0].request.content.decode().count(reasoning) == int(forwarded)
if sequential:
assert messages[3].get("reasoning_content") == (next_reasoning if forward is True else None)
assert route.calls[0].request.content.decode().count(next_reasoning) == int(forward is True)
assert messages[3].get(output_field) == (next_reasoning if forwarded else None)
assert route.calls[0].request.content.decode().count(next_reasoning) == int(forwarded)
assert input_items == original

View file

@ -29885,6 +29885,12 @@ export interface components {
} | null;
/** Quality Router Default Model */
quality_router_default_model?: string | null;
/**
* Reasoning Content Field
* @default reasoning_content
* @enum {string}
*/
reasoning_content_field: "reasoning_content" | "reasoning";
/** Region Name */
region_name?: string | null;
/** Regional Endpoint Uplift Multiplier */
@ -40104,6 +40110,12 @@ export interface components {
} | null;
/** Quality Router Default Model */
quality_router_default_model?: string | null;
/**
* Reasoning Content Field
* @default reasoning_content
* @enum {string}
*/
reasoning_content_field: "reasoning_content" | "reasoning";
/** Region Name */
region_name?: string | null;
/** Regional Endpoint Uplift Multiplier */