diff --git a/backend/routes/allowlist.py b/backend/routes/allowlist.py index 51a4d8f716c..d8296774409 100644 --- a/backend/routes/allowlist.py +++ b/backend/routes/allowlist.py @@ -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 diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 9e04ca79822..10e790ed970 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -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) diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index b8441d2bc6d..b36c4d6db6f 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -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) ) diff --git a/litellm/litellm_core_utils/llm_response_utils/get_api_base.py b/litellm/litellm_core_utils/llm_response_utils/get_api_base.py index 4d731b5e63a..ef54231351a 100644 --- a/litellm/litellm_core_utils/llm_response_utils/get_api_base.py +++ b/litellm/litellm_core_utils/llm_response_utils/get_api_base.py @@ -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 diff --git a/litellm/litellm_core_utils/reasoning_content_utils.py b/litellm/litellm_core_utils/reasoning_content_utils.py new file mode 100644 index 00000000000..be1d227a3e2 --- /dev/null +++ b/litellm/litellm_core_utils/reasoning_content_utils.py @@ -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 diff --git a/litellm/llms/hosted_vllm/chat/transformation.py b/litellm/llms/hosted_vllm/chat/transformation.py index 43eb2af171e..d6360982c48 100644 --- a/litellm/llms/hosted_vllm/chat/transformation.py +++ b/litellm/llms/hosted_vllm/chat/transformation.py @@ -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] diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 3b38825c83d..bf61ceb43ab 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -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") diff --git a/litellm/main.py b/litellm/main.py index 8c9d7f2513d..10db315c66e 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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 }, ) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 07be73d7573..feb9d718099 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -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, diff --git a/litellm/router.py b/litellm/router.py index bb118639839..c4e911fa4e0 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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( diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py index 20214078852..2db37371c63 100644 --- a/litellm/types/litellm_params.py +++ b/litellm/types/litellm_params.py @@ -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 diff --git a/litellm/types/router.py b/litellm/types/router.py index 2ab1a1185ed..7fa49352a9c 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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): diff --git a/tests/router_unit_tests/test_router_forward_reasoning_content.py b/tests/router_unit_tests/test_router_forward_reasoning_content.py new file mode 100644 index 00000000000..116acba2382 --- /dev/null +++ b/tests/router_unit_tests/test_router_forward_reasoning_content.py @@ -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 diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 3bf51f02d34..b5d83461e2e 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 5f7807650e1..e756a039608 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -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: diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index fc750d88e42..e3eddb3e9a5 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -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" diff --git a/tests/unit/caching/test_caching.py b/tests/unit/caching/test_caching.py index 0e0f2b7eac6..fa0d4db742c 100644 --- a/tests/unit/caching/test_caching.py +++ b/tests/unit/caching/test_caching.py @@ -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 diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 1a56227b008..d16de968c25 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -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, ) diff --git a/tests/unit/interactions/test_openapi_compliance.py b/tests/unit/interactions/test_openapi_compliance.py index 247d02298aa..b373fd95b28 100644 --- a/tests/unit/interactions/test_openapi_compliance.py +++ b/tests/unit/interactions/test_openapi_compliance.py @@ -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]}...") diff --git a/tests/unit/litellm_core_utils/llm_response_utils/test_get_api_base.py b/tests/unit/litellm_core_utils/llm_response_utils/test_get_api_base.py index 90fd5ba6c88..5193ebdb78e 100644 --- a/tests/unit/litellm_core_utils/llm_response_utils/test_get_api_base.py +++ b/tests/unit/litellm_core_utils/llm_response_utils/test_get_api_base.py @@ -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"), ], ) diff --git a/tests/unit/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py b/tests/unit/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py index c792d5dffcd..afe09bbf2e5 100644 --- a/tests/unit/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py +++ b/tests/unit/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py @@ -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 diff --git a/tests/unit/llms/openai/chat/test_openai_gpt_transformation.py b/tests/unit/llms/openai/chat/test_openai_gpt_transformation.py index 85a04778e5c..5c028000fe3 100644 --- a/tests/unit/llms/openai/chat/test_openai_gpt_transformation.py +++ b/tests/unit/llms/openai/chat/test_openai_gpt_transformation.py @@ -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): diff --git a/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 7b9de4644b4..1487578af0b 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -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""" diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 4c36f99d2c8..5ac5649bba7 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -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"}, diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index ab4f6c12431..de833717c75 100644 --- a/tests/unit/types/test_litellm_params.py +++ b/tests/unit/types/test_litellm_params.py @@ -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", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index f709583c7dc..f75e4fd67a1 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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 */