mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge 9f6bfb26ac into 0980f756bd
This commit is contained in:
commit
69fd88afb8
26 changed files with 1387 additions and 140 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
54
litellm/litellm_core_utils/reasoning_content_utils.py
Normal file
54
litellm/litellm_core_utils/reasoning_content_utils.py
Normal file
|
|
@ -0,0 +1,54 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from copy import deepcopy
|
||||
from types import MappingProxyType
|
||||
from typing import (
|
||||
Final,
|
||||
cast, # noqa: TID251 # Preserves arbitrary provider fields without lossy TypedDict validation.
|
||||
)
|
||||
|
||||
from litellm.exceptions import BadRequestError
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
||||
def should_normalize_reasoning_content(field: object, *, model: str, provider: str) -> bool:
|
||||
if field is None or field == "reasoning_content":
|
||||
return False
|
||||
if field == "reasoning":
|
||||
return True
|
||||
raise BadRequestError(
|
||||
message="reasoning_content_field must be reasoning_content or reasoning",
|
||||
model=model,
|
||||
llm_provider=provider,
|
||||
)
|
||||
|
||||
|
||||
def normalize_reasoning_content(
|
||||
messages: Sequence[AllMessageValues], *, forward: bool = True, normalize: bool = True, strings_only: bool = False
|
||||
) -> list[AllMessageValues]: # mutable-ok: provider request contract
|
||||
def normalize_message(message: AllMessageValues) -> AllMessageValues:
|
||||
if message["role"] != "assistant":
|
||||
return message
|
||||
if not normalize and forward:
|
||||
return message
|
||||
history: Final[Mapping[str, object]] = message
|
||||
removed_fields: Final = ("reasoning", "reasoning_content") if normalize else ("reasoning_content",)
|
||||
reasoning: Final = (
|
||||
history.get("reasoning") if history.get("reasoning") is not None else history.get("reasoning_content")
|
||||
)
|
||||
normalized: Final[Mapping[str, object]] = MappingProxyType(
|
||||
{
|
||||
**MappingProxyType({key: value for key, value in history.items() if key not in removed_fields}),
|
||||
**(
|
||||
MappingProxyType({"reasoning": reasoning})
|
||||
if normalize
|
||||
and forward
|
||||
and reasoning is not None
|
||||
and (not strings_only or isinstance(reasoning, str))
|
||||
else MappingProxyType({})
|
||||
),
|
||||
}
|
||||
)
|
||||
result: Final = dict(normalized) # mutable-ok: provider request contract
|
||||
return cast(AllMessageValues, result) # cast-ok: only optional reasoning keys change
|
||||
|
||||
return [normalize_message(message) for message in deepcopy(messages)] # mutable-ok: provider request contract
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
266
tests/router_unit_tests/test_router_forward_reasoning_content.py
Normal file
266
tests/router_unit_tests/test_router_forward_reasoning_content.py
Normal file
|
|
@ -0,0 +1,266 @@
|
|||
import asyncio
|
||||
import json
|
||||
from copy import deepcopy
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
|
||||
|
||||
URL: Final = "https://reasoning-test.invalid/v1/chat/completions"
|
||||
MODEL: Final = "hosted_vllm/reasoning-test"
|
||||
REASONING: Final = "Inspect both tool results before answering."
|
||||
|
||||
|
||||
def _messages():
|
||||
return [
|
||||
{"role": "user", "content": "Compare both records"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"reasoning_content": REASONING,
|
||||
"tool_calls": [
|
||||
{"id": f"call_{index}", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}
|
||||
for index in (1, 2)
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "first record"},
|
||||
{"role": "tool", "tool_call_id": "call_2", "content": "second record"},
|
||||
]
|
||||
|
||||
|
||||
def _route(mock: respx.MockRouter):
|
||||
return mock.post(URL).respond(
|
||||
200,
|
||||
json={
|
||||
"id": "chatcmpl-reasoning-test",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "reasoning-test",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Compared"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _assert_wire(request: httpx.Request, enabled: bool):
|
||||
body: Final = json.loads(request.content)
|
||||
assert body["model"] == "reasoning-test"
|
||||
assert "forward_reasoning_content" not in body
|
||||
assert "forward_reasoning_content" not in request.content.decode()
|
||||
messages: Final = body["messages"]
|
||||
assert [message["role"] for message in messages] == ["user", "assistant", "tool", "tool"]
|
||||
assert [tool["id"] for tool in messages[1]["tool_calls"]] == ["call_1", "call_2"]
|
||||
assert [message["tool_call_id"] for message in messages[2:]] == ["call_1", "call_2"]
|
||||
assert [message["content"] for message in messages[2:]] == ["first record", "second record"]
|
||||
assert messages[1].get("reasoning_content") == (REASONING if enabled else None)
|
||||
assert request.content.decode().count(REASONING) == int(enabled)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("async_mode", [False, True], ids=["completion", "acompletion"])
|
||||
@pytest.mark.parametrize("forward", [None, False, True], ids=["absent", "false", "true"])
|
||||
async def test_direct_completion_reasoning_flag_reaches_final_wire(
|
||||
async_mode: bool, forward: bool | None, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
messages: Final = _messages()
|
||||
original: Final = deepcopy(messages)
|
||||
kwargs: Final = {
|
||||
"model": MODEL,
|
||||
"api_base": URL.removesuffix("/chat/completions"),
|
||||
"api_key": "test-key",
|
||||
"messages": messages,
|
||||
**({} if forward is None else {"forward_reasoning_content": forward}),
|
||||
}
|
||||
with respx.mock(assert_all_called=True) as mock:
|
||||
route: Final = _route(mock)
|
||||
response: Final = await litellm.acompletion(**kwargs) if async_mode else litellm.completion(**kwargs)
|
||||
assert response.choices[0].message.content == "Compared"
|
||||
assert route.call_count == 1
|
||||
_assert_wire(route.calls[0].request, forward is not False)
|
||||
assert messages == original
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("async_mode", [False, True], ids=["completion", "acompletion"])
|
||||
async def test_router_aliases_isolate_reasoning_flag_on_same_backend(async_mode: bool, monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
model_list: Final = [
|
||||
{
|
||||
"model_name": alias,
|
||||
"litellm_params": {
|
||||
"model": MODEL,
|
||||
"api_base": URL.removesuffix("/chat/completions"),
|
||||
"api_key": "test-key",
|
||||
**({} if forward is None else {"forward_reasoning_content": forward}),
|
||||
},
|
||||
}
|
||||
for alias, forward in (("default", None), ("disabled", False), ("enabled", True))
|
||||
]
|
||||
original_models: Final = deepcopy(model_list)
|
||||
router: Final = Router(model_list=model_list, num_retries=0)
|
||||
messages: Final = _messages()
|
||||
original_messages: Final = deepcopy(messages)
|
||||
with respx.mock(assert_all_called=True) as mock:
|
||||
route: Final = _route(mock)
|
||||
for alias in ("enabled", "default", "disabled", "enabled"):
|
||||
response = (
|
||||
await router.acompletion(model=alias, messages=messages)
|
||||
if async_mode
|
||||
else router.completion(model=alias, messages=messages)
|
||||
)
|
||||
assert response.choices[0].message.content == "Compared"
|
||||
_assert_wire(route.calls[-1].request, alias != "disabled")
|
||||
assert messages == original_messages
|
||||
assert route.call_count == 4
|
||||
assert model_list == original_models
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("async_mode", [False, True], ids=["sync", "async"])
|
||||
@pytest.mark.parametrize("surface", ["sdk", "router", "responses", "router-responses"])
|
||||
@pytest.mark.parametrize("provider", ["hosted_vllm", "openai"])
|
||||
@pytest.mark.parametrize("normalize", [False, True])
|
||||
async def test_local_cache_separates_forwarded_reasoning_history(
|
||||
async_mode: bool, surface: str, provider: str, normalize: bool, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
if provider == "openai" and not normalize:
|
||||
pytest.skip("OpenAI does not apply hosted_vllm forwarding policy")
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local", namespace="reasoning-cache-test"))
|
||||
messages: Final = _messages()
|
||||
original_messages: Final = deepcopy(messages)
|
||||
input_items: Final = [
|
||||
{"role": "user", "content": "Compare both records"},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "rs_previous",
|
||||
"summary": [],
|
||||
"content": [{"type": "reasoning_text", "text": REASONING}],
|
||||
},
|
||||
{"type": "function_call", "call_id": "call_1", "name": "lookup", "arguments": "{}"},
|
||||
{"type": "function_call", "call_id": "call_2", "name": "lookup", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "call_1", "output": "first record"},
|
||||
{"type": "function_call_output", "call_id": "call_2", "output": "second record"},
|
||||
]
|
||||
original_input: Final = deepcopy(input_items)
|
||||
router: Final = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": alias,
|
||||
"litellm_params": {
|
||||
"model": f"{provider}/reasoning-test",
|
||||
"api_base": URL.removesuffix("/chat/completions"),
|
||||
"api_key": "test-key",
|
||||
"use_chat_completions_api": True,
|
||||
**(
|
||||
{
|
||||
"forward_reasoning_content": True,
|
||||
**(
|
||||
{}
|
||||
if enabled is None
|
||||
else {"reasoning_content_field": "reasoning" if enabled else "reasoning_content"}
|
||||
),
|
||||
}
|
||||
if normalize
|
||||
else ({} if enabled is None else {"forward_reasoning_content": enabled})
|
||||
),
|
||||
},
|
||||
}
|
||||
for alias, enabled in (("default", None), ("disabled", False), ("enabled", True))
|
||||
],
|
||||
cache_responses=True,
|
||||
caching_groups=[("default", "disabled", "enabled")],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
def backend(request: httpx.Request) -> httpx.Response:
|
||||
body: Final = json.loads(request.content)
|
||||
assert "forward_reasoning_content" not in body
|
||||
assert "reasoning_content_field" not in body
|
||||
enabled: Final = body["messages"][1].get("reasoning" if normalize else "reasoning_content") == REASONING
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "provider-cache-reasoning",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "reasoning-test",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "forwarded" if enabled else "omitted"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
|
||||
},
|
||||
)
|
||||
|
||||
with respx.mock(assert_all_called=True) as mock:
|
||||
route: Final = mock.post(URL).mock(side_effect=backend)
|
||||
for alias, forward, expected_calls in (
|
||||
("default", None, 1),
|
||||
("disabled", False, 1 if normalize else 2),
|
||||
("enabled", True, 2),
|
||||
("disabled", False, 2),
|
||||
("enabled", True, 2),
|
||||
):
|
||||
kwargs: Final = {
|
||||
"model": f"{provider}/reasoning-test",
|
||||
"api_base": URL.removesuffix("/chat/completions"),
|
||||
"api_key": "test-key",
|
||||
"caching": True,
|
||||
**(
|
||||
{
|
||||
"forward_reasoning_content": True,
|
||||
**(
|
||||
{}
|
||||
if forward is None
|
||||
else {"reasoning_content_field": "reasoning" if forward else "reasoning_content"}
|
||||
),
|
||||
}
|
||||
if normalize
|
||||
else ({} if forward is None else {"forward_reasoning_content": forward})
|
||||
),
|
||||
}
|
||||
if surface == "router-responses":
|
||||
response = (
|
||||
await router.aresponses(model=alias, input=input_items)
|
||||
if async_mode
|
||||
else router.responses(model=alias, input=input_items)
|
||||
)
|
||||
elif surface == "router":
|
||||
response = (
|
||||
await router.acompletion(model=alias, messages=messages)
|
||||
if async_mode
|
||||
else router.completion(model=alias, messages=messages)
|
||||
)
|
||||
elif surface == "responses":
|
||||
response = (
|
||||
await litellm.aresponses(input=input_items, use_chat_completions_api=True, **kwargs)
|
||||
if async_mode
|
||||
else litellm.responses(input=input_items, use_chat_completions_api=True, **kwargs)
|
||||
)
|
||||
else:
|
||||
response = (
|
||||
await litellm.acompletion(messages=messages, **kwargs)
|
||||
if async_mode
|
||||
else litellm.completion(messages=messages, **kwargs)
|
||||
)
|
||||
await asyncio.gather(*tuple(_PENDING_CACHE_WRITES))
|
||||
content: Final = (
|
||||
response.output[0].content[0].text
|
||||
if surface in ("responses", "router-responses")
|
||||
else response.choices[0].message.content
|
||||
)
|
||||
assert content == ("forwarded" if (forward is True if normalize else forward is not False) else "omitted")
|
||||
assert route.call_count == expected_calls
|
||||
assert messages == original_messages
|
||||
assert input_items == original_input
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]}...")
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue