mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat: configure historical assistant reasoning field
This commit is contained in:
parent
8a82ea72ac
commit
9e47c58e1d
13 changed files with 338 additions and 21 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
37
litellm/litellm_core_utils/reasoning_content_utils.py
Normal file
37
litellm/litellm_core_utils/reasoning_content_utils.py
Normal 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
|
||||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue