This commit is contained in:
Jordi Ibáñez 2026-10-01 10:57:55 +02:00 • committed by GitHub
commit 69fd88afb8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
26 changed files with 1387 additions and 140 deletions

View file

@ -60,6 +60,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
# Tools / agents (registry & policy admin)
"/v1/tool/",
"/v1/agents",
"/v1/traces",
# Guardrails admin
"/v2/guardrails/",
# MCP server admin + BYOK OAuth flow (UI-initiated) + dynamic per-server endpoints

View file

@ -18,7 +18,7 @@ from enum import Enum
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
from pydantic import BaseModel
from pydantic import BaseModel, TypeAdapter
import litellm
from litellm._logging import verbose_logger
@ -59,6 +59,9 @@ def _native_response(result: object) -> object:
return result
_LITELLM_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
def print_verbose(print_statement):
try:
verbose_logger.debug(print_statement)
@ -390,6 +393,20 @@ class Cache:
param_value = kwargs[param]
cache_key += f"{param}: {param_value}"
nested_litellm_params: Final = _LITELLM_PARAMS_ADAPTER.validate_python(
kwargs.get("litellm_params") or MappingProxyType({})
)
forward_reasoning_content: Final = kwargs.get(
"forward_reasoning_content", nested_litellm_params.get("forward_reasoning_content")
)
if forward_reasoning_content is False:
cache_key += "forward_reasoning_content: False"
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

@ -32,6 +32,8 @@ AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset(
PROVIDER_AFFINITY_HEADER_KWARG_KEY: Final = "provider_affinity_header"
REASONING_TRANSPORT_KWARGS_KEYS: Final = 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
OPTIONAL_KWARGS_KEYS: Final = (
@ -70,6 +72,7 @@ OPTIONAL_KWARGS_KEYS: Final = (
}
)
| AWS_CREDENTIAL_KWARGS_KEYS
| REASONING_TRANSPORT_KWARGS_KEYS
| frozenset(CustomPricingLiteLLMParams.model_fields)
)

View file

@ -1,5 +1,9 @@
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from pydantic import TypeAdapter
import litellm
from litellm import verbose_logger
@ -9,6 +13,8 @@ from ...litellm_core_utils.get_llm_provider_logic import (
)
from ...types.router import LiteLLM_Params
_LITELLM_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
def _api_base_without_login(provider: str) -> str | None:
if provider == "github_copilot":
@ -50,9 +56,11 @@ def get_api_base(model: str, optional_params: dict | LiteLLM_Params) -> str | No
if isinstance(optional_params, LiteLLM_Params):
_optional_params = optional_params
elif "model" in optional_params:
_optional_params = LiteLLM_Params(**optional_params)
_optional_params = LiteLLM_Params.model_validate(optional_params)
else: # prevent needing to copy and pop the dict
_optional_params = LiteLLM_Params(model=model, **optional_params) # convert to pydantic object
_optional_params = LiteLLM_Params.model_validate(
_LITELLM_PARAMS_ADAPTER.validate_python(MappingProxyType({"model": model, **optional_params}))
) # convert to pydantic object
except Exception:
return None
# get llm provider

View file

@ -0,0 +1,54 @@
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.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, strings_only: bool = False
) -> list[AllMessageValues]: # mutable-ok: provider request contract
def normalize_message(message: AllMessageValues) -> AllMessageValues:
if message["role"] != "assistant":
return message
if not normalize and forward:
return message
history: Final[Mapping[str, object]] = message
removed_fields: Final = ("reasoning", "reasoning_content") if normalize else ("reasoning_content",)
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 removed_fields}),
**(
MappingProxyType({"reasoning": reasoning})
if normalize
and forward
and reasoning is not None
and (not strings_only or isinstance(reasoning, str))
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

@ -4,12 +4,17 @@ Translate from OpenAI's `/v1/chat/completions` to VLLM's `/v1/chat/completions`
import json
from collections.abc import Coroutine
from copy import deepcopy
from typing import Final, Literal, cast, overload
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,
should_normalize_reasoning_content,
)
from litellm.litellm_core_utils.reasoning_effort_utils import (
reasoning_effort_from_thinking_budget,
)
@ -142,6 +147,36 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
return ChatCompletionVideoObject(type="video_url", video_url=ChatCompletionVideoUrlObject(url=file_data))
raise ValueError("file_id or file_data is required")
def transform_request(
self,
model: str,
messages: list[AllMessageValues], # 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 not False,
strings_only=True,
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)
async def async_transform_request(
self,
model: str,
messages: list[AllMessageValues], # 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
)
@overload
def _transform_messages(
self, messages: list[AllMessageValues], model: str, is_async: Literal[True]

View file

@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast, overload
from urllib.parse import urlparse
import httpx
from pydantic import TypeAdapter
import litellm
from litellm.constants import OPENAI_SYSTEM_MESSAGES_FIRST_PROVIDERS
@ -32,6 +33,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,
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
@ -71,6 +76,7 @@ else:
_NO_TOOLS_UPDATE: Final[Mapping[str, object]] = MappingProxyType({})
_LITELLM_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
@ -487,8 +493,17 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
Returns:
dict: The transformed request. Sent as the body of the API call.
"""
transport_params: Final = _LITELLM_PARAMS_ADAPTER.validate_python(litellm_params)
request_messages: Final = (
normalize_reasoning_content(messages)
if transport_params.get("custom_llm_provider") == "openai"
and should_normalize_reasoning_content(
transport_params.get("reasoning_content_field"), model=model, provider="openai"
)
else messages
)
messages = self._transform_messages(
messages=self._prompt_cache_ordered_messages(messages, litellm_params), model=model
messages=self._prompt_cache_ordered_messages(request_messages, transport_params), model=model
)
if not self._should_preserve_cache_control_for_endpoint(
litellm_params.get("custom_llm_provider"), litellm_params.get("api_base")
@ -518,8 +533,17 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
litellm_params: dict,
headers: dict,
) -> dict:
transport_params: Final = _LITELLM_PARAMS_ADAPTER.validate_python(litellm_params)
request_messages: Final = (
normalize_reasoning_content(messages)
if transport_params.get("custom_llm_provider") == "openai"
and should_normalize_reasoning_content(
transport_params.get("reasoning_content_field"), model=model, provider="openai"
)
else messages
)
transformed_messages = await self._transform_messages(
messages=self._prompt_cache_ordered_messages(messages, litellm_params), model=model, is_async=True
messages=self._prompt_cache_ordered_messages(request_messages, transport_params), 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

@ -84,6 +84,7 @@ from litellm.litellm_core_utils.get_litellm_params import (
AWS_CREDENTIAL_KWARGS_KEYS,
OPTIONAL_KWARGS_KEYS,
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
REASONING_TRANSPORT_KWARGS_KEYS,
InvalidControlOption,
parse_control_options,
with_control_options,
@ -5702,7 +5703,11 @@ def completion(
gigachat_access_token=kwargs.get("gigachat_access_token"),
**{
key: kwargs[key]
for key in (*AWS_CREDENTIAL_KWARGS_KEYS, PROVIDER_AFFINITY_HEADER_KWARG_KEY)
for key in (
*AWS_CREDENTIAL_KWARGS_KEYS,
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
*REASONING_TRANSPORT_KWARGS_KEYS,
)
if key in kwargs
},
)

View file

@ -2197,7 +2197,7 @@ async def test_model_connection(
await ModelManagementAuthChecks.can_user_make_model_call(
model_params=Deployment(
model_name="test_model",
litellm_params=LiteLLM_Params(**litellm_params),
litellm_params=LiteLLM_Params.model_validate(litellm_params),
model_info=resolved_model_info,
),
user_api_key_dict=user_api_key_dict,

View file

@ -3959,7 +3959,7 @@ class Router:
model_info["original_model_id"] = original_model_id
deployment_pydantic_obj: Final = Deployment(
model_name=model_group,
litellm_params=LiteLLM_Params(**dynamic_litellm_params),
litellm_params=LiteLLM_Params.model_validate(dynamic_litellm_params),
model_info=model_info,
)
Router._register_deployment_pricing(deployment=deployment_pydantic_obj)
@ -9329,7 +9329,7 @@ class Router:
continue
deployment = Deployment(
model_name=model_name,
litellm_params=(lp if not isinstance(lp, dict) else LiteLLM_Params(**lp)),
litellm_params=(lp if not isinstance(lp, dict) else LiteLLM_Params.model_validate(lp)),
model_info=(entry.get("model_info") if isinstance(entry, dict) else entry.model_info),
)
if self._has_registered_strategy(self.adaptive_routers, model_name, self._deployment_tags(deployment)):
@ -10703,7 +10703,7 @@ class Router:
if isinstance(litellm_params_data, LiteLLM_Params):
litellm_params = litellm_params_data
elif isinstance(litellm_params_data, dict) and "model" in litellm_params_data:
litellm_params = LiteLLM_Params(**litellm_params_data)
litellm_params = LiteLLM_Params.model_validate(litellm_params_data)
else:
raise ValueError(
f"Deployment missing valid litellm_params. "
@ -12546,7 +12546,7 @@ class Router:
if allowed_model_region is not None:
if not is_region_allowed(
litellm_params=LiteLLM_Params(**_litellm_params),
litellm_params=LiteLLM_Params.model_validate(_litellm_params),
allowed_model_region=allowed_model_region,
):
invalid_model_indices.add(idx)
@ -12564,7 +12564,7 @@ class Router:
_,
) = litellm.get_llm_provider(
model=_dep_model_for_params,
litellm_params=LiteLLM_Params(**_litellm_params),
litellm_params=LiteLLM_Params.model_validate(_litellm_params),
)
except Exception as e: # noqa: BLE001 # best-effort filter: an unresolvable provider must not fail the request
verbose_router_logger.debug(

View file

@ -236,6 +236,8 @@ class PromptOptions:
@dataclass(frozen=True, slots=True, kw_only=True)
class ResponseOptions:
merge_reasoning_content_in_choices: bool | None = None
forward_reasoning_content: bool | None = None
reasoning_content_field: str | None = None
enable_json_schema_validation: bool | None = None
complete_response: bool | None = None
keepalive_seconds: float | None = None

View file

@ -526,6 +526,11 @@ class LiteLLM_Params(GenericLiteLLMParams):
model: str
model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True)
forward_reasoning_content: bool | None = None
reasoning_content_field: str | None = Field(
default=None,
description="Historical assistant reasoning field: reasoning_content (default) or reasoning.",
)
def __contains__(self, key) -> bool:
# Define custom behavior for the 'in' operator
@ -548,6 +553,11 @@ class updateLiteLLMParams(GenericLiteLLMParams):
# This class is used to update the LiteLLM_Params
# only differece is model is optional
model: str | None = None
forward_reasoning_content: bool | None = None
reasoning_content_field: str | None = Field(
default=None,
description="Historical assistant reasoning field: reasoning_content (default) or reasoning.",
)
class updateDeployment(BaseModel):

View file

@ -0,0 +1,266 @@
import asyncio
import json
from copy import deepcopy
from typing import Final
import httpx
import pytest
import respx
import litellm
from litellm import Router
from litellm.caching.caching import Cache
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
URL: Final = "https://reasoning-test.invalid/v1/chat/completions"
MODEL: Final = "hosted_vllm/reasoning-test"
REASONING: Final = "Inspect both tool results before answering."
def _messages():
return [
{"role": "user", "content": "Compare both records"},
{
"role": "assistant",
"content": None,
"reasoning_content": REASONING,
"tool_calls": [
{"id": f"call_{index}", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}
for index in (1, 2)
],
},
{"role": "tool", "tool_call_id": "call_1", "content": "first record"},
{"role": "tool", "tool_call_id": "call_2", "content": "second record"},
]
def _route(mock: respx.MockRouter):
return mock.post(URL).respond(
200,
json={
"id": "chatcmpl-reasoning-test",
"object": "chat.completion",
"created": 1,
"model": "reasoning-test",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Compared"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
},
)
def _assert_wire(request: httpx.Request, enabled: bool):
body: Final = json.loads(request.content)
assert body["model"] == "reasoning-test"
assert "forward_reasoning_content" not in body
assert "forward_reasoning_content" not in request.content.decode()
messages: Final = body["messages"]
assert [message["role"] for message in messages] == ["user", "assistant", "tool", "tool"]
assert [tool["id"] for tool in messages[1]["tool_calls"]] == ["call_1", "call_2"]
assert [message["tool_call_id"] for message in messages[2:]] == ["call_1", "call_2"]
assert [message["content"] for message in messages[2:]] == ["first record", "second record"]
assert messages[1].get("reasoning_content") == (REASONING if enabled else None)
assert request.content.decode().count(REASONING) == int(enabled)
@pytest.mark.asyncio
@pytest.mark.parametrize("async_mode", [False, True], ids=["completion", "acompletion"])
@pytest.mark.parametrize("forward", [None, False, True], ids=["absent", "false", "true"])
async def test_direct_completion_reasoning_flag_reaches_final_wire(
async_mode: bool, forward: bool | None, monkeypatch: pytest.MonkeyPatch
):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
messages: Final = _messages()
original: Final = deepcopy(messages)
kwargs: Final = {
"model": MODEL,
"api_base": URL.removesuffix("/chat/completions"),
"api_key": "test-key",
"messages": messages,
**({} if forward is None else {"forward_reasoning_content": forward}),
}
with respx.mock(assert_all_called=True) as mock:
route: Final = _route(mock)
response: Final = await litellm.acompletion(**kwargs) if async_mode else litellm.completion(**kwargs)
assert response.choices[0].message.content == "Compared"
assert route.call_count == 1
_assert_wire(route.calls[0].request, forward is not False)
assert messages == original
@pytest.mark.asyncio
@pytest.mark.parametrize("async_mode", [False, True], ids=["completion", "acompletion"])
async def test_router_aliases_isolate_reasoning_flag_on_same_backend(async_mode: bool, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
model_list: Final = [
{
"model_name": alias,
"litellm_params": {
"model": MODEL,
"api_base": URL.removesuffix("/chat/completions"),
"api_key": "test-key",
**({} if forward is None else {"forward_reasoning_content": forward}),
},
}
for alias, forward in (("default", None), ("disabled", False), ("enabled", True))
]
original_models: Final = deepcopy(model_list)
router: Final = Router(model_list=model_list, num_retries=0)
messages: Final = _messages()
original_messages: Final = deepcopy(messages)
with respx.mock(assert_all_called=True) as mock:
route: Final = _route(mock)
for alias in ("enabled", "default", "disabled", "enabled"):
response = (
await router.acompletion(model=alias, messages=messages)
if async_mode
else router.completion(model=alias, messages=messages)
)
assert response.choices[0].message.content == "Compared"
_assert_wire(route.calls[-1].request, alias != "disabled")
assert messages == original_messages
assert route.call_count == 4
assert model_list == original_models
@pytest.mark.asyncio
@pytest.mark.parametrize("async_mode", [False, True], ids=["sync", "async"])
@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, 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()
original_messages: Final = deepcopy(messages)
input_items: Final = [
{"role": "user", "content": "Compare both records"},
{
"type": "reasoning",
"id": "rs_previous",
"summary": [],
"content": [{"type": "reasoning_text", "text": REASONING}],
},
{"type": "function_call", "call_id": "call_1", "name": "lookup", "arguments": "{}"},
{"type": "function_call", "call_id": "call_2", "name": "lookup", "arguments": "{}"},
{"type": "function_call_output", "call_id": "call_1", "output": "first record"},
{"type": "function_call_output", "call_id": "call_2", "output": "second record"},
]
original_input: Final = deepcopy(input_items)
router: Final = Router(
model_list=[
{
"model_name": alias,
"litellm_params": {
"model": f"{provider}/reasoning-test",
"api_base": URL.removesuffix("/chat/completions"),
"api_key": "test-key",
"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))
],
cache_responses=True,
caching_groups=[("default", "disabled", "enabled")],
num_retries=0,
)
def backend(request: httpx.Request) -> httpx.Response:
body: Final = json.loads(request.content)
assert "forward_reasoning_content" not in body
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={
"id": "provider-cache-reasoning",
"object": "chat.completion",
"created": 1,
"model": "reasoning-test",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "forwarded" if enabled else "omitted"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
},
)
with respx.mock(assert_all_called=True) as mock:
route: Final = mock.post(URL).mock(side_effect=backend)
for alias, forward, expected_calls in (
("default", None, 1),
("disabled", False, 1 if normalize else 2),
("enabled", True, 2),
("disabled", False, 2),
("enabled", True, 2),
):
kwargs: Final = {
"model": f"{provider}/reasoning-test",
"api_base": URL.removesuffix("/chat/completions"),
"api_key": "test-key",
"caching": True,
**(
{
"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-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
else router.completion(model=alias, messages=messages)
)
elif surface == "responses":
response = (
await litellm.aresponses(input=input_items, use_chat_completions_api=True, **kwargs)
if async_mode
else litellm.responses(input=input_items, use_chat_completions_api=True, **kwargs)
)
else:
response = (
await litellm.acompletion(messages=messages, **kwargs)
if async_mode
else litellm.completion(messages=messages, **kwargs)
)
await asyncio.gather(*tuple(_PENDING_CACHE_WRITES))
content: Final = (
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 if normalize else forward is not False) else "omitted")
assert route.call_count == expected_calls
assert messages == original_messages
assert input_items == original_input

View file

@ -1039,6 +1039,20 @@ def _raw_batches_request(body: Dict[str, Any]) -> MagicMock:
request.headers = {"Content-Type": "application/json"}
request.client = MagicMock()
request.client.host = "127.0.0.1"
request.scope = {
"type": "http",
"asgi": {"version": "3.0", "spec_version": "2.3"},
"http_version": "1.1",
"method": "POST",
"scheme": "http",
"path": "/v1/batches",
"raw_path": b"/v1/batches",
"query_string": b"",
"root_path": "",
"headers": [(b"content-type", b"application/json"), (b"host", b"localhost")],
"client": ("127.0.0.1", 54321),
"server": ("localhost", 8000),
}
request.body = AsyncMock(return_value=json.dumps(body).encode())
return request

View file

@ -1340,7 +1340,9 @@ class TestUpdateModel:
"""
@pytest.mark.asyncio
async def test_update_model_clears_cache_after_db_write(self):
@pytest.mark.parametrize("reasoning_field", [None, "reasoning_content", "reasoning"])
@pytest.mark.parametrize("forward", [None, False, True])
async def test_update_model_clears_cache_after_db_write(self, reasoning_field, forward, monkeypatch):
"""
Regression test for the stale-router bug: POST /model/update must refresh
the in-memory router after persisting to LiteLLM_ProxyModelTable, otherwise
@ -1356,12 +1358,17 @@ class TestUpdateModel:
updateLiteLLMParams,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-reasoning-field")
model_id = "db-model-under-test"
existing_row = MagicMock()
existing_row.litellm_params = {
"model": "openai/gpt-4o-mini",
"api_key": "sk-existing",
"reasoning_content_field": encrypt_value_helper("reasoning"),
"forward_reasoning_content": True,
}
existing_row.model_dump.return_value = {
"model_name": "gpt-4o-mini",
@ -1396,10 +1403,6 @@ class TestUpdateModel:
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
new=AsyncMock(return_value=None),
),
patch( # test-quality-ok: [TQ008] isolate persistence from encryption implementation
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
side_effect=lambda value: value,
),
patch( # test-quality-ok: [TQ008] isolate persistence from router reload implementation
"litellm.proxy.management_endpoints.model_management_endpoints.clear_cache",
new=AsyncMock(
@ -1409,7 +1412,11 @@ class TestUpdateModel:
):
await update_model(
model_params=updateDeployment(
litellm_params=updateLiteLLMParams(guardrails=["g1"]),
litellm_params=updateLiteLLMParams(
**({"guardrails": ["g1"]} if reasoning_field is None and forward is None else {}),
**({} if reasoning_field is None else {"reasoning_content_field": reasoning_field}),
**({} if forward is None else {"forward_reasoning_content": forward}),
),
model_info=ModelInfo(id=model_id),
),
user_api_key_dict=admin_user,
@ -1417,6 +1424,11 @@ class TestUpdateModel:
mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once()
mock_clear_cache.assert_awaited_once_with()
stored = json.loads(mock_prisma.db.litellm_proxymodeltable.update.call_args.kwargs["data"]["litellm_params"])
assert decrypt_value_helper(stored["reasoning_content_field"], key="reasoning_content_field") == (reasoning_field or "reasoning")
if reasoning_field is None:
assert stored["reasoning_content_field"] == existing_row.litellm_params["reasoning_content_field"]
assert stored["forward_reasoning_content"] is (True if forward is None else forward)
@pytest.mark.asyncio
async def test_update_model_legacy_null_credential_name_is_not_a_detach_for_non_admin(self):
@ -7431,6 +7443,76 @@ class TestAccessGroupModelSync:
evict.assert_not_awaited()
@pytest.mark.parametrize("field", [None, "reasoning_content", "reasoning"])
@pytest.mark.parametrize("forward", [None, False, True])
def test_model_patch_preserves_reasoning_field_unless_explicit(field, forward, monkeypatch):
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.proxy.management_endpoints.model_management_endpoints import update_db_model
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-reasoning-field")
deployment = Deployment(
model_name="reasoning-test",
litellm_params=LiteLLM_Params(model="openai/reasoning-test", reasoning_content_field="reasoning", forward_reasoning_content=True),
model_info=ModelInfo(id="reasoning-row"),
)
params = updateLiteLLMParams(
**({"tpm": 123} if field is None and forward is None else {}),
**({} if field is None else {"reasoning_content_field": field}),
**({} if forward is None else {"forward_reasoning_content": forward}),
)
if forward is None:
assert "forward_reasoning_content" not in params.model_dump(exclude_none=True)
if field is None:
assert "reasoning_content_field" not in params.model_dump(exclude_none=True)
result = update_db_model(db_model=deployment, updated_patch=updateDeployment(litellm_params=params))
stored = json.loads(result["litellm_params"])
if field is None and forward is None:
assert stored["tpm"] == 123
assert stored["forward_reasoning_content"] is (True if forward is None else forward)
assert (
stored["reasoning_content_field"] if field is None
else decrypt_value_helper(value=stored["reasoning_content_field"], key="reasoning_content_field")
) == (field or "reasoning")
assert deployment.litellm_params.reasoning_content_field == "reasoning"
assert deployment.litellm_params.forward_reasoning_content is True
@pytest.mark.asyncio
@pytest.mark.parametrize("field", ["reasoning_content", "reasoning"])
async def test_reasoning_field_encrypted_db_round_trip(field, monkeypatch):
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.proxy.management_endpoints.model_management_endpoints import get_db_model, update_db_model
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-reasoning-field")
initial = Deployment(
model_name="reasoning-test",
litellm_params=LiteLLM_Params(model="openai/reasoning-test", forward_reasoning_content=True),
model_info=ModelInfo(id="reasoning-row"),
)
first = update_db_model(
db_model=initial,
updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(reasoning_content_field=field)),
)
encrypted = json.loads(first["litellm_params"])
assert encrypted["reasoning_content_field"] != field
raw = {"model_name": first["model_name"], "litellm_params": encrypted, "model_info": {"id": "reasoning-row"}}
deployment = Deployment(**raw)
assert deployment.litellm_params.reasoning_content_field == encrypted["reasoning_content_field"]
row = MagicMock()
row.model_dump.return_value = raw
prisma = MagicMock()
prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=row)
loaded = await get_db_model("reasoning-row", prisma)
assert loaded.litellm_params.reasoning_content_field == encrypted["reasoning_content_field"]
second = update_db_model(
db_model=loaded, updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(tpm=123))
)
stored = json.loads(second["litellm_params"])
assert stored["reasoning_content_field"] == encrypted["reasoning_content_field"]
assert stored["forward_reasoning_content"] is True
assert decrypt_value_helper(stored["reasoning_content_field"], key="reasoning_content_field") == field
class TestTeamMemberAutoRouterWrites:
@pytest.fixture(autouse=True)
def _salt(self, monkeypatch: pytest.MonkeyPatch) -> None:

View file

@ -18,13 +18,20 @@ async def test_trace_reader_projects_connection_and_parameters(recording_server:
recording_server.enqueue(ResponseSpec(body={"data": [{"trace_id": "trace-1"}]}))
reader_url: Final = recording_server.base_url.replace("http://", "http://reader:p%40ss%2Fword%25@")
storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, reader_url + "?database=wrong")
rows: Final = json.loads(await storage.query("trace_spans", {"trace_id": "trace-1"}))
rows: Final = json.loads(
await storage.query("trace_spans", {"trace_id": "trace-1", "team_ids": [], "api_key_hash": "", "trace_ref": ""})
)
request: Final = recording_server.requests[0]
parameters: Final = parse_qs(urlsplit(request.path).query)
parameters: Final = parse_qs(urlsplit(request.path).query, keep_blank_values=True)
assert rows == {"data": [{"trace_id": "trace-1"}]}
assert b"o.TraceId = {trace_id:String}" in request.raw_body
assert b"FROM otel_traces AS o" in request.raw_body
assert b"WHERE o.TraceId = {trace_id:String}" in request.raw_body
assert b"trace-1" not in request.raw_body
assert parameters["database"] == ["trace_test"]
assert parameters["param_trace_id"] == ["trace-1"]
assert parameters["param_team_ids"] == ["[]"]
assert parameters["param_api_key_hash"] == [""]
assert parameters["param_trace_ref"] == [""]
assert parameters["readonly"] == ["1"]
assert "user" not in parameters
assert "password" not in parameters
@ -36,11 +43,11 @@ async def test_trace_reader_rejects_success_status_with_embedded_error(recording
recording_server.enqueue(ResponseSpec(body={"data": [], "exception": "query failed"}))
storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, recording_server.base_url)
with pytest.raises(RuntimeError, match="invalid or failed JSON"):
await storage.query("trace_spans", {})
await storage.query("trace_spans", {"trace_id": "trace-1", "team_ids": [], "api_key_hash": "", "trace_ref": ""})
@pytest.mark.asyncio
async def test_reader_rejects_arbitrary_sql_before_sending(recording_server: RecordingServer) -> None:
async def test_trace_reader_rejects_arbitrary_sql_before_sending(recording_server: RecordingServer) -> None:
recording_server.expected_requests = 0
storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, recording_server.base_url)
with pytest.raises(ValueError, match="unknown ClickHouse read query"):
@ -61,7 +68,9 @@ async def test_schema_binding_rejects_non_positive_retention() -> None:
@pytest.mark.asyncio
async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement(recording_server: RecordingServer) -> None:
async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement(
recording_server: RecordingServer,
) -> None:
recording_server.expected_requests = 2
recording_server.enqueue(ResponseSpec(body=""))
recording_server.enqueue(ResponseSpec(status=403, body="denied"))
@ -73,25 +82,29 @@ async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement
assert recording_server.requests[0].raw_body.startswith(b"CREATE DATABASE IF NOT EXISTS")
assert recording_server.requests[1].raw_body.startswith(b"CREATE TABLE IF NOT EXISTS")
assert "readonly" not in parse_qs(urlsplit(recording_server.requests[0].path).query)
assert recording_server.requests[0].headers["authorization"] == "Basic " + base64.b64encode(
b"writer:p@ss/word%"
).decode()
assert (
recording_server.requests[0].headers["authorization"]
== "Basic " + base64.b64encode(b"writer:p@ss/word%").decode()
)
@pytest.mark.asyncio
async def test_insert_encodes_and_sends_rows(recording_server: RecordingServer) -> None:
recording_server.enqueue(ResponseSpec(body=""))
storage: Final = NativeTraceStorage("trace_test", recording_server.base_url)
before: Final = time.time_ns() // 1_000_000
before_insert_ms: Final = time.time_ns() // 1_000_000
await storage.insert_rows("otel_traces", [{"Timestamp": 1_234_567_890, "Input": "hello", "EngineReceivedMs": -1}])
after: Final = time.time_ns() // 1_000_000
after_insert_ms: Final = time.time_ns() // 1_000_000
request: Final = recording_server.requests[0]
row: Final = json.loads(gzip.decompress(request.raw_body))
assert before <= row["EngineReceivedMs"] <= after
assert type(row["EngineReceivedMs"]) is int
assert before_insert_ms <= row["EngineReceivedMs"] <= after_insert_ms
assert row == {
"Input": "hello",
"Timestamp": "1970-01-01T00:00:01.23456789Z",
"EngineReceivedMs": row["EngineReceivedMs"],
}
assert parse_qs(urlsplit(request.path).query)["query"] == ["INSERT INTO `trace_test`.otel_traces FORMAT JSONEachRow"]
assert parse_qs(urlsplit(request.path).query)["query"] == [
"INSERT INTO `trace_test`.otel_traces FORMAT JSONEachRow"
]
assert request.headers["content-encoding"] == "gzip"

View file

@ -120,6 +120,55 @@ def _semantic_cache(**cache_kwargs):
)
@pytest.mark.parametrize("semantic", [False, True])
def test_reasoning_forwarding_cache_scope_preserves_groups_namespace_and_tenant(semantic):
cache = _semantic_cache(namespace="reasoning-test") if semantic else Cache(type="local", namespace="reasoning-test")
def key(alias="first", tenant="tenant-a", namespace="reasoning-test", nested=False, forward=None, prompt="hi"):
return cache.get_cache_key(
model="hosted_vllm/reasoning-test",
messages=[{"role": "user", "content": prompt}],
metadata={"model_group": alias, "caching_groups": [("first", "second")], "user_api_key": tenant},
cache={"namespace": namespace},
**(
{"litellm_params": {"forward_reasoning_content": forward}}
if nested
else {"forward_reasoning_content": forward}
),
)
default = key()
enabled = key(forward=True)
assert default == enabled == key(nested=True, forward=True)
assert key(forward=False) == key(nested=True, forward=False)
assert key(forward=False) != default
assert enabled == key(nested=True, forward=True) == key(alias="second", forward=True)
assert enabled.startswith("reasoning-test:")
assert enabled != key(namespace="other-namespace", forward=True)
if semantic:
assert enabled != key(tenant="tenant-b", forward=True)
assert enabled == key(prompt="hello", forward=True)
else:
assert enabled != key(prompt="hello", forward=True)
def test_reasoning_forwarding_cache_key_preserves_default_key_and_top_level_precedence():
cache = Cache(type="local", namespace="reasoning-cache-test")
request = {"model": "hosted_vllm/reasoning-test", "messages": [{"role": "user", "content": "hi"}]}
legacy = "reasoning-cache-test:fca1120c8360f4b9ca0cd9b52f981f290a6eec25a8c6256033a81edcc713618c"
assert cache.get_cache_key(**request) == legacy
assert cache.get_cache_key(**request, forward_reasoning_content=True) == legacy
disabled = cache.get_cache_key(**request, forward_reasoning_content=False)
assert disabled != legacy
assert (
cache.get_cache_key(
**request, forward_reasoning_content=False, litellm_params={"forward_reasoning_content": True}
)
== disabled
)
assert cache.get_cache_key(**request, litellm_params={"forward_reasoning_content": True}) == legacy
@pytest.mark.parametrize(
"cache_type",
[LiteLLMCacheType.REDIS_SEMANTIC, LiteLLMCacheType.VALKEY_SEMANTIC],
@ -284,6 +333,48 @@ def test_exact_cache_key_includes_anthropic_messages_params(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 normalized != cache.get_cache_key(
**request, reasoning_content_field="reasoning", forward_reasoning_content=False
)
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",
)
@pytest.mark.asyncio
async def test_embedding_cache_skips_write_when_one_input_yields_many_embeddings(monkeypatch):
"""A cross-encoder behind /embeddings returns one score per document for a single

View file

@ -597,7 +597,9 @@ class TestExecuteSessionOperationSurfacesTransportError:
raise _FakeExceptionGroup("transport", [_FakeExceptionGroup("reader", failures)])
self._make_session(session_class, initialize)
expected: Final = close_error if failure_phase == "early" else connect_error if failure_phase == "mixed" else cancelled
expected: Final = (
close_error if failure_phase == "early" else connect_error if failure_phase == "mixed" else cancelled
)
with pytest.raises(type(expected)) as caught:
await client._execute_session_operation(self._make_transport(close_transport), AsyncMock(), http_client)
assert caught.value is expected
@ -617,7 +619,6 @@ class TestExecuteSessionOperationSurfacesTransportError:
result = await client._execute_session_operation(transport_ctx, _op)
assert result == "done"
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_session_entry_failure_still_closes_transport(self, session_class):
@ -1235,11 +1236,7 @@ def test_mcp_extra_matches_proxy_extra_and_supports_streamable_http():
sdk2_names: Final = frozenset(("mcp", "httpx2", "pydantic"))
mcp_extra: Final = {Requirement(req).name: req for req in extras["mcp"]}
assert mcp_extra == {
name: req
for req in extras["proxy"]
if (name := Requirement(req).name) in sdk2_names
}
assert mcp_extra == {name: req for req in extras["proxy"] if (name := Requirement(req).name) in sdk2_names}
specifier: Final = Requirement(mcp_extra["mcp"]).specifier
assert not specifier.contains("1.28.1")
@ -1639,7 +1636,9 @@ async def test_http_response_handler_preserves_success_and_http_errors(status_co
@pytest.mark.asyncio
async def test_http_status_check_allows_auth_refresh_before_rejecting() -> None:
from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import ClientCredentialsBearerAuth
from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import (
ClientCredentialsBearerAuth,
)
seen = []
@ -1892,14 +1891,20 @@ async def test_sse_read_failure_is_preserved() -> None:
@pytest.mark.parametrize("protocol_version", ["auto", "2025-06-18"])
@pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio])
@pytest.mark.parametrize("mode", ["ok", "closed", "silent"])
async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str, protocol_version: str) -> None:
async def test_transport_completion_and_normal_messages(
transport: MCPTransport, mode: str, protocol_version: str
) -> None:
from mcp import ClientSession
from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message
logging_callback: Final = AsyncMock()
read_timeout: Final = 0.2 if mode == "silent" else 30
client: Final = MCPClient(
server_url="https://example.com/sse", transport_type=transport, timeout=read_timeout, logging_callback=logging_callback, protocol_version=protocol_version
server_url="https://example.com/sse",
transport_type=transport,
timeout=read_timeout,
logging_callback=logging_callback,
protocol_version=protocol_version,
)
async def operation(session: ClientSession) -> CallToolResult:
@ -2461,8 +2466,15 @@ def test_client_import_before_proxy_credentials_succeeds_in_fresh_process():
import subprocess
result = subprocess.run(
[sys.executable, "-c", "import litellm.experimental_mcp_client.client; from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager; print(MCPServerManager.__name__)"],
capture_output=True, text=True, timeout=60, check=False,
[
sys.executable,
"-c",
"import litellm.experimental_mcp_client.client; from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager; print(MCPServerManager.__name__)",
],
capture_output=True,
text=True,
timeout=60,
check=False,
)
assert result.returncode == 0, result.stderr
assert result.stdout.strip() == "MCPServerManager"
@ -2495,8 +2507,10 @@ async def test_request_auth_preview_uses_the_same_effective_headers_as_egress()
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
client: Final = MCPClient(
server_url="https://upstream.example/mcp", auth_type=MCPAuth.bearer_token,
resolved_auth=StaticHeaderAuth("Bearer resolved"), extra_headers={"X-Trace": "trace"},
server_url="https://upstream.example/mcp",
auth_type=MCPAuth.bearer_token,
resolved_auth=StaticHeaderAuth("Bearer resolved"),
extra_headers={"X-Trace": "trace"},
)
request: Final = await client.prepare_request_auth()
assert request.method == "POST"
@ -2520,18 +2534,29 @@ async def test_expired_session_preserves_sdk_error_and_next_operation_reinitiali
return httpx2.Response(202)
requests.append((payload["method"], request.headers.get("mcp-session-id")))
if payload["method"] == "initialize":
return httpx2.Response(200, headers={"mcp-session-id": f"session-{len(requests)}"}, json={
"jsonrpc": "2.0", "id": payload["id"], "result": {
"protocolVersion": "2025-06-18", "capabilities": {},
"serverInfo": {"name": "expiry-test", "version": "1"},
return httpx2.Response(
200,
headers={"mcp-session-id": f"session-{len(requests)}"},
json={
"jsonrpc": "2.0",
"id": payload["id"],
"result": {
"protocolVersion": "2025-06-18",
"capabilities": {},
"serverInfo": {"name": "expiry-test", "version": "1"},
},
},
})
)
if len(requests) == 2:
if rpc_error:
return httpx2.Response(404, json={
"jsonrpc": "2.0", "id": payload["id"],
"error": {"code": METHOD_NOT_FOUND, "message": "Tool catalog unavailable"},
})
return httpx2.Response(
404,
json={
"jsonrpc": "2.0",
"id": payload["id"],
"error": {"code": METHOD_NOT_FOUND, "message": "Tool catalog unavailable"},
},
)
return httpx2.Response(404)
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload["id"], "result": {"tools": []}})
@ -2547,7 +2572,12 @@ async def test_expired_session_preserves_sdk_error_and_next_operation_reinitiali
streamable_http_client(client.server_url, http_client=http_client), lambda session: session.list_tools()
)
assert result.tools == []
assert requests == [("initialize", None), ("tools/list", "session-1"), ("initialize", None), ("tools/list", "session-3")]
assert requests == [
("initialize", None),
("tools/list", "session-1"),
("initialize", None),
("tools/list", "session-3"),
]
@pytest.mark.asyncio
@ -2604,7 +2634,9 @@ def test_public_mcp_import_preserves_incompatible_sdk_error() -> None:
@pytest.mark.parametrize("grouped", (False, True))
@pytest.mark.parametrize("raise_on_error", (False, True))
@pytest.mark.parametrize("termination", ("ok", "failure", "hang"))
async def test_outer_deadline_delivers_session_termination(termination: str, grouped: bool, raise_on_error: bool) -> None:
async def test_outer_deadline_delivers_session_termination(
termination: str, grouped: bool, raise_on_error: bool
) -> None:
deleted: Final = asyncio.Event()
started: Final = asyncio.Event()
@ -2644,12 +2676,23 @@ async def test_outer_deadline_delivers_session_termination(termination: str, gro
client: Final = _MockTransportClient(respond, server_url="https://example.com/mcp", timeout=30)
async def invoke():
with anyio.fail_after(0.2):
pending: Final = client.call_tool(CallToolRequestParams(name="slow", arguments={}), raise_on_error=raise_on_error)
if grouped:
await asyncio.gather(pending)
else:
pending: Final = asyncio.ensure_future(
client.call_tool(CallToolRequestParams(name="slow", arguments={}), raise_on_error=raise_on_error)
)
try:
with anyio.fail_after(2.0):
await started.wait()
with anyio.fail_after(0.2):
if grouped:
await asyncio.gather(pending)
else:
await pending
finally:
pending.cancel()
try:
await pending
except BaseException:
pass
before: Final = anyio.current_time()
with pytest.raises(TimeoutError):
@ -2849,7 +2892,9 @@ async def test_cancellation_delivers_termination_over_tcp(
listener: Final = await asyncio.start_server(handle_connection, "127.0.0.1", 0)
port: Final = listener.sockets[0].getsockname()[1]
client: Final = MCPClient(
server_url=f"http://127.0.0.1:{port}/mcp", protocol_version=protocol_version, timeout=2 if cancel_mode == "read_timeout" else 30
server_url=f"http://127.0.0.1:{port}/mcp",
protocol_version=protocol_version,
timeout=2 if cancel_mode == "read_timeout" else 30,
)
async def calls():
@ -2933,16 +2978,32 @@ async def test_configured_upstream_revision_is_offered_and_checked(revision, acc
assert payload.params["protocolVersion"] == offered
assert ("sampling" in payload.params["capabilities"]) == callbacks
assert ("elicitation" in payload.params["capabilities"]) == callbacks
return httpx2.Response(200, json={
"jsonrpc": "2.0", "id": payload.id,
"result": {"protocolVersion": offered if accepted else "unsupported",
"capabilities": {"tools": {}}, "serverInfo": {"name": "upstream", "version": "1"}},
})
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"protocolVersion": offered if accepted else "unsupported",
"capabilities": {"tools": {}},
"serverInfo": {"name": "upstream", "version": "1"},
},
},
)
assert accepted, "No operation may execute after failed version negotiation"
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"tools": [{"name": "echo", "inputSchema": {"type": "object"}}]}})
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {"tools": [{"name": "echo", "inputSchema": {"type": "object"}}]},
},
)
client = _MockTransportClient(
respond, server_url="https://example.com/mcp", protocol_version=revision,
respond,
server_url="https://example.com/mcp",
protocol_version=revision,
sampling_callback=AsyncMock() if callbacks else None,
elicitation_callback=AsyncMock() if callbacks else None,
)

View file

@ -10,7 +10,7 @@ Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v
import json
import os
import re
from typing import Any, Dict
from typing import Any, Dict, Final
from unittest.mock import MagicMock, patch
import httpx
@ -33,30 +33,10 @@ def _load_openapi_spec_dict() -> Dict[str, Any]:
return response.json()
except Exception as e: # pragma: no cover - defensive, env-dependent
pytest.skip(
f"Skipping Google Interactions OpenAPI compliance tests - "
f"unable to load spec from {OPENAPI_SPEC_URL}: {e}"
f"Skipping Google Interactions OpenAPI compliance tests - unable to load spec from {OPENAPI_SPEC_URL}: {e}"
)
def _model_create_request_schema(spec_dict: Dict[str, Any]) -> Dict[str, Any]:
schemas = spec_dict["components"]["schemas"]
create_path = next(path for path in spec_dict["paths"] if path.endswith("/interactions"))
body_schema = spec_dict["paths"][create_path]["post"]["requestBody"]["content"]["application/json"]["schema"]
variants = [schemas[option["$ref"].split("/")[-1]] for option in body_schema.get("oneOf", []) if "$ref" in option]
return next(variant for variant in variants if "model" in variant.get("properties", {}))
def _interaction_resource_path(spec_dict: Dict[str, Any], method: str) -> str | None:
return next(
(
path
for path, methods in spec_dict["paths"].items()
if re.search(r"/interactions/\{[^}]+\}$", path) and method in methods
),
None,
)
def _declared_type_value(variant_schema: Dict[str, Any]) -> Any:
"""The single `type` value a union variant pins, whether spelled as a const or a 1-item enum."""
type_property = variant_schema.get("properties", {}).get("type", {})
@ -64,6 +44,56 @@ def _declared_type_value(variant_schema: Dict[str, Any]) -> Any:
return type_property.get("const") or (enum_values[0] if len(enum_values) == 1 else None)
def _resolve_local_ref(spec_dict: dict[str, Any], schema: dict[str, Any]) -> dict[str, Any]:
"""Resolve component references used by operations, schemas, and parameters."""
if "$ref" not in schema:
return schema
reference: Final = schema["$ref"]
assert reference.startswith("#/components/"), f"Expected a local component reference: {reference}"
category, name = reference.removeprefix("#/components/").split("/")
return spec_dict["components"][category][name.replace("~1", "/").replace("~0", "~")]
def _interaction_operation(
spec_dict: dict[str, Any], method: str, *, individual: bool = False
) -> tuple[str, dict[str, Any]]:
"""Match collection or item routes exactly, independent of placeholder names."""
pattern: Final = r"(?:/[^/]+)*/interactions" + (r"/(\{[^/{}]+\})" if individual else "")
matches: Final = tuple(
(path, path_item, match)
for path, path_item in spec_dict["paths"].items()
if (match := re.fullmatch(pattern, path)) and method in path_item
)
assert len(matches) == 1, f"Expected one {method.upper()} interactions endpoint, got {matches}"
path, path_item, match = matches[0]
operation: Final = path_item[method]
if individual:
parameter_name: Final = match.group(1)[1:-1]
parameters: Final = {
(parameter["name"], parameter["in"]): parameter
for raw_parameter in (*path_item.get("parameters", ()), *operation.get("parameters", ()))
for parameter in (_resolve_local_ref(spec_dict, raw_parameter),)
}
parameter: Final = parameters.get((parameter_name, "path"))
assert parameter is not None, f"{path} must declare its interaction ID path parameter"
assert parameter.get("required") is True, f"{path} must require its interaction ID"
parameter_schema: Final = _resolve_local_ref(spec_dict, parameter["schema"])
assert parameter_schema.get("type") == "string", f"{path} must accept a string interaction ID"
return path, operation
def _model_request_schema(spec_dict: dict[str, Any]) -> dict[str, Any]:
"""Find the model variant of the JSON body declared by the create operation."""
_, operation = _interaction_operation(spec_dict, "post")
request_body: Final = _resolve_local_ref(spec_dict, operation["requestBody"])
assert request_body.get("required") is True, "Creating an interaction must require a request body"
schema: Final = _resolve_local_ref(spec_dict, request_body["content"]["application/json"]["schema"])
variants: Final = tuple(_resolve_local_ref(spec_dict, variant) for variant in schema.get("oneOf", (schema,)))
model_variants: Final = tuple(variant for variant in variants if "model" in variant.get("properties", {}))
assert len(model_variants) == 1, f"Expected one model request variant, got {model_variants}"
return model_variants[0]
@pytest.fixture(scope="module")
def spec_dict() -> Dict[str, Any]:
"""Load raw spec dict for manual validation."""
@ -80,10 +110,14 @@ class TestRequestCompliance:
"""Tests that our request bodies match the OpenAPI spec."""
def test_create_model_interaction_request_schema(self, spec_dict):
schema = _model_create_request_schema(spec_dict)
"""Verify the model request schema declared by POST /interactions."""
schema = _model_request_schema(spec_dict)
assert "model" in schema["required"]
assert "input" in schema["properties"]
for field in ("model", "input"):
assert field in schema["properties"]
assert schema["properties"][field].get("readOnly") is not True
assert _resolve_local_ref(spec_dict, schema["properties"][field]).get("readOnly") is not True
# Check our supported optional fields exist in spec
our_optional_fields = [
@ -106,13 +140,8 @@ class TestRequestCompliance:
def test_input_types_match_spec(self, spec_dict):
"""Verify input field supports string, Content, Content[], Turn[]."""
schema = _model_create_request_schema(spec_dict)
input_schema = schema["properties"]["input"]
# The input property may be inline oneOf or a $ref to InteractionsInput
if "$ref" in input_schema:
ref_name = input_schema["$ref"].split("/")[-1]
input_schema = spec_dict["components"]["schemas"][ref_name]
schema = _model_request_schema(spec_dict)
input_schema = _resolve_local_ref(spec_dict, schema["properties"]["input"])
# Should be oneOf with multiple types
assert "oneOf" in input_schema
@ -143,22 +172,18 @@ class TestRequestCompliance:
discriminator = content_schema.get("discriminator")
if discriminator is not None:
assert (
discriminator.get("propertyName") == "type"
), f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'"
assert discriminator.get("propertyName") == "type", (
f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'"
)
variant_names = [
option["$ref"].split("/")[-1]
for option in content_schema.get("oneOf", [])
if "$ref" in option
option["$ref"].split("/")[-1] for option in content_schema.get("oneOf", []) if "$ref" in option
]
assert variant_names, f"Content is not a union of named variants: {content_schema}"
mapping = (discriminator or {}).get("mapping") or {}
type_values = {
variant: mapping_value
for mapping_value, ref in mapping.items()
for variant in [ref.split("/")[-1]]
variant: mapping_value for mapping_value, ref in mapping.items() for variant in [ref.split("/")[-1]]
} or {
variant: _declared_type_value(spec_dict["components"]["schemas"].get(variant, {}))
for variant in variant_names
@ -209,7 +234,9 @@ class TestRequestCompliance:
for option in spec_dict["components"]["schemas"]["Step"]["oneOf"]
if "$ref" in option
}
assert {"UserInputStep", "ModelOutputStep"} <= step_variants, f"Step union is missing role steps: {step_variants}"
assert {"UserInputStep", "ModelOutputStep"} <= step_variants, (
f"Step union is missing role steps: {step_variants}"
)
for step_name, type_value in [("UserInputStep", "user_input"), ("ModelOutputStep", "model_output")]:
step_schema = spec_dict["components"]["schemas"][step_name]
@ -279,9 +306,7 @@ class TestResponseCompliance:
expected_fields = ["total_input_tokens", "total_output_tokens", "total_tokens"]
for field in expected_fields:
assert (
field in usage_schema["properties"]
), f"Usage field '{field}' not in spec"
assert field in usage_schema["properties"], f"Usage field '{field}' not in spec"
print(f"✓ Usage field '{field}' exists")
@ -300,9 +325,7 @@ class TestToolsCompliance:
"""Verify FunctionDeclaration schema for function tools."""
if "FunctionDeclaration" in spec_dict["components"]["schemas"]:
func_schema = spec_dict["components"]["schemas"]["FunctionDeclaration"]
assert "name" in func_schema.get(
"properties", {}
) or "name" in func_schema.get("required", [])
assert "name" in func_schema.get("properties", {}) or "name" in func_schema.get("required", [])
print("✓ FunctionDeclaration schema found")
else:
print("⚠ FunctionDeclaration schema not found (may be nested)")
@ -313,33 +336,94 @@ class TestEndpointCompliance:
def test_create_endpoint_exists(self, spec_dict):
"""Verify POST /interactions endpoint exists."""
paths = spec_dict["paths"]
# Find the create interactions endpoint
create_path = None
for path, methods in paths.items():
if "interactions" in path and "post" in methods:
create_path = path
break
assert create_path is not None, "POST /interactions endpoint not found"
create_path, _ = _interaction_operation(spec_dict, "post")
print(f"✓ Create endpoint: POST {create_path}")
def test_get_endpoint_exists(self, spec_dict):
"""Verify GET /interactions/{id} endpoint exists."""
get_path = _interaction_resource_path(spec_dict, "get")
assert get_path is not None, "GET /interactions/{id} endpoint not found"
get_path, _ = _interaction_operation(spec_dict, "get", individual=True)
print(f"✓ Get endpoint: GET {get_path}")
def test_delete_endpoint_exists(self, spec_dict):
"""Verify DELETE /interactions/{id} endpoint exists."""
delete_path = _interaction_resource_path(spec_dict, "delete")
assert delete_path is not None, "DELETE /interactions/{id} endpoint not found"
delete_path, _ = _interaction_operation(spec_dict, "delete", individual=True)
print(f"✓ Delete endpoint: DELETE {delete_path}")
class TestOperationResolution:
"""Keep structural resolution strict without depending on generated names."""
@pytest.mark.parametrize("as_union", [False, True])
def test_model_schema_comes_from_create_operation(self, as_union):
model_schema: Final = {"properties": {"model": {"type": "string"}}, "required": ["model"]}
reference: Final = {"$ref": "#/components/schemas/RenamedModelRequest"}
body_schema: Final = (
{"oneOf": [{"properties": {"agent": {"type": "string"}}}, reference]} if as_union else reference
)
spec: Final = {
"paths": {
"/{version}/interactions": {
"post": {
"requestBody": {"required": True, "content": {"application/json": {"schema": body_schema}}}
}
}
},
"components": {
"schemas": {"RenamedModelRequest": model_schema, "CreateModelInteractionParams": {"properties": {}}}
},
}
assert _model_request_schema(spec) is model_schema
@pytest.mark.parametrize("method,shared", [("get", False), ("delete", True)])
def test_item_route_accepts_a_renamed_declared_identifier(self, method, shared):
parameter: Final = {"name": "renamedId", "in": "path", "required": True, "schema": {"type": "string"}}
parameters: Final = [{"$ref": "#/components/parameters/Identifier"}]
operation: Final = {"parameters": [] if shared else parameters}
path: Final = "/{version}/interactions/{renamedId}"
spec: Final = {
"paths": {path: {"parameters": parameters if shared else [], method: operation}},
"components": {"parameters": {"Identifier": parameter}},
}
assert _interaction_operation(spec, method, individual=True) == (path, operation)
@pytest.mark.parametrize(
"path,parameter,error",
[
(
"/interactions/{id}/cancel",
{"required": True, "type": "string"},
"Expected one GET interactions endpoint",
),
(
"/other_interactions/{id}",
{"required": True, "type": "string"},
"Expected one GET interactions endpoint",
),
("/interactions/{id}", {"required": False, "type": "string"}, "must require its interaction ID"),
("/interactions/{id}", {"required": True, "type": "integer"}, "must accept a string interaction ID"),
],
)
def test_item_route_rejects_incompatible_contracts(self, path, parameter, error):
spec: Final = {
"paths": {
path: {
"get": {
"parameters": [
{
"name": "id",
"in": "path",
"required": parameter["required"],
"schema": {"type": parameter["type"]},
}
]
}
}
}
}
with pytest.raises(AssertionError, match=error):
_interaction_operation(spec, "get", individual=True)
if __name__ == "__main__":
# Quick manual test
import httpx
@ -356,6 +440,4 @@ if __name__ == "__main__":
if method in ["get", "post", "delete", "put", "patch"]:
print(f" {method.upper()} {path}")
print(
f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}..."
)
print(f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}...")

View file

@ -1,4 +1,5 @@
import json
from typing import Final
import pytest
@ -36,6 +37,21 @@ class TestDeclaredAuthenticatingProvider:
the declaration without resolving. The recorder appends before raising, and get_api_base
swallows resolver errors, so an empty list proves the lookup never ran."""
@pytest.mark.parametrize("include_model", [False, True])
def test_invalid_retry_text_does_not_resolve_provider(self, include_model, resolution_lookups):
params: Final = {
"max_retries": "2.0",
"self": "reserved-placeholder",
**({"model": "openai/demo"} if include_model else {}),
}
original: Final = dict(params)
api_base: Final = litellm.get_api_base(model="openai/demo", optional_params=params)
assert api_base is None
assert resolution_lookups == []
assert params == original
@pytest.mark.parametrize(
"model, custom_llm_provider, expected",
[
@ -82,7 +98,10 @@ class TestDeclaredAuthenticatingProvider:
@pytest.mark.parametrize(
"model, expected",
[
("gemini/gemini-2.5-pro", "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:generateContent"),
(
"gemini/gemini-2.5-pro",
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:generateContent",
),
("openai/gpt-4o", "https://api.openai.com"),
],
)

View file

@ -1,4 +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
@ -11,6 +18,90 @@ from litellm.constants import (
from litellm.llms.hosted_vllm.chat.transformation import HostedVLLMChatConfig
@pytest.mark.parametrize("params", [{}, {"forward_reasoning_content": False}, {"forward_reasoning_content": True}])
@pytest.mark.parametrize("is_async", [False, True])
@pytest.mark.asyncio
async def test_forward_reasoning_content_preserves_only_explicit_history(
params: dict[str, bool], is_async: bool
) -> None:
config = HostedVLLMChatConfig()
messages = [
{"role": "user", "content": "Check the counter"},
{
"role": "assistant",
"content": None,
"reasoning_content": " synthetic history\n",
"thinking_blocks": [{"type": "thinking", "thinking": "Do not convert this", "signature": "sig"}],
"tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "read", "arguments": "{}"}}],
},
{"role": "tool", "tool_call_id": "call_1", "content": "7"},
{
"role": "assistant",
"content": None,
"thinking_blocks": [{"type": "thinking", "thinking": "Never synthesize history", "signature": "sig"}],
"tool_calls": [{"id": "call_2", "type": "function", "function": {"name": "verify", "arguments": "{}"}}],
},
{"role": "tool", "tool_call_id": "call_2", "content": "verified"},
]
original = deepcopy(messages)
arguments = dict(
model="qwen3.8-flash-next", messages=messages, optional_params={}, litellm_params=params, headers={}
)
result = await config.async_transform_request(**arguments) if is_async else config.transform_request(**arguments)
expected = deepcopy(original)
expected[1].pop("thinking_blocks")
expected[3].pop("thinking_blocks")
if params.get("forward_reasoning_content") is False:
expected[1].pop("reasoning_content")
assert result["messages"] == expected
assert messages == original
assert "forward_reasoning_content" not in result
@pytest.mark.asyncio
async def test_forward_reasoning_content_reused_config_and_caller_are_isolated():
config = HostedVLLMChatConfig()
messages = [{"role": "assistant", "content": "answer", "reasoning_content": "synthetic history"}]
original = deepcopy(messages)
for enabled in (False, True, False, True):
for transform in (config.transform_request, config.async_transform_request):
result = transform(
model="qwen3.8-flash-next",
messages=messages,
optional_params={},
litellm_params={"forward_reasoning_content": enabled},
headers={},
)
if transform == config.async_transform_request:
result = await result
assert result["messages"] == (original if enabled else [{"role": "assistant", "content": "answer"}])
assert messages == original
@pytest.mark.asyncio
async def test_forward_reasoning_content_keeps_async_content_conversion():
class AsyncContentConfig(HostedVLLMChatConfig):
async def _async_transform_content_item(self, content_item):
return {"type": "image_url", "image_url": {"url": "data:image/png;base64,c3ludGhldGlj"}}
config = AsyncContentConfig()
messages = [
{"role": "user", "content": [{"type": "image_url", "image_url": {"url": "https://example.invalid/image.png"}}]},
{"role": "assistant", "content": "answer", "reasoning_content": "synthetic history"},
]
original = deepcopy(messages)
result = await config.async_transform_request(
model="qwen3.8-flash-next",
messages=messages,
optional_params={},
litellm_params={"forward_reasoning_content": True},
headers={},
)
assert result["messages"][0]["content"][0]["image_url"]["url"] == "data:image/png;base64,c3ludGhldGlj"
assert result["messages"][1]["reasoning_content"] == "synthetic history"
assert messages == original
def test_hosted_vllm_chat_transformation_file_url():
config = HostedVLLMChatConfig()
video_url = "https://example.com/video.mp4"
@ -397,3 +488,217 @@ 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 not False
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])
@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()
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": 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: str,
is_async: bool,
via_router: bool,
bridge: bool,
field: str,
forward: bool,
monkeypatch: pytest.MonkeyPatch,
) -> None:
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
@pytest.mark.parametrize(
"history",
[
{"reasoning_content": 42},
{"reasoning_content": ["step"]},
{"reasoning_content": {"text": "step"}},
{"reasoning": 42},
{"reasoning_content": "source", "reasoning": 42},
{"reasoning_content": ""},
{"reasoning": ""},
],
)
@pytest.mark.parametrize("is_async", [False, True])
@pytest.mark.asyncio
async def test_normalized_hosted_reasoning_only_forwards_strings(history, is_async):
config = HostedVLLMChatConfig()
messages = [{"role": "assistant", "content": "answer", **history}]
original = deepcopy(messages)
arguments = dict(
model="hosted_vllm/test",
messages=messages,
optional_params={},
litellm_params={"reasoning_content_field": "reasoning"},
headers={},
)
result = await config.async_transform_request(**arguments) if is_async else config.transform_request(**arguments)
value = history.get("reasoning", history.get("reasoning_content"))
expected = {"role": "assistant", "content": "answer", **({"reasoning": value} if isinstance(value, str) else {})}
assert result["messages"] == [expected]
assert messages == original

View file

@ -4,6 +4,7 @@ Tests for OpenAI GPT transformation (litellm/llms/openai/chat/gpt_transformation
import pytest
from copy import deepcopy
from typing import Final
@ -1206,6 +1207,34 @@ class TestSystemMessagesFirst:
)
assert tuple(m["content"] for m in messages) == self.ORIGINAL_ORDER
@pytest.mark.asyncio
@pytest.mark.parametrize("is_async", [False, True])
@pytest.mark.parametrize("enabled", [False, True])
async def test_reasoning_normalization_preserves_prompt_cache_ordering(
self, monkeypatch: pytest.MonkeyPatch, is_async: bool, enabled: bool
) -> None:
monkeypatch.setattr(litellm, "openai_system_messages_first", enabled)
messages: Final = [
{**message, **({"reasoning_content": "thinking"} if message["role"] == "assistant" else {})}
for message in self.MESSAGES
]
original: Final = deepcopy(messages)
kwargs: Final = {
"model": "reasoning-test",
"messages": messages,
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai", "reasoning_content_field": "reasoning"},
"headers": {},
}
request: Final = (
await self.config.async_transform_request(**kwargs) if is_async else self.config.transform_request(**kwargs)
)
assert tuple(m["content"] for m in request["messages"]) == (self.ORDERED if enabled else self.ORIGINAL_ORDER)
assert next(m for m in request["messages"] if m["role"] == "assistant") == {
"role": "assistant", "content": "reply", "reasoning": "thinking"
}
assert messages == original
@pytest.mark.asyncio
async def test_async_transform_request_moves_system_messages_first(self, monkeypatch):
class UninstantiatedOpenAIGPTConfig(OpenAIGPTConfig):

View file

@ -2,7 +2,9 @@ import json
from copy import deepcopy
from typing import Final, Literal
import httpx
import pytest
import respx
from openai.types.responses.response_function_web_search import (
ActionFind,
ActionOpenPage,
@ -30,6 +32,114 @@ from litellm.types.utils import (
)
@pytest.mark.asyncio
@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,
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."
next_reasoning: Final = "Use the first result to inspect the second."
input_items: Final = [
{"role": "user", "content": "Compare both records"},
{
"type": "reasoning",
"id": "rs_previous",
"summary": [],
"content": [{"type": "reasoning_text", "text": reasoning}],
},
{"type": "function_call", "call_id": "call_1", "name": "lookup", "arguments": "{}"},
*(
[
{"type": "function_call_output", "call_id": "call_1", "output": "first record"},
{
"type": "reasoning",
"id": "rs_next",
"summary": [],
"content": [{"type": "reasoning_text", "text": next_reasoning}],
},
]
if sequential
else []
),
{"type": "function_call", "call_id": "call_2", "name": "lookup", "arguments": "{}"},
*([] if sequential else [{"type": "function_call_output", "call_id": "call_1", "output": "first record"}]),
{"type": "function_call_output", "call_id": "call_2", "output": "second record"},
]
original: Final = deepcopy(input_items)
kwargs: Final = {
"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 not False
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(
200,
json={
"id": "chatcmpl-bridge-reasoning",
"object": "chat.completion",
"created": 1,
"model": "reasoning-test",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "Compared"}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
},
)
)
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"]
)
assert [tool["id"] for message in messages for tool in message.get("tool_calls", [])] == ["call_1", "call_2"]
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(output_field) == (reasoning if forwarded else None)
assert route.calls[0].request.content.decode().count(reasoning) == int(forwarded)
if sequential:
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
class TestLiteLLMCompletionResponsesConfig:
def test_transform_input_file_item_to_file_item_with_file_id(self):
"""Test transformation of input_file item with file_id to Chat Completion file format"""

View file

@ -761,30 +761,31 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"supports_computer_use": {"type": "boolean"},
"cache_creation_input_audio_token_cost": {"type": "number"},
"cache_creation_input_token_cost": {"type": "number"},
"cache_creation_input_token_cost_ultrafast": {"type": "number"},
"cache_creation_input_token_cost_above_1hr": {"type": "number"},
"cache_creation_input_token_cost_above_32k_tokens": {"type": "number"},
"cache_creation_input_token_cost_above_128k_tokens": {"type": "number"},
"cache_creation_input_token_cost_above_200k_tokens": {"type": "number"},
"cache_creation_input_token_cost_above_256k_tokens": {"type": "number"},
"cache_creation_input_token_cost_above_272k_tokens": {"type": "number"},
"cache_creation_input_token_cost_above_272k_tokens_flex": {"type": "number"},
"cache_creation_input_token_cost_above_272k_tokens_ultrafast": {"type": "number"},
"cache_creation_input_token_cost_above_272k_tokens_flex": {"type": "number"},
"cache_creation_input_token_cost_above_272k_tokens_priority": {"type": "number"},
"cache_creation_input_token_cost_above_200k_tokens_batches": {"type": "number"},
"cache_creation_input_token_cost_above_272k_tokens_batches": {"type": "number"},
"cache_creation_input_token_cost_batches": {"type": "number"},
"cache_creation_input_token_cost_flex": {"type": "number"},
"cache_creation_input_token_cost_priority": {"type": "number"},
"cache_creation_input_token_cost_ultrafast": {"type": "number"},
"cache_read_input_token_cost": {"type": "number"},
"cache_read_input_token_cost_ultrafast": {"type": "number"},
"cache_read_input_token_cost_above_32k_tokens": {"type": "number"},
"cache_read_input_token_cost_above_128k_tokens": {"type": "number"},
"cache_read_input_token_cost_above_200k_tokens": {"type": "number"},
"cache_read_input_token_cost_above_200k_tokens_batches": {"type": "number"},
"cache_read_input_token_cost_above_256k_tokens": {"type": "number"},
"cache_read_input_token_cost_above_272k_tokens": {"type": "number"},
"cache_read_input_token_cost_above_272k_tokens_flex": {"type": "number"},
"cache_read_input_token_cost_above_272k_tokens_ultrafast": {"type": "number"},
"cache_read_input_token_cost_above_272k_tokens_flex": {"type": "number"},
"cache_read_input_token_cost_above_512k_tokens": {"type": "number"},
"input_cost_per_token_above_272k_tokens_ultrafast": {"type": "number"},
"cache_read_input_token_cost_batches": {"type": "number"},
@ -813,7 +814,6 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"cache_read_input_token_cost_flex": {"type": "number"},
"cache_read_input_token_cost_priority": {"type": "number"},
"cache_read_input_token_cost_balanced": {"type": "number"},
"cache_read_input_token_cost_ultrafast": {"type": "number"},
"cache_read_input_token_cost_above_200k_tokens_priority": {"type": "number"},
"cache_read_input_token_cost_above_272k_tokens_priority": {"type": "number"},
"input_cost_per_token_flex": {"type": "number"},

View file

@ -182,6 +182,8 @@ OPTION_NAMES: Final = (
"assistant_continue_message",
"disable_add_transform_inline_image_block",
"merge_reasoning_content_in_choices",
"forward_reasoning_content",
"reasoning_content_field",
"enable_json_schema_validation",
"complete_response",
"stream_chunk_size",

View file

@ -34095,6 +34095,8 @@ export interface components {
default_api_key_tpm_limit?: number | null;
/** Drop Params */
drop_params?: boolean | string | null;
/** Forward Reasoning Content */
forward_reasoning_content?: boolean | null;
/** Gcs Bucket Name */
gcs_bucket_name?: string | null;
/** Google Maps Grounding Cost Per Query */
@ -34300,6 +34302,11 @@ export interface components {
} | null;
/** Quality Router Default Model */
quality_router_default_model?: string | null;
/**
* Reasoning Content Field
* @description Historical assistant reasoning field: reasoning_content (default) or reasoning.
*/
reasoning_content_field?: string | null;
/** Region Name */
region_name?: string | null;
/** Regional Endpoint Uplift Multiplier */
@ -48453,6 +48460,8 @@ export interface components {
default_api_key_tpm_limit?: number | null;
/** Drop Params */
drop_params?: boolean | string | null;
/** Forward Reasoning Content */
forward_reasoning_content?: boolean | null;
/** Gcs Bucket Name */
gcs_bucket_name?: string | null;
/** Google Maps Grounding Cost Per Query */
@ -48658,6 +48667,11 @@ export interface components {
} | null;
/** Quality Router Default Model */
quality_router_default_model?: string | null;
/**
* Reasoning Content Field
* @description Historical assistant reasoning field: reasoning_content (default) or reasoning.
*/
reasoning_content_field?: string | null;
/** Region Name */
region_name?: string | null;
/** Regional Endpoint Uplift Multiplier */