Merge branch 'litellm_internal_staging' into litellm_/dark-mode-logo-strategy-7b99f2

This commit is contained in:
Yuneng Jiang 2026-08-27 16:09:25 -07:00
commit 1a9045efd4
No known key found for this signature in database
112 changed files with 4571 additions and 673 deletions

View file

@ -6,10 +6,10 @@
"limit": 2564
},
"reportAssignmentType": {
"limit": 320
"limit": 319
},
"reportAttributeAccessIssue": {
"limit": 483
"limit": 480
},
"reportCallIssue": {
"limit": 113
@ -30,7 +30,7 @@
"limit": 7
},
"reportGeneralTypeIssues": {
"limit": 154
"limit": 105
},
"reportIncompatibleMethodOverride": {
"limit": 56
@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44528
"limit": 44526
},
"reportUnknownLambdaType": {
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38804
"limit": 38782
},
"reportUnknownParameterType": {
"limit": 19829
},
"reportUnknownVariableType": {
"limit": 30355
"limit": 30349
},
"reportUnnecessaryCast": {
"limit": 117
@ -123,7 +123,7 @@
"limit": 5
},
"reportUnnecessaryIsInstance": {
"limit": 833
"limit": 831
},
"reportUntypedBaseClass": {
"limit": 0

View file

@ -73,6 +73,11 @@ ARRAY_KEYS: dict[str, JsonSchema] = {
"description": "Output modalities the model can produce.",
"items": {"type": "string", "enum": ["text", "image", "audio", "video", "code"]},
},
"reasoning_effort_levels": {
"type": "array",
"description": "Exact reasoning_effort levels this deployment accepts; wins over supports_* flags.",
"items": {"type": "string", "enum": ["none", "minimal", "low", "medium", "high", "xhigh", "max"]},
},
"supported_regions": {
"type": "array",
"description": "Cloud regions the model is available in ('global' or region ids).",

View file

@ -59,9 +59,11 @@ if TYPE_CHECKING:
from litellm.types.llms.openai import (
ALL_RESPONSES_API_TOOL_PARAMS,
AllMessageValues,
ChatCompletionFileObject,
ChatCompletionImageObject,
ChatCompletionRedactedThinkingBlock,
ChatCompletionThinkingBlock,
ChatCompletionToolReferenceObject,
OpenAIMessageContentListBlock,
)
from litellm.types.utils import Choices
@ -175,6 +177,16 @@ def _map_incomplete_reason_to_finish_reason(incomplete_reason: str | None) -> Li
return "length"
def _input_file_from_file_value(file_value: object) -> dict[str, object]:
if not isinstance(file_value, dict):
return {"type": "input_file"}
file_dict: Final = cast("dict[str, object]", file_value) # cast-ok: runtime dict checked
return {
"type": "input_file",
**{key: file_dict[key] for key in ("file_id", "file_data", "filename") if key in file_dict},
}
def _incomplete_reason_from_response_payload(response_payload: object) -> str | None:
if not isinstance(response_payload, Mapping):
return None
@ -957,7 +969,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
content: str
| list[object]
| Iterable[
Union["OpenAIMessageContentListBlock", "ChatCompletionThinkingBlock", "ChatCompletionRedactedThinkingBlock"]
Union[
"OpenAIMessageContentListBlock",
"ChatCompletionThinkingBlock",
"ChatCompletionRedactedThinkingBlock",
"ChatCompletionToolReferenceObject",
]
]
| None,
role: str,
@ -1006,17 +1023,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
result.append(converted)
verbose_logger.debug("Chat provider: image -> %s", converted)
elif item_type == "file":
# Map Chat Completion file to Responses API input_file
# {"type": "file", "file": {"file_data": "...", "filename": "..."}}
# -> {"type": "input_file", "file_data": "...", "filename": "..."}
file_data = item.get("file", {})
converted = {"type": "input_file"}
if isinstance(file_data, dict):
for key in ["file_id", "file_data", "filename"]:
if key in file_data:
converted[key] = file_data[key]
converted = _input_file_from_file_value(
cast("ChatCompletionFileObject", item).get("file"), # cast-ok: type tag checked
)
result.append(converted)
verbose_logger.debug("Chat provider: file -> %s", converted)
elif item_type == "tool_reference":
verbose_logger.debug(
"Chat provider: tool_reference has no responses API equivalent; skipped"
)
elif item_type in [
"input_text",
"input_image",

View file

@ -76,7 +76,10 @@ from litellm.llms.perplexity.cost_calculator import (
from litellm.llms.tencent.cost_calculator import (
cost_per_token as tencent_cost_per_token,
)
from litellm.llms.together_ai.cost_calculator import get_model_params_and_category
from litellm.llms.together_ai.cost_calculator import (
get_model_params_and_category,
has_together_registry_pricing,
)
from litellm.llms.vertex_ai.cost_calculator import (
cost_per_character as google_cost_per_character,
)
@ -1569,10 +1572,9 @@ def completion_cost(
return MCPCostCalculator.calculate_mcp_tool_call_cost(litellm_logging_obj=litellm_logging_obj)
# Calculate cost based on prompt_tokens, completion_tokens
if "togethercomputer" in model or "together_ai" in model or custom_llm_provider == "together_ai":
# together ai prices based on size of llm
# get_model_params_and_category takes a model name and returns the category of LLM size it is in model_prices_and_context_window.json
if (
"togethercomputer" in model or "together_ai" in model or custom_llm_provider == "together_ai"
) and not has_together_registry_pricing(model, litellm.model_cost):
model = get_model_params_and_category(model, call_type=CallTypes(call_type))
# replicate llms are calculate based on time for request running

View file

@ -56,6 +56,9 @@ from litellm.types.mcp import (
MCPStdioConfig,
MCPTransport,
MCPTransportType,
credential_redirect_hook,
has_header,
without_header,
)
@ -273,6 +276,7 @@ class MCPClient:
transport_type: MCPTransportType = MCPTransport.http,
auth_type: MCPAuthType = None,
auth_value: str | dict[str, str] | None = None,
auth_header_name: str | None = None,
timeout: float | None = None,
stdio_config: MCPStdioConfig | None = None,
extra_headers: dict[str, str] | None = None,
@ -288,6 +292,11 @@ class MCPClient:
self.auth_type: MCPAuthType = auth_type
self.timeout: float = timeout if timeout is not None else MCP_CLIENT_TIMEOUT
self._mcp_auth_value: str | dict[str, str] | None = None
# The one place this client decides which header its credential occupies: the operator's
# configured slot on the v1 path, or the slot the v2 resolver's auth object already owns.
# Every consumer reads this rather than re-deriving it, since each re-derivation so far
# picked up a different bug.
self._credential_slot: str | None = auth_header_name or getattr(resolved_auth, "header_name", None)
self.stdio_config: MCPStdioConfig | None = stdio_config
self.extra_headers: dict[str, str] | None = extra_headers
self.ssl_verify: VerifyTypes | None = ssl_verify
@ -501,26 +510,33 @@ class MCPClient:
else:
self._mcp_auth_value = mcp_auth_value
def _header_slot(self, default: str) -> str:
return self._credential_slot or default
def _get_auth_headers(self) -> dict:
"""Generate authentication headers based on auth type."""
headers: Final = {}
if self._mcp_auth_value:
if isinstance(self._mcp_auth_value, str):
if self.auth_type == MCPAuth.bearer_token:
headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}"
static_bearer: Final = strip_auth_scheme(self._mcp_auth_value, "Bearer")
headers[self._header_slot("Authorization")] = f"Bearer {static_bearer}"
elif self.auth_type == MCPAuth.basic:
headers["Authorization"] = f"Basic {self._mcp_auth_value}"
headers[self._header_slot("Authorization")] = f"Basic {self._mcp_auth_value}"
elif self.auth_type == MCPAuth.api_key:
headers["X-API-Key"] = self._mcp_auth_value
headers[self._header_slot("X-API-Key")] = self._mcp_auth_value
elif self.auth_type == MCPAuth.authorization:
# This auth type means the caller owns the whole header value.
headers["Authorization"] = self._mcp_auth_value
headers[self._header_slot("Authorization")] = self._mcp_auth_value
elif self.auth_type == MCPAuth.oauth2:
headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}"
oauth2_bearer: Final = strip_auth_scheme(self._mcp_auth_value, "Bearer")
headers[self._header_slot("Authorization")] = f"Bearer {oauth2_bearer}"
elif self.auth_type == MCPAuth.token:
headers["Authorization"] = f"token {strip_auth_scheme(self._mcp_auth_value, 'token')}"
scheme_token: Final = strip_auth_scheme(self._mcp_auth_value, "token")
headers[self._header_slot("Authorization")] = f"token {scheme_token}"
elif self.auth_type == MCPAuth.oauth2_token_exchange:
headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}"
exchanged_bearer: Final = strip_auth_scheme(self._mcp_auth_value, "Bearer")
headers[self._header_slot("Authorization")] = f"Bearer {exchanged_bearer}"
elif isinstance(self._mcp_auth_value, dict):
headers.update(self._mcp_auth_value)
# Note: aws_sigv4 auth is not handled here — SigV4 requires per-request
@ -528,7 +544,14 @@ class MCPClient:
# of static headers. See MCPSigV4Auth and _create_httpx_client_factory().
# update the headers with the extra headers
if self.extra_headers:
headers.update(self.extra_headers)
# Mirrors _resolve_v2_auth: when the operator named a slot for the credential the
# gateway resolved, no injected header may shadow it, case-insensitively, since HTTP
# header names are. Without a configured slot the old precedence stands unchanged.
slot: Final = self._credential_slot
injected: Final = (
without_header(self.extra_headers, slot) if slot and has_header(headers, slot) else self.extra_headers
)
headers.update(injected or {})
return _strip_header_whitespace(headers)
def _create_httpx_client_factory(self) -> Callable[..., httpx.AsyncClient]:
@ -556,12 +579,14 @@ class MCPClient:
# SigV4 aws_auth. Both are None for the common case — no behavior change.
fallback_auth: Final = self._resolved_auth if self._resolved_auth is not None else self._aws_auth
effective_auth: Final = auth if auth is not None else fallback_auth
guard: Final = credential_redirect_hook(self.server_url, self._credential_slot)
return httpx.AsyncClient(
headers=headers,
timeout=timeout,
auth=effective_auth,
verify=ssl_config,
follow_redirects=True,
event_hooks={"request": [guard]} if guard else {},
)
return factory

View file

@ -376,7 +376,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
# 2. list of objects - only apply to last item per Anthropic spec
elif isinstance(message_content, list):
if len(message_content) > 0 and isinstance(message_content[-1], dict):
message_content[-1]["cache_control"] = control
message_content[-1]["cache_control"] = control # pyright: ignore[reportGeneralTypeIssues] # loose runtime dict
return message
@staticmethod

View file

@ -0,0 +1,193 @@
"""Provider-agnostic SRT/WebVTT subtitle synthesis from timestamped transcription tokens."""
from collections.abc import Sequence
from dataclasses import dataclass
from itertools import accumulate, chain
from typing import Final
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
CUE_MAX_TOKENS: Final = 15
CUE_MAX_DURATION_MS: Final = 5000
SRT_RESPONSE_FORMAT: Final = "srt"
VTT_RESPONSE_FORMAT: Final = "vtt"
SUBTITLE_RESPONSE_FORMATS: Final = frozenset((SRT_RESPONSE_FORMAT, VTT_RESPONSE_FORMAT))
@dataclass(frozen=True, slots=True)
class SubtitleToken:
text: str
start_ms: int | None = None
end_ms: int | None = None
speaker: str | int | None = None
@dataclass(frozen=True, slots=True)
class SubtitleCue:
start_ms: int
end_ms: int
text: str
@dataclass(frozen=True, slots=True)
class _CueAccumulator:
texts: tuple[str, ...] = ()
start_ms: int | None = None
end_ms: int | None = None
speaker: str | int | None = None
def _completed_cue(accumulator: _CueAccumulator) -> tuple[SubtitleCue, ...]:
if not accumulator.texts or accumulator.start_ms is None:
return ()
text: Final = "".join(accumulator.texts).strip()
if not text:
return ()
end_ms: Final = accumulator.end_ms if accumulator.end_ms is not None else accumulator.start_ms
return (SubtitleCue(start_ms=accumulator.start_ms, end_ms=end_ms, text=text),)
def _cue_break_reached(accumulator: _CueAccumulator, token: SubtitleToken) -> bool:
if len(accumulator.texts) >= CUE_MAX_TOKENS:
return True
return (
accumulator.start_ms is not None
and token.start_ms is not None
and token.start_ms - accumulator.start_ms >= CUE_MAX_DURATION_MS
)
_AbsorbStep = tuple[tuple[SubtitleCue, ...], _CueAccumulator]
def _absorb_token(accumulator: _CueAccumulator, token: SubtitleToken) -> _AbsorbStep:
if token.start_ms is None and accumulator.start_ms is None:
return (), accumulator
if token.speaker is not None and token.speaker != accumulator.speaker:
return _completed_cue(accumulator), _CueAccumulator(
texts=(token.text,),
start_ms=token.start_ms,
end_ms=token.end_ms,
speaker=token.speaker,
)
if _cue_break_reached(accumulator, token):
return _completed_cue(accumulator), _CueAccumulator(
texts=(token.text,),
start_ms=token.start_ms,
end_ms=token.end_ms,
speaker=accumulator.speaker,
)
return (), _CueAccumulator(
texts=(*accumulator.texts, token.text),
start_ms=accumulator.start_ms if accumulator.start_ms is not None else token.start_ms,
end_ms=token.end_ms if token.end_ms is not None else accumulator.end_ms,
speaker=accumulator.speaker,
)
def _absorb_step(carry: _AbsorbStep, token: SubtitleToken) -> _AbsorbStep:
return _absorb_token(carry[1], token)
def group_subtitle_tokens_into_cues(tokens: Sequence[SubtitleToken]) -> tuple[SubtitleCue, ...]:
steps: Final = tuple(accumulate(tokens, _absorb_step, initial=((), _CueAccumulator())))
completed: Final = chain.from_iterable(emitted for emitted, _ in steps)
return (*completed, *_completed_cue(steps[-1][1]))
def _format_timestamp(total_ms: int, millis_separator: str) -> str:
clamped: Final = max(total_ms, 0)
hours, hour_remainder = divmod(clamped, 3_600_000)
minutes, minute_remainder = divmod(hour_remainder, 60_000)
seconds, millis = divmod(minute_remainder, 1_000)
return f"{hours:02d}:{minutes:02d}:{seconds:02d}{millis_separator}{millis:03d}"
def _render_srt(cues: Sequence[SubtitleCue]) -> str:
lines: Final = tuple(
line
for index, cue in enumerate(cues, start=1)
for line in (
str(index),
f"{_format_timestamp(cue.start_ms, ',')} --> {_format_timestamp(cue.end_ms, ',')}",
cue.text,
"",
)
)
return "\n".join(lines)
def _render_vtt(cues: Sequence[SubtitleCue]) -> str:
cue_lines: Final = tuple(
line
for cue in cues
for line in (
f"{_format_timestamp(cue.start_ms, '.')} --> {_format_timestamp(cue.end_ms, '.')}",
cue.text,
"",
)
)
return "\n".join(("WEBVTT", "", *cue_lines))
def render_subtitle_tokens_as_srt(tokens: Sequence[SubtitleToken]) -> str:
"""Render tokens as an SRT document; empty string when no token has timestamp data."""
cues: Final = group_subtitle_tokens_into_cues(tokens)
if not cues:
return ""
return _render_srt(cues)
def render_subtitle_tokens_as_vtt(tokens: Sequence[SubtitleToken]) -> str:
"""Render tokens as a WebVTT document; the WEBVTT header is emitted even without cues."""
return _render_vtt(group_subtitle_tokens_into_cues(tokens))
class TranscriptionWordTiming(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
word: str = ""
start: float | None = None
end: float | None = None
speaker: str | None = None
_WORD_TIMINGS_ADAPTER: Final = TypeAdapter(tuple[TranscriptionWordTiming, ...])
def _seconds_to_ms(seconds: float | None) -> int | None:
if seconds is None:
return None
return round(seconds * 1000)
def _word_to_subtitle_token(word: TranscriptionWordTiming) -> SubtitleToken:
return SubtitleToken(
text=f"{word.word} ",
start_ms=_seconds_to_ms(word.start),
end_ms=_seconds_to_ms(word.end),
speaker=word.speaker,
)
def _parse_word_timings(words: object) -> tuple[TranscriptionWordTiming, ...]:
try:
return _WORD_TIMINGS_ADAPTER.validate_python(words)
except ValidationError:
return ()
def synthesize_subtitle_document(words: object, response_format: str) -> str | None:
"""
Build an SRT/VTT document from OpenAI verbose_json-style word dicts
(word/start/end in float seconds, optional speaker). Returns None when the
format is not a subtitle format or the words carry no usable timestamps.
"""
if response_format not in SUBTITLE_RESPONSE_FORMATS:
return None
tokens: Final = tuple(_word_to_subtitle_token(word) for word in _parse_word_timings(words))
cues: Final = group_subtitle_tokens_into_cues(tokens)
if not cues:
return None
return _render_srt(cues) if response_format == SRT_RESPONSE_FORMAT else _render_vtt(cues)

View file

@ -1747,6 +1747,46 @@ def hoist_images_from_tool_messages(
]
def _is_tool_reference_part(part: object) -> bool:
return isinstance(part, dict) and part.get("type") == "tool_reference"
def _tool_message_carries_tool_reference(message: AllMessageValues) -> bool:
if message.get("role") != "tool":
return False
content = message.get("content")
return isinstance(content, list) and any(_is_tool_reference_part(part) for part in content)
def _drop_tool_reference_parts(message: AllMessageValues) -> AllMessageValues:
if not _tool_message_carries_tool_reference(message):
return message
content = cast(list, message.get("content")) # cast-ok: shape checked by _tool_message_carries_tool_reference
remaining_parts = [ # mutable-ok: tool message content must stay a json list
part for part in content if not _is_tool_reference_part(part)
]
new_content = remaining_parts if remaining_parts else ""
rewritten = {**message, "content": new_content} # mutable-ok: chat messages are plain json dicts
return cast(AllMessageValues, rewritten) # cast-ok: dict spread keeps keys like cache_control
def drop_tool_reference_parts_from_tool_messages(
messages: list[AllMessageValues], # mutable-ok: message pipelines type messages as mutable lists
) -> list[AllMessageValues]: # mutable-ok: message pipelines type messages as mutable lists
"""
Remove tool_reference content parts from role:"tool" messages.
The OpenAI chat spec only accepts text in tool messages, so a tool_reference
part carried through the Anthropic adapter makes strict providers reject the
request. The reference names an already-declared tool rather than carrying
content, so it is dropped; a reference-only result keeps its tool message with
empty text so the preceding tool_call stays answered.
"""
if not any(_tool_message_carries_tool_reference(message) for message in messages):
return messages
return [_drop_tool_reference_parts(message) for message in messages] # mutable-ok: pipelines mutate message lists
def _attempt_json_repair(s: str) -> Any | None:
"""
Attempt to repair truncated JSON produced by LLM tool calls.

View file

@ -1412,7 +1412,7 @@ def convert_to_gemini_tool_call_result(
)
except Exception as e:
verbose_logger.warning("Failed to process image in tool response: %s", e)
elif content_type in ("file", "input_file"):
elif content_type in ("file", "input_file"): # pyright: ignore[reportUnnecessaryContains] # loose runtime dict
# Extract file for inline_data (for tool results with PDF, audio, video, etc.)
file_data = content.get("file_data", "")
if not file_data:
@ -1564,14 +1564,23 @@ def convert_to_anthropic_tool_result(
}
"""
anthropic_content: (
str | list[AnthropicMessagesToolResultContent | AnthropicMessagesImageParam | AnthropicMessagesDocumentParam]
str
| list[
AnthropicMessagesToolResultContent
| AnthropicMessagesImageParam
| AnthropicMessagesDocumentParam
| ToolReference
]
) = ""
if isinstance(message["content"], str):
anthropic_content = message["content"]
elif isinstance(message["content"], list):
content_list: Final = message["content"]
anthropic_content_list: list[
AnthropicMessagesToolResultContent | AnthropicMessagesImageParam | AnthropicMessagesDocumentParam
AnthropicMessagesToolResultContent
| AnthropicMessagesImageParam
| AnthropicMessagesDocumentParam
| ToolReference
] = []
for content in content_list:
if content["type"] == "text":
@ -1614,6 +1623,8 @@ def convert_to_anthropic_tool_result(
original_content_element=content,
)
anthropic_content_list.append(cast(AnthropicMessagesImageParam, _anthropic_image_param))
elif content["type"] == "tool_reference":
anthropic_content_list.append(ToolReference(type="tool_reference", tool_name=content["tool_name"]))
elif content["type"] == "file":
file_content = cast(ChatCompletionFileObject, content)
_file_block = anthropic_process_openai_file_message(file_content)

View file

@ -8,11 +8,12 @@ import time
import traceback
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Any, Final, NoReturn, Protocol, TypeVar, cast
import anyio
import httpx
from pydantic import BaseModel
from pydantic import BaseModel, ValidationError
from typing_extensions import NotRequired, TypedDict
import litellm
@ -182,6 +183,23 @@ class _VertexChunkLike(Protocol):
candidates: Sequence[_VertexCandidateLike]
class _ParsedChunkHiddenParams(BaseModel):
provider_specific_fields: Mapping[str, object] | None = None
def _provider_hidden_params(chunk: object) -> Mapping[str, object] | None:
hidden: Final[object] = getattr(chunk, "_hidden_params", None)
if not isinstance(hidden, dict):
return None
try:
parsed: Final = _ParsedChunkHiddenParams.model_validate(hidden)
except ValidationError:
return None
if not parsed.provider_specific_fields:
return None
return MappingProxyType({"provider_specific_fields": dict(parsed.provider_specific_fields)})
class CustomStreamWrapper:
def __init__(
self,
@ -801,7 +819,7 @@ class CustomStreamWrapper:
except Exception as e:
raise e
def model_response_creator(self, chunk: dict | None = None, hidden_params: dict | None = None):
def model_response_creator(self, chunk: dict | None = None, hidden_params: Mapping[str, object] | None = None):
_model: Final = self._cached_model_name
_logging_obj_llm_provider: Final = self._cached_logging_llm_provider
@ -1504,7 +1522,7 @@ class CustomStreamWrapper:
def chunk_creator(self, chunk: Any):
if hasattr(chunk, "id"):
self.response_id = chunk.id
model_response = self.model_response_creator()
model_response = self.model_response_creator(hidden_params=_provider_hidden_params(chunk))
response_obj: dict[str, Any] = {}
try:
# return this for all models

View file

@ -24,10 +24,12 @@ from litellm._logging import verbose_proxy_logger
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
LiteLLMAnthropicMessagesAdapter,
is_provider_native_tool_dict,
)
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.llms.base_llm.guardrail_translation.utils import (
anthropic_tool_name,
anthropic_tool_names,
effective_scan_only_tool_results_for_guardrail,
effective_skip_system_message_for_guardrail,
effective_skip_tool_message_for_guardrail,
@ -360,7 +362,13 @@ class AnthropicMessagesHandler(BaseTranslation):
structured_messages: Final = [full_structured_messages[index] for index in scoped_message_indices]
tools_to_check: Final[list[ChatCompletionToolParam]] = (
[] if scan_only_tool_results else chat_completion_compatible_request.get("tools", [])
[]
if scan_only_tool_results
else [
tool
for tool in chat_completion_compatible_request.get("tools", [])
if not is_provider_native_tool_dict(tool)
]
)
# Step 1: Extract all text content and images
@ -419,7 +427,10 @@ class AnthropicMessagesHandler(BaseTranslation):
tool_name=anthropic_tool_name,
)
if scan_only_tool_results
else anthropic_tools
else [
*(tool for tool in data.get("tools") or [] if is_provider_native_tool_dict(tool)),
*anthropic_tools,
]
)
guardrailed_structured_messages: Final = guardrailed_inputs.get("structured_messages")
@ -677,12 +688,9 @@ class AnthropicMessagesHandler(BaseTranslation):
)
def extract_request_tool_names(self, data: dict) -> list[str]:
"""Extract tool names from Anthropic messages request (tools[].name)."""
names: Final[list[str]] = []
for tool in data.get("tools") or []:
if isinstance(tool, dict) and tool.get("name"):
names.append(str(tool["name"]))
return names
"""Extract every tool name in an Anthropic messages request: tools[].name, plus
tools[].function.name for OpenAI-format tools the bridge forwards verbatim."""
return [name for tool in data.get("tools") or [] for name in anthropic_tool_names(tool)]
@classmethod
def _extract_input_text_and_images(

View file

@ -1,8 +1,8 @@
import copy
import hashlib
import json
from collections.abc import AsyncIterator, Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast
from collections.abc import AsyncIterator, Iterator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypeVar, cast
import litellm
from litellm.llms.anthropic.experimental_pass_through.utils import (
@ -18,6 +18,22 @@ TOOL_NAME_PREFIX_LENGTH: Final = OPENAI_MAX_TOOL_NAME_LENGTH - TOOL_NAME_HASH_LE
PROVIDERS_PROXYING_AN_UNKNOWN_BACKEND: Final = frozenset({"litellm_proxy"})
_ANTHROPIC_TOOL_SCHEMA_KEYS: Final = frozenset(
{"name", "type", "input_schema", "description", "cache_control", "strict"}
)
def _is_openai_function_tool(tool: Mapping[str, object]) -> bool:
return tool.get("type") == "function" and "function" in tool
def is_provider_native_tool_dict(tool: Mapping[str, object]) -> bool:
if len(tool) != 1:
return False
key, value = next(iter(tool.items()))
return key not in _ANTHROPIC_TOOL_SCHEMA_KEYS and isinstance(value, dict)
def truncate_tool_name(name: str) -> str:
"""
Truncate tool names that exceed OpenAI's 64-character limit.
@ -126,7 +142,9 @@ from litellm.types.llms.openai import (
ChatCompletionToolMessage,
ChatCompletionToolParam,
ChatCompletionToolParamFunctionChunk,
ChatCompletionToolReferenceObject,
ChatCompletionUserMessage,
ToolMessageContentPart,
)
from litellm.types.utils import Choices, ModelResponse, StreamingChoices, Usage
@ -135,6 +153,8 @@ from .streaming_iterator import AnthropicStreamWrapper
if TYPE_CHECKING:
from litellm.types.llms.anthropic import ContentBlockContentBlockDict
ToolResultContent: TypeAlias = str | list[ToolMessageContentPart]
class AnthropicAdapter:
def __init__(self) -> None:
@ -412,90 +432,13 @@ class LiteLLMAnthropicMessagesAdapter:
self._add_cache_control_if_applicable(content, doc_obj, model)
new_user_content_list.append(doc_obj)
elif content.get("type") == "tool_result":
if "content" not in content:
tool_result = ChatCompletionToolMessage(
role="tool",
tool_call_id=content.get("tool_use_id", ""),
content="",
)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result)
elif isinstance(content.get("content"), str):
tool_result = ChatCompletionToolMessage(
role="tool",
tool_call_id=content.get("tool_use_id", ""),
content=str(content.get("content", "")),
)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result)
elif isinstance(content.get("content"), list):
# Combine all content items into a single tool message
# to avoid creating multiple tool_result blocks with the same ID
# (each tool_use must have exactly one tool_result)
content_items = list(content.get("content", []))
# Single-item text keeps the backward-compatible string format; a single
# image or document becomes a structured image_url part
if len(content_items) == 1:
c = content_items[0]
if isinstance(c, str):
tool_result = ChatCompletionToolMessage(
role="tool",
tool_call_id=content.get("tool_use_id", ""),
content=c,
)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result)
elif isinstance(c, dict):
if c.get("type") == "text":
tool_result = ChatCompletionToolMessage(
role="tool",
tool_call_id=content.get("tool_use_id", ""),
content=c.get("text", ""),
)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result)
elif c.get("type") in ("image", "document"):
image_part = self._tool_result_image_part(c.get("source"))
tool_result = ChatCompletionToolMessage(
role="tool",
tool_call_id=content.get("tool_use_id", ""),
content=[image_part] # mutable-ok: content must be a json list
if image_part
else "",
)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result)
else:
# For multiple content items, combine into a single tool message
# with list content to preserve all items while having one tool_use_id
combined_content_parts: list[
ChatCompletionTextObject | ChatCompletionImageObject
] = []
for c in content_items:
if isinstance(c, str):
combined_content_parts.append(ChatCompletionTextObject(type="text", text=c))
elif isinstance(c, dict):
if c.get("type") == "text":
combined_content_parts.append(
ChatCompletionTextObject(
type="text",
text=c.get("text", ""),
)
)
elif c.get("type") in ("image", "document"):
image_part = self._tool_result_image_part(c.get("source"))
if image_part:
combined_content_parts.append(image_part)
# Create a single tool message with combined content
if combined_content_parts:
tool_result = ChatCompletionToolMessage(
role="tool",
tool_call_id=content.get("tool_use_id", ""),
content=combined_content_parts,
)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result)
tool_result = ChatCompletionToolMessage(
role="tool",
tool_call_id=content.get("tool_use_id", ""),
content=self._tool_result_content(content.get("content")),
)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result)
if len(tool_message_list) > 0:
new_messages.extend(tool_message_list)
@ -771,6 +714,10 @@ class LiteLLMAnthropicMessagesAdapter:
new_tools.append(tool)
continue
if _is_openai_function_tool(tool) or is_provider_native_tool_dict(tool):
new_tools.append(cast(ChatCompletionToolParam, tool)) # cast-ok: passed through verbatim to provider
continue
raw_name = tool.get("name")
if raw_name is None or (isinstance(raw_name, str) and not str(raw_name).strip()):
original_name = f"litellm_unnamed_tool_{idx}"
@ -1032,7 +979,25 @@ class LiteLLMAnthropicMessagesAdapter:
anthropic_message_request: AnthropicMessagesRequest,
new_kwargs: ChatCompletionRequest,
) -> None:
"""Translate Anthropic thinking to either thinking or reasoning_effort."""
"""Translate Anthropic thinking to either thinking or reasoning_effort.
A Claude-family target keeps ``thinking`` verbatim, since every bridged provider serving one
speaks that param. Carrying its adaptive effort tier alongside takes two different params,
because the two are not interchangeable at the provider mapping below.
Bedrock takes ``output_config`` directly, which attaches the tier and leaves ``thinking``
alone. Every other bridged Claude target takes ``reasoning_effort``, and used to be sent no
tier at all, so an adaptive request arrived byte-identical whichever effort the caller
asked for. That tier stays a plain string there, since the summary it would otherwise be
wrapped with already travels inside the forwarded ``thinking`` block, and the wrapped dict
is rejected outright by some of these providers.
``reasoning_effort`` is not a substitute for ``output_config`` on the Bedrock side: an
application inference profile ARN resolves to neither, so the tier is dropped, and providers
that rebuild ``output_config`` from it overwrite a caller-set ``thinking.display`` doing so.
An adaptive request with no tier stays untouched either way, so the provider's own default
still applies.
"""
if "thinking" not in anthropic_message_request:
return
@ -1041,35 +1006,38 @@ class LiteLLMAnthropicMessagesAdapter:
return
model: Final = new_kwargs.get("model", "")
if self.is_anthropic_claude_model(model) or self.is_bedrock_arn_model(model):
is_bedrock_target: Final = model.startswith(("bedrock/", "converse/", "invoke/")) or self.is_bedrock_arn_model(
model
)
is_claude_target: Final = self.is_anthropic_claude_model(model) or self.is_bedrock_arn_model(model)
output_config: Final = anthropic_message_request.get("output_config")
if is_claude_target:
new_kwargs["thinking"] = thinking
# Adaptive thinking without its effort tier makes Bedrock Converse
# return zero reasoning blocks, so forward output_config (minus
# `format`, already translated to response_format) for Bedrock
# targets only: other bridged providers reject the raw param, and
# get_llm_provider strips the `bedrock/` prefix before this runs.
if model.startswith(("bedrock/", "converse/", "invoke/")) or self.is_bedrock_arn_model(model):
claude_output_config: Final = anthropic_message_request.get("output_config")
if isinstance(claude_output_config, dict):
effort_config: Final = {k: v for k, v in claude_output_config.items() if k != "format"}
if is_bedrock_target:
if isinstance(output_config, dict):
effort_config: Final = {k: v for k, v in output_config.items() if k != "format"}
if effort_config:
new_kwargs["output_config"] = effort_config # rebind-ok: out-param store like thinking above
return
thinking_type: Final = thinking.get("type") if isinstance(thinking, dict) else None
declared_effort: Final = (
output_config.get("effort") if thinking_type == "adaptive" and isinstance(output_config, dict) else None
)
if is_claude_target and not declared_effort:
return
reasoning_effort = self.translate_anthropic_thinking_to_reasoning_effort(cast(AnthropicThinkingParam, thinking))
reasoning_effort: Final = declared_effort or self.translate_anthropic_thinking_to_reasoning_effort(
cast(AnthropicThinkingParam, thinking)
)
if not reasoning_effort:
return
thinking_type: Final = thinking.get("type") if isinstance(thinking, dict) else None
# For adaptive thinking, override with output_config.effort if available
if thinking_type == "adaptive":
output_config: Final = anthropic_message_request.get("output_config")
if isinstance(output_config, dict) and output_config.get("effort"):
reasoning_effort = output_config["effort"]
new_kwargs["reasoning_effort"] = self._apply_reasoning_summary_wrapping(
reasoning_effort, cast(dict[str, object], thinking)
new_kwargs["reasoning_effort"] = (
reasoning_effort
if is_claude_target
else self._apply_reasoning_summary_wrapping(reasoning_effort, cast(dict[str, object], thinking))
)
def _translate_output_format_to_openai(
@ -1210,6 +1178,39 @@ class LiteLLMAnthropicMessagesAdapter:
return None
def _tool_result_content(self, raw_content: object) -> ToolResultContent:
if isinstance(raw_content, str):
return raw_content
if not isinstance(raw_content, list):
return ""
items: Final = cast(Sequence[object], raw_content) # cast-ok: untrusted client payload
parts: Final = tuple(part for part in (self._tool_result_part(item) for item in items) if part is not None)
match parts:
case ():
return ""
case ({"type": "text", "text": str(text)},):
return text
case _:
return list(parts) # mutable-ok: content must be a json list
def _tool_result_part(self, item: object) -> ToolMessageContentPart | None:
if isinstance(item, str):
return ChatCompletionTextObject(type="text", text=item)
if not isinstance(item, dict):
return None
block: Final = cast(Mapping[str, object], item) # cast-ok: untrusted client payload
match block.get("type"):
case "text":
return ChatCompletionTextObject(type="text", text=str(block.get("text") or ""))
case "image" | "document":
return self._tool_result_image_part(block.get("source"))
case "tool_reference":
return ChatCompletionToolReferenceObject(
type="tool_reference", tool_name=str(block.get("tool_name") or "")
)
case _:
return None
def _tool_result_image_part(self, image_source: object) -> ChatCompletionImageObject | None:
if not isinstance(image_source, dict):
return None

View file

@ -1,4 +1,6 @@
import os
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
import litellm
@ -23,6 +25,29 @@ def is_reasoning_auto_summary_enabled() -> bool:
return litellm.reasoning_auto_summary or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true"
_DECLARED_DEGRADATION_CHAINS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
{"max": ("max", "xhigh", "high"), "xhigh": ("xhigh", "high"), "minimal": ("minimal", "low")}
)
def _effort_from_declaration(model_info: ModelInfo, effort: str) -> str | None:
"""A declared level set is the WHOLE answer for this gate, so a level it omits degrades even
where a per-level flag would have allowed it. Honoring both would let /model_group/info and
this path disagree about the same entry. None means the entry declares nothing, and the flag
chain below decides as before.
A declaration that omits every level in a chain still lands on that chain's terminal, which can
itself be undeclared. Picking a nearer declared level instead would need a strength ordering,
and the advertisement order is presentation only by design, so the terminal stays the answer."""
from litellm.router_utils.reasoning_effort_capability import declared_reasoning_efforts
declared: Final = declared_reasoning_efforts(model_info)
if declared is None:
return None
chain: Final = _DECLARED_DEGRADATION_CHAINS[effort]
return next((level for level in chain if level in declared), chain[-1])
def normalize_reasoning_effort_value(
effort: str,
model: str,
@ -48,6 +73,10 @@ def normalize_reasoning_effort_value(
except Exception:
model_info = None
declared_effort: Final = _effort_from_declaration(model_info, effort) if model_info is not None else None
if declared_effort is not None:
return declared_effort
if effort == "max":
if model_info and model_info.get("supports_max_reasoning_effort"):
return "max"

View file

@ -4,6 +4,7 @@ from httpx._models import Headers, Response
import litellm
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_tool_reference_parts_from_tool_messages,
hoist_images_from_tool_messages,
)
from litellm.litellm_core_utils.prompt_templates.factory import (
@ -252,7 +253,8 @@ class AzureOpenAIConfig(BaseConfig):
litellm_params: dict,
headers: dict,
) -> dict:
azure_messages: Final = convert_to_azure_openai_messages(hoist_images_from_tool_messages(messages))
stripped_messages: Final = drop_tool_reference_parts_from_tool_messages(messages)
azure_messages: Final = convert_to_azure_openai_messages(hoist_images_from_tool_messages(stripped_messages))
return {
"model": model,
"messages": azure_messages,

View file

@ -40,6 +40,16 @@ class BaseAudioTranscriptionConfig(BaseConfig, ABC):
def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]:
pass
@property
def supports_subtitle_synthesis(self) -> bool:
"""
Opt-in for providers without a native srt/vtt response body: when True
and the user asked for response_format srt/vtt, the http handler
synthesizes the subtitle document from the word timestamps the
provider's TranscriptionResponse carries in `words`.
"""
return False
def get_complete_url(
self,
api_base: str | None,

View file

@ -209,9 +209,20 @@ def openai_tool_name(tool: object) -> str | None:
return flat_name if isinstance(flat_name, str) else None
def anthropic_tool_names(tool: object) -> tuple[str, ...]:
"""Every name a /v1/messages tool dict can act under: the flat Anthropic ``name`` plus
``function.name`` for OpenAI-format tools the bridge forwards verbatim. Allowlist checks
must see both, or a decoy flat name could smuggle a disallowed ``function.name`` through."""
if not isinstance(tool, dict):
return ()
function: Final = tool.get("function") if tool.get("type") == "function" else None
function_name: Final = function.get("name") if isinstance(function, dict) else None
return tuple(name for name in (tool.get("name"), function_name) if isinstance(name, str) and name)
def anthropic_tool_name(tool: object) -> str | None:
name: Final = tool.get("name") if isinstance(tool, dict) else None
return name if isinstance(name, str) else None
names: Final = anthropic_tool_names(tool)
return names[0] if names else None
def merge_returned_tools_into_request_tools(

View file

@ -25,6 +25,10 @@ from litellm.litellm_core_utils.agentic_loop_settings import (
validated_max_agentic_loops,
)
from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
SUBTITLE_RESPONSE_FORMATS,
synthesize_subtitle_document,
)
from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields
from litellm.litellm_core_utils.realtime_errors import realtime_error_event, websocket_close_reason
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
@ -1296,9 +1300,23 @@ class BaseLLMHTTPHandler:
api_key: str | None,
) -> TranscriptionResponse:
"""Shared logic for transforming audio transcription responses."""
return provider_config.transform_audio_transcription_response(
transformed: Final = provider_config.transform_audio_transcription_response(
raw_response=response,
)
if not provider_config.supports_subtitle_synthesis:
return transformed
requested_format: Final = optional_params.get("response_format")
if not isinstance(requested_format, str) or requested_format not in SUBTITLE_RESPONSE_FORMATS:
return transformed
document: Final = synthesize_subtitle_document(
words=transformed.get("words"),
response_format=requested_format,
)
if document is not None:
transformed.text = document
if "words" in transformed:
delattr(transformed, "words")
return transformed
def audio_transcriptions(
self,

View file

@ -11,7 +11,7 @@ Request format:
"input": {
"messages": [{"role": "user", "content": [{"text": "<prompt>"}]}]
},
"parameters": {"size": "1024*1024", ...}
"parameters": {"size": "1024*1024", "n": 1, ...}
}
Response format:
@ -19,7 +19,7 @@ Response format:
"output": {
"choices": [{"message": {"content": [{"image": "<url>"}]}}]
},
"usage": {"input_tokens": 0, "output_tokens": 0, "width": 1024, "height": 1024, "image_count": 1}
"usage": {"output_width": 1024, "output_height": 1024, "output_image_count": 1}
}
"""
@ -46,6 +46,8 @@ else:
DEFAULT_API_BASE: Final = "https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation"
CHAT_COMPATIBLE_MODE_PATH: Final = "/compatible-mode/v1"
# Maps OpenAI size strings (WxH) to DashScope size strings (W*H)
OPENAI_TO_DASHSCOPE_SIZE: Final[dict] = {
"256x256": "256*256",
@ -59,7 +61,8 @@ OPENAI_TO_DASHSCOPE_SIZE: Final[dict] = {
class DashScopeImageGenerationConfig(BaseImageGenerationConfig):
"""
Configuration for DashScope image generation (qwen-image-2.0, qwen-image-2.0-pro).
Configuration for DashScope image generation (qwen-image-2.0, qwen-image-2.0-pro,
qwen-image-3.0, qwen-image-3.0-pro).
"""
def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]:
@ -82,8 +85,8 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig):
if k == "size":
# Convert "WxH" → "W*H"
mapped["size"] = OPENAI_TO_DASHSCOPE_SIZE.get(v, v.replace("x", "*"))
elif k == "n":
mapped["image_count"] = v
else:
mapped[k] = v
return mapped
def get_complete_url(
@ -95,7 +98,10 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig):
litellm_params: dict,
stream: bool | None = None,
) -> str:
return api_base or get_secret_str("DASHSCOPE_API_BASE_IMAGE") or DEFAULT_API_BASE
image_api_base: Final = (
api_base if api_base and not api_base.rstrip("/").endswith(CHAT_COMPATIBLE_MODE_PATH) else None
)
return image_api_base or get_secret_str("DASHSCOPE_API_BASE_IMAGE") or DEFAULT_API_BASE
def validate_environment(
self,

View file

@ -4,6 +4,7 @@ from typing import Final
from httpx import Headers, Response
from litellm.litellm_core_utils.audio_utils.subtitle_utils import SUBTITLE_RESPONSE_FORMATS
from litellm.litellm_core_utils.audio_utils.utils import (
normalize_transcription_language_to_bcp47,
process_audio_file,
@ -48,6 +49,10 @@ class GeminiAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
) -> list[OpenAIAudioTranscriptionOptionalParams]: # mutable-ok: BaseAudioTranscriptionConfig signature
return ["language", "response_format", "timestamp_granularities"] # mutable-ok: base contract returns a list
@property
def supports_subtitle_synthesis(self) -> bool:
return True
def map_openai_params(
self,
non_default_params: Mapping[str, object],
@ -215,16 +220,17 @@ def _language_config(language: object) -> GeminiTranscriptionConfig:
return language_config
def _timestamp_config(timestamp_granularities: object) -> GeminiTranscriptionConfig:
if isinstance(timestamp_granularities, list) and "word" in timestamp_granularities:
return _WORD_TIMESTAMP_CONFIG
return _EMPTY_TRANSCRIPTION_CONFIG
def _timestamp_config(timestamp_granularities: object, response_format: object) -> GeminiTranscriptionConfig:
wants_word_timestamps: Final = (
isinstance(timestamp_granularities, list) and "word" in timestamp_granularities
) or (isinstance(response_format, str) and response_format in SUBTITLE_RESPONSE_FORMATS)
return _WORD_TIMESTAMP_CONFIG if wants_word_timestamps else _EMPTY_TRANSCRIPTION_CONFIG
def _build_transcription_config(optional_params: Mapping[str, object]) -> GeminiTranscriptionConfig:
transcription_config: Final[GeminiTranscriptionConfig] = {
**_language_config(optional_params.get("language")),
**_timestamp_config(optional_params.get("timestamp_granularities")),
**_timestamp_config(optional_params.get("timestamp_granularities"), optional_params.get("response_format")),
}
return transcription_config

View file

@ -8,7 +8,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
from litellm.litellm_core_utils.prompt_templates.image_handling import (
convert_url_to_base64,
)
from litellm.types.llms.openai import AllMessageValues, ChatCompletionFileObject
from litellm.types.llms.openai import AllMessageValues, ChatCompletionFileObject, ChatCompletionImageObject
from litellm.types.llms.vertex_ai import ContentType, PartType
from litellm.utils import supports_reasoning
@ -16,6 +16,13 @@ from ...vertex_ai.gemini.transformation import _gemini_convert_messages_with_his
from ...vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
def _image_url_fields(img_element: ChatCompletionImageObject) -> tuple[str | None, str | None, str | None]:
image_value: Final = img_element.get("image_url")
if isinstance(image_value, dict):
return image_value.get("url"), image_value.get("format"), image_value.get("detail")
return image_value, None, None
class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
"""
Reference: https://ai.google.dev/api/rest/v1beta/GenerationConfig
@ -118,16 +125,8 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
_parts: list[PartType] = []
for element in _message_content:
if element.get("type") == "image_url":
img_element = element
_image_url: str | None = None
format: str | None = None
detail: str | None = None
if isinstance(img_element.get("image_url"), dict):
_image_url = img_element["image_url"].get("url")
format = img_element["image_url"].get("format")
detail = img_element["image_url"].get("detail")
else:
_image_url = img_element.get("image_url")
img_element = cast(ChatCompletionImageObject, element) # cast-ok: runtime type tag checked
_image_url, format, detail = _image_url_fields(img_element)
if _image_url and "https://" in _image_url:
image_obj = convert_to_anthropic_image_obj(_image_url, format=format)
converted_image_url = convert_generic_image_chunk_to_openai_image_obj(image_obj)

View file

@ -292,7 +292,7 @@ class MistralConfig(OpenAIGPTConfig):
file_id = file_content.get("file", {}).get("file_id")
if file_id:
# Replace 'file' with 'file_id'
file_content["file_id"] = file_id
file_content["file_id"] = file_id # pyright: ignore[reportGeneralTypeIssues] # legacy in-place rewrite of the block shape
file_content.pop("file", None)
return messages

View file

@ -18,6 +18,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
_should_convert_tool_call_to_json_mode,
)
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_tool_reference_parts_from_tool_messages,
get_tool_call_names,
hoist_images_from_tool_messages,
)
@ -336,7 +337,8 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
self, messages: list[AllMessageValues], model: str, is_async: bool = False
) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]:
"""OpenAI no longer supports image_url as a string, so we need to convert it to a dict"""
hoisted_messages: Final = hoist_images_from_tool_messages(messages)
stripped_messages: Final = drop_tool_reference_parts_from_tool_messages(messages)
hoisted_messages: Final = hoist_images_from_tool_messages(stripped_messages)
async def _async_transform():
for message in hoisted_messages:

View file

@ -4,6 +4,11 @@ Shared utilities for the Soniox provider (https://soniox.com).
from typing import Any, Final
from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
SubtitleToken,
render_subtitle_tokens_as_srt,
render_subtitle_tokens_as_vtt,
)
from litellm.llms.base_llm.chat.transformation import BaseLLMException
# Soniox API base URL.
@ -109,121 +114,13 @@ def render_soniox_tokens(tokens: list[dict[str, Any]]) -> str:
return "".join(text_parts)
# ---------------------------------------------------------------------------
# SRT / VTT subtitle rendering
# ---------------------------------------------------------------------------
# Maximum number of tokens to group into a single subtitle cue.
_CUE_MAX_TOKENS: Final[int] = 15
# Maximum duration (in ms) for a single cue before forcing a break.
_CUE_MAX_DURATION_MS: Final[int] = 5000
def _format_timestamp_srt(ms: int) -> str:
"""Format milliseconds as SRT timestamp: HH:MM:SS,mmm"""
ms = max(ms, 0)
hours: Final = ms // 3_600_000
ms %= 3_600_000
minutes: Final = ms // 60_000
ms %= 60_000
seconds: Final = ms // 1_000
millis: Final = ms % 1_000
return f"{hours:02d}:{minutes:02d}:{seconds:02d},{millis:03d}"
def _format_timestamp_vtt(ms: int) -> str:
"""Format milliseconds as VTT timestamp: HH:MM:SS.mmm"""
ms = max(ms, 0)
hours: Final = ms // 3_600_000
ms %= 3_600_000
minutes: Final = ms // 60_000
ms %= 60_000
seconds: Final = ms // 1_000
millis: Final = ms % 1_000
return f"{hours:02d}:{minutes:02d}:{seconds:02d}.{millis:03d}"
def _group_tokens_into_cues(
tokens: list[dict[str, Any]],
) -> list[dict[str, Any]]:
"""
Group Soniox tokens into subtitle cues.
Each cue has:
- start_ms: int
- end_ms: int
- text: str
Grouping heuristics:
- A new cue starts when token count exceeds _CUE_MAX_TOKENS.
- A new cue starts when duration exceeds _CUE_MAX_DURATION_MS.
- A new cue starts when the speaker changes (if diarization is on).
- Tokens without timestamps are appended to the current cue.
"""
cues: Final[list[dict[str, Any]]] = []
current_tokens: list[str] = []
current_start: int | None = None
current_end: int | None = None
current_speaker: Any | None = None
def _flush() -> None:
if current_tokens and current_start is not None:
text: Final = "".join(current_tokens).strip()
if text:
cues.append(
{
"start_ms": current_start,
"end_ms": (current_end if current_end is not None else current_start),
"text": text,
}
)
for token in tokens:
start_ms = token.get("start_ms")
end_ms = token.get("end_ms")
text = token.get("text", "")
speaker = token.get("speaker")
# Skip tokens with no timestamp data entirely if we have no cue started
if start_ms is None and current_start is None:
continue
# Speaker change forces a new cue
if speaker is not None and speaker != current_speaker:
_flush()
current_tokens = []
current_start = start_ms
current_end = end_ms
current_speaker = speaker
current_tokens.append(text)
continue
# Duration or token count exceeded -> flush
should_break = False
if (
len(current_tokens) >= _CUE_MAX_TOKENS
or current_start is not None
and start_ms is not None
and (start_ms - current_start) >= _CUE_MAX_DURATION_MS
):
should_break = True
if should_break:
_flush()
current_tokens = []
current_start = start_ms
current_end = end_ms
current_tokens.append(text)
else:
if current_start is None:
current_start = start_ms
if end_ms is not None:
current_end = end_ms
current_tokens.append(text)
_flush()
return cues
def _soniox_token_to_subtitle_token(token: dict[str, Any]) -> SubtitleToken:
return SubtitleToken(
text=token.get("text", ""),
start_ms=token.get("start_ms"),
end_ms=token.get("end_ms"),
speaker=token.get("speaker"),
)
def render_soniox_tokens_as_srt(tokens: list[dict[str, Any]]) -> str:
@ -232,20 +129,7 @@ def render_soniox_tokens_as_srt(tokens: list[dict[str, Any]]) -> str:
Returns an empty string if no tokens have timestamp data.
"""
cues: Final = _group_tokens_into_cues(tokens)
if not cues:
return ""
lines: Final[list[str]] = []
for idx, cue in enumerate(cues, start=1):
start = _format_timestamp_srt(cue["start_ms"])
end = _format_timestamp_srt(cue["end_ms"])
lines.append(str(idx))
lines.append(f"{start} --> {end}")
lines.append(cue["text"])
lines.append("") # blank line between cues
return "\n".join(lines)
return render_subtitle_tokens_as_srt(tuple(_soniox_token_to_subtitle_token(token) for token in tokens))
def render_soniox_tokens_as_vtt(tokens: list[dict[str, Any]]) -> str:
@ -254,14 +138,4 @@ def render_soniox_tokens_as_vtt(tokens: list[dict[str, Any]]) -> str:
Returns the VTT header even if no cues are present.
"""
cues: Final = _group_tokens_into_cues(tokens)
lines: Final[list[str]] = ["WEBVTT", ""]
for cue in cues:
start = _format_timestamp_vtt(cue["start_ms"])
end = _format_timestamp_vtt(cue["end_ms"])
lines.append(f"{start} --> {end}")
lines.append(cue["text"])
lines.append("") # blank line between cues
return "\n".join(lines)
return render_subtitle_tokens_as_vtt(tuple(_soniox_token_to_subtitle_token(token) for token in tokens))

View file

@ -4,7 +4,8 @@ Translates from OpenAI's `/v1/chat/completions` to Together AI's `/v1/chat/compl
Docs: https://docs.together.ai/docs/chat-overview
"""
from collections.abc import Callable, Container, Coroutine
from collections.abc import Callable, Container, Coroutine, Mapping
from types import MappingProxyType
from typing import (
Final,
Literal,
@ -12,11 +13,13 @@ from typing import (
overload,
)
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_logger
from litellm.exceptions import UnsupportedParamsError
from litellm.types.llms.openai import AllMessageValues
from litellm.utils import supports_function_calling, supports_response_schema
from litellm.utils import supports_function_calling, supports_reasoning, supports_response_schema
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
@ -38,6 +41,34 @@ def _registry_verdict(model: str, flag: str, check: Callable[[str], bool]) -> bo
return None
ADJUSTABLE_EFFORT_REASONING_MODELS: Final = frozenset(
{
"openai/gpt-oss-120b",
"openai/gpt-oss-20b",
}
)
HYBRID_REASONING_MODELS: Final = frozenset(
{
"MiniMaxAI/MiniMax-M3",
"Qwen/Qwen3.5-9B",
"Qwen/Qwen3.6-Plus",
"deepseek-ai/DeepSeek-V4-Pro",
"moonshotai/Kimi-K3",
"nvidia/nemotron-3-ultra-550b-a55b",
"zai-org/GLM-5.2",
}
)
HIGH_MAX_EFFORT_MODEL_PREFIX: Final = "deepseek-ai/DeepSeek-V4-Pro"
EFFORT_TRANSLATION: Final = MappingProxyType({"minimal": "low", "xhigh": "high", "max": "high"})
HIGH_MAX_EFFORT_TRANSLATION: Final = MappingProxyType(
{"minimal": "high", "low": "high", "medium": "high", "xhigh": "max"}
)
class TogetherReasoningToggle(TypedDict):
enabled: ReadOnly[bool]
def _function_calling_verdict(model: str) -> bool | None:
return _registry_verdict(
model,
@ -83,6 +114,36 @@ def _tool_params_to_drop(passed_params: Container[str], model: str, drop_params:
)
def _supports_together_reasoning(model: str) -> bool:
if model in ADJUSTABLE_EFFORT_REASONING_MODELS or model in HYBRID_REASONING_MODELS:
return True
if model.startswith(HIGH_MAX_EFFORT_MODEL_PREFIX):
return True
return supports_reasoning(model, custom_llm_provider="together_ai")
def _adjustable_effort(effort: str, model: str) -> str:
if effort == "none":
verbose_logger.debug(
"together_ai model %s cannot disable reasoning; mapping reasoning_effort=none to low", model
)
return "low"
return EFFORT_TRANSLATION.get(effort, effort)
def _reasoning_effort_payload(effort: str, model: str) -> Mapping[str, object]:
if effort == "default":
return MappingProxyType({})
if model in ADJUSTABLE_EFFORT_REASONING_MODELS:
return MappingProxyType({"reasoning_effort": _adjustable_effort(effort, model)})
if effort == "none":
disable_reasoning: Final[TogetherReasoningToggle] = {"enabled": False}
return MappingProxyType({"reasoning": disable_reasoning})
if model.startswith(HIGH_MAX_EFFORT_MODEL_PREFIX):
return MappingProxyType({"reasoning_effort": HIGH_MAX_EFFORT_TRANSLATION.get(effort, effort)})
return MappingProxyType({"reasoning_effort": EFFORT_TRANSLATION.get(effort, effort)})
def _drop_response_format(passed_params: Container[str], model: str, drop_params: bool) -> bool:
if "response_format" not in passed_params:
return False
@ -153,6 +214,15 @@ class TogetherAIChatConfig(OpenAIGPTConfig):
return super()._transform_messages(stripped, model, is_async=True)
return super()._transform_messages(stripped, model, is_async=False)
def get_supported_openai_params(self, model: str) -> list: # mutable-ok: inherited contract
supported_params: Final = super().get_supported_openai_params(model)
if not _supports_together_reasoning(model):
return supported_params
return [ # mutable-ok: the inherited contract returns a plain list; building fresh avoids mutating the base class's value
*supported_params,
"reasoning_effort",
]
def map_openai_params(
self,
non_default_params: dict,
@ -165,4 +235,10 @@ class TogetherAIChatConfig(OpenAIGPTConfig):
mapped_openai_params.pop(param)
if _drop_response_format(mapped_openai_params, model, drop_params):
mapped_openai_params.pop("response_format")
effort: Final = mapped_openai_params.get("reasoning_effort")
if not isinstance(effort, str):
return mapped_openai_params
mapped_openai_params.pop("reasoning_effort")
for key, value in _reasoning_effort_payload(effort, model).items():
mapped_openai_params.setdefault(key, value)
return mapped_openai_params

View file

@ -3,6 +3,7 @@ Handles calculating cost for together ai models
"""
import re
from collections.abc import Mapping
from typing import Final
from litellm.constants import (
@ -18,6 +19,12 @@ from litellm.constants import (
from litellm.types.utils import CallTypes
def has_together_registry_pricing(model: str, cost_map: Mapping[str, object]) -> bool:
stripped: Final = model.removeprefix("together_ai/")
entry: Final = cost_map.get(f"together_ai/{stripped}")
return isinstance(entry, Mapping) and "input_cost_per_token" in entry
# Extract the number of billion parameters from the model name
# only used for together_computer LLMs
def get_model_params_and_category(model_name, call_type: CallTypes) -> str:

View file

@ -9315,6 +9315,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-kimi-k3-through-fireworks-ai-on-microsoft-foundry/4540187",
"supported_modalities": [
"text",
@ -14647,6 +14652,22 @@
"/v1/images/generations"
]
},
"dashscope/qwen-image-3.0": {
"litellm_provider": "dashscope",
"mode": "image_generation",
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supported_endpoints": [
"/v1/images/generations"
]
},
"dashscope/qwen-image-3.0-pro": {
"litellm_provider": "dashscope",
"mode": "image_generation",
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supported_endpoints": [
"/v1/images/generations"
]
},
"databricks/databricks-bge-large-en": {
"cache_creation_input_token_cost": 1.0003e-07,
"cache_read_input_token_cost": 1.0003e-07,
@ -31881,6 +31902,11 @@
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://platform.kimi.ai/docs/pricing/chat-k3",
"supports_function_calling": true,
"supports_reasoning": true,
@ -36489,6 +36515,14 @@
"litellm_provider": "perplexity",
"mode": "responses",
"output_cost_per_token": 1.5e-05,
"reasoning_effort_levels": [
"minimal",
"low",
"medium",
"high",
"xhigh",
"max"
],
"source": "https://docs.perplexity.ai/docs/agent-api/models",
"supports_web_search": true,
"supports_reasoning": true,
@ -38791,14 +38825,14 @@
"supports_reasoning": true
},
"together_ai/Qwen/Qwen3.7-Max": {
"cache_read_input_token_cost": 1.3e-07,
"input_cost_per_token": 1.25e-06,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 2.5e-06,
"litellm_provider": "together_ai",
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"max_tokens": 1000000,
"mode": "chat",
"output_cost_per_token": 3.75e-06,
"output_cost_per_token": 7.5e-06,
"source": "https://docs.together.ai/docs/serverless-models",
"supports_prompt_caching": true
},
@ -38970,6 +39004,11 @@
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.together.ai/docs/serverless-models",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
@ -39053,6 +39092,24 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"together_ai/zai-org/GLM-5.3-Flash": {
"cache_read_input_token_cost": 3e-08,
"input_cost_per_token": 1.5e-07,
"litellm_provider": "together_ai",
"max_input_tokens": 1048575,
"max_output_tokens": 1048575,
"max_tokens": 1048575,
"mode": "chat",
"output_cost_per_token": 5e-07,
"source": "https://docs.together.ai/docs/serverless-models",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"tts-1": {
"input_cost_per_character": 1.5e-05,
"litellm_provider": "openai",
@ -51740,6 +51797,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -51804,6 +51866,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -51820,6 +51887,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.25e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -51836,6 +51908,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -52008,6 +52085,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.25e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -52024,6 +52106,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,

View file

@ -34,6 +34,7 @@ from mcp.types import (
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl, BaseModel
from typing_extensions import ReadOnly
import litellm
from litellm._logging import verbose_logger
@ -72,6 +73,7 @@ from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
MCPPerUserTokenCache,
mcp_per_user_token_cache,
resolve_mcp_auth,
resolved_token_header,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
_redact_mcp_resource_url,
@ -99,6 +101,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_
build_token_exchanger,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
DEFAULT_CREDENTIAL_HEADER,
AuthorizationCodeConfig,
ClientCredentialsConfig,
CredError,
@ -153,6 +156,8 @@ from litellm.types.mcp import (
MCPAuth,
MCPStdioConfig,
MCPTokenEndpointAuthMethod,
has_header,
without_header,
)
from litellm.types.mcp_server.mcp_server_manager import (
MCPInfo,
@ -349,6 +354,7 @@ class MCPServerConfig(TypedDict, total=False):
audience: str
subject_token_type: str
upstream_resource: str
upstream_token_header: ReadOnly[str]
id_jag_resource_token_endpoint: str
id_jag_resource: str
client_private_key: str
@ -828,18 +834,6 @@ def _should_strip_caller_authorization(
)
def _without_authorization(
headers: dict[str, str] | None,
) -> dict[str, str] | None:
"""A copy of ``headers`` with any ``Authorization`` key removed (case-insensitive), or
None if nothing remains. Drops only the credential, keeping other forwarded headers.
"""
if not headers:
return None
filtered: Final = {k: v for k, v in headers.items() if k.lower() != "authorization"}
return filtered or None
def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.
@ -914,7 +908,9 @@ def _resolve_openapi_tool_auth(
if isinstance(per_server, dict):
authorization: Final = next((v for k, v in per_server.items() if k.lower() == "authorization"), None)
merged: Final = merge_mcp_headers(extra_headers=forwarded, static_headers=_without_authorization(per_server))
merged: Final = merge_mcp_headers(
extra_headers=forwarded, static_headers=without_header(per_server, DEFAULT_CREDENTIAL_HEADER)
)
if authorization is None:
byok: Final = _format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None
return byok, merged, mcp_auth_header
@ -981,7 +977,7 @@ def _client_forwarded_authorization_headers(
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
):
return _without_authorization(extra_headers)
return without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER)
return extra_headers
@ -994,7 +990,7 @@ def _take_forwarded_authorization(
if not headers:
return None, headers
value: Final = next((v for k, v in headers.items() if k.lower() == "authorization"), None)
return value, _without_authorization(headers)
return value, without_header(headers, DEFAULT_CREDENTIAL_HEADER)
def _passthrough_token_from_mcp_auth_header(
@ -2166,6 +2162,7 @@ class MCPServerManager:
DEFAULT_SUBJECT_TOKEN_TYPE,
),
upstream_resource=server_config.get("upstream_resource", None),
upstream_token_header=server_config.get("upstream_token_header", None),
# ID-JAG fields
id_jag_resource_token_endpoint=server_config.get("id_jag_resource_token_endpoint", None),
id_jag_resource=server_config.get("id_jag_resource", None),
@ -2698,6 +2695,7 @@ class MCPServerManager:
or (credentials_dict.get("subject_token_type") if credentials_dict else None)
or DEFAULT_SUBJECT_TOKEN_TYPE,
upstream_resource=(credentials_dict.get("upstream_resource") if credentials_dict else None),
upstream_token_header=(credentials_dict.get("upstream_token_header") if credentials_dict else None),
# ID-JAG fields — read from credentials JSON blob
id_jag_resource_token_endpoint=(
credentials_dict.get("id_jag_resource_token_endpoint") if credentials_dict else None
@ -3525,10 +3523,9 @@ class MCPServerManager:
case Ok(auth):
# NoOpAuth has no header_name and so never conflicts.
header_name: Final[str | None] = getattr(auth, "header_name", None)
conflicts: Final = bool(
header_name and extra_headers and any(key.lower() == header_name.lower() for key in extra_headers)
)
if not conflicts:
if header_name is None or not extra_headers:
return auth, extra_headers
if not has_header(extra_headers, header_name):
return auth, extra_headers
if isinstance(
spec.config,
@ -3540,9 +3537,10 @@ class MCPServerManager:
# guardrail such as MCPJWTSigner, static_headers, or any other injected
# Authorization must NOT shadow it (otherwise the upstream gets e.g. the
# signer's JWT instead of the minted token and rejects it, and for M2M the
# one-shot 401 refetch is lost with it). Drop the conflicting header so the
# resolved token reaches upstream.
return auth, _without_authorization(extra_headers)
# one-shot 401 refetch is lost with it). Drop only the header the resolved
# credential is about to occupy, so a static credential the operator aimed at a
# DIFFERENT header still reaches upstream.
return auth, without_header(extra_headers, header_name)
# Other modes: an Authorization already supplied via extra_headers (a forwarded caller
# header or static_headers) is intentional and wins; v1 applies those last.
return None, extra_headers
@ -3650,6 +3648,7 @@ class MCPServerManager:
):
spec = None
auth_value: Final = await resolve_mcp_auth(resolved_server, mcp_auth_header) if spec is None else None
auth_header_name: Final = resolved_token_header(resolved_server, mcp_auth_header) if spec is None else None
# Create sampling and elicitation callbacks for this client
sampling_cb = (
@ -3758,6 +3757,7 @@ class MCPServerManager:
transport_type=transport,
auth_type=resolved_server.auth_type,
auth_value=auth_value,
auth_header_name=auth_header_name,
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
extra_headers=extra_headers,
aws_auth=aws_auth,
@ -5306,7 +5306,7 @@ class MCPServerManager:
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
):
extra_headers = _without_authorization(extra_headers)
extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER)
elif mcp_server.is_client_forwarded_token:
extra_headers = _client_forwarded_authorization_headers(
mcp_server=mcp_server,

View file

@ -7,6 +7,7 @@ with ``client_id``, ``client_secret``, and ``token_url``.
import asyncio
import hashlib
from collections.abc import Mapping
from typing import TYPE_CHECKING, Final
import httpx
@ -313,9 +314,26 @@ async def resolve_mcp_auth(
1. ``mcp_auth_header`` — per-request/per-user override
2. OAuth2 client_credentials token — auto-fetched and cached
3. ``server.authentication_token`` — static token from config/DB
``resolved_token_header`` answers, for the same two inputs, which header the value belongs in.
"""
if mcp_auth_header:
return mcp_auth_header
if server.has_client_credentials:
return await mcp_oauth2_token_cache.async_get_token(server)
return server.authentication_token
def resolved_token_header(
server: "MCPServer",
mcp_auth_header: str | Mapping[str, str] | None = None,
) -> str | None:
"""Which upstream header the value ``resolve_mcp_auth`` just returned belongs in.
``None`` means keep the auth_type default. A caller-supplied ``mcp_auth_header`` is the caller's
own credential aimed at the slot the upstream normally uses, so it never moves; only the values
the gateway resolved from its own config (the minted M2M token, the static token) follow
``upstream_token_header``. Same inputs and same branch order as ``resolve_mcp_auth``, so the two
cannot disagree about which case they are in.
"""
return None if mcp_auth_header else server.upstream_token_header

View file

@ -47,12 +47,14 @@ def sanitize_openapi_tool_name(raw_name: str) -> str:
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.url_utils import async_safe_get
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
from litellm.types.mcp import credential_redirect_hook, custom_credential_slot
class _OpenAPIJSONSchema(TypedDict, total=False):
@ -119,6 +121,10 @@ _request_resolved_auth_headers: Final[contextvars.ContextVar[dict[str, str] | No
"_request_resolved_auth_headers", default=None
)
_request_upstream_url: Final[contextvars.ContextVar[str | None]] = contextvars.ContextVar(
"_request_upstream_url", default=None
)
def _sanitize_path_parameter_value(param_value: object, param_name: str) -> str:
"""Ensure path params cannot introduce directory traversal."""
@ -349,6 +355,35 @@ def build_input_schema(operation: _OpenAPIOperation) -> dict[str, object]:
}
async def _drop_credential_across_origin(request: httpx.Request) -> None:
"""Apply this request's cross-origin credential guard, if it needs one.
Reads the per-request context rather than closing over it so the hook is one stable object, which
keeps the guarded client cacheable. A closure would key a new entry per call, and the handler it
built would never be closed.
"""
guard: Final = credential_redirect_hook(
_request_upstream_url.get() or "", custom_credential_slot(_request_resolved_auth_headers.get())
)
if guard is not None:
await guard(request)
def _upstream_client() -> AsyncHTTPHandler:
"""The HTTP client for one upstream call, guarded when a credential rides a custom slot.
A resolved credential outside ``Authorization`` is not stripped across origins by the client
itself, so this arm installs the same hook the MCP client uses. Both variants come from the
shared cache, so a guarded call reuses its connection pool like any other.
"""
if custom_credential_slot(_request_resolved_auth_headers.get()) is None:
return get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
return get_async_httpx_client(
llm_provider=httpxSpecialProvider.MCP,
params={"event_hooks": {"request": [_drop_credential_across_origin]}},
)
def _merge_openapi_tool_request_headers(
static_headers: dict[str, str],
) -> dict[str, str]:
@ -510,8 +545,9 @@ def create_tool_function(
except (json.JSONDecodeError, TypeError):
json_body = {"data": body_value}
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
client: Final = _upstream_client()
upstream: Final = server_label or f"{original_method.upper()} {path}"
url_token: Final = _request_upstream_url.set(url)
try:
if original_method == "get":
@ -529,6 +565,8 @@ def create_tool_function(
except MaskedHTTPStatusError as e:
_raise_for_upstream_failure(e.response, upstream, relays_upstream_auth)
raise
finally:
_request_upstream_url.reset(url_token)
_raise_for_upstream_failure(response, upstream, relays_upstream_auth)
return response.text

View file

@ -21,6 +21,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
Result,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
DEFAULT_CREDENTIAL_HEADER,
Ambient,
ApiKeyConfig,
ApiKeySource,
@ -35,6 +36,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ClientCredentialsConfig,
ClientSecretAuth,
CredError,
HeaderCarrier,
IdJagConfig,
NoneConfig,
PassthroughConfig,
@ -45,9 +47,11 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
Subject,
TokenExchangeConfig,
parse_auth_spec_kind,
validate_header_name,
)
__all__ = [
"DEFAULT_CREDENTIAL_HEADER",
"Ambient",
"ApiKeyConfig",
"ApiKeySource",
@ -63,6 +67,7 @@ __all__ = [
"ClientSecretAuth",
"CredError",
"Error",
"HeaderCarrier",
"IdJagConfig",
"NoOpAuth",
"NoneConfig",
@ -78,4 +83,5 @@ __all__ = [
"TokenExchangeConfig",
"UpstreamCredentialProvider",
"parse_auth_spec_kind",
"validate_header_name",
]

View file

@ -20,6 +20,7 @@ from typing_extensions import assert_never
from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
DEFAULT_CREDENTIAL_HEADER,
ApiKeyConfig,
AuthorizationCodeConfig,
ClientAuth,
@ -45,6 +46,15 @@ _TOKEN_EXCHANGE_SUBJECT_TOKEN_DEFAULT: Final = "urn:ietf:params:oauth:token-type
_ID_JAG_SUBJECT_TOKEN_DEFAULT: Final = "urn:ietf:params:oauth:token-type:id_token"
def token_header(server: MCPServer) -> str:
"""The upstream header this server's resolved credential occupies.
One owner for every arm, so no spec builder spells the default itself and a server can never
hand two arms different answers.
"""
return server.upstream_token_header or DEFAULT_CREDENTIAL_HEADER
def to_subject(user_api_key_auth: UserAPIKeyAuth | None, subject_token: str | None) -> Subject:
"""Map v1's authenticated principal onto the resolver's Subject.
@ -122,7 +132,7 @@ def _oauth2_spec(server: MCPServer, resource: str) -> ServerSpec | None:
return ServerSpec(
server_id=server.server_id,
resource=resource,
config=AuthorizationCodeConfig(),
config=AuthorizationCodeConfig(header_name=token_header(server)),
)
return None
@ -140,6 +150,7 @@ def _client_credentials_spec(server: MCPServer, resource: str) -> ServerSpec:
server_id=server.server_id,
resource=resource,
config=ClientCredentialsConfig(
header_name=token_header(server),
client_id=server.client_id,
client_secret=SecretStr(server.client_secret) if server.client_secret else None,
token_url=server.effective_token_url,
@ -173,6 +184,7 @@ def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec | None:
server_id=server.server_id,
resource=resource,
config=TokenExchangeConfig(
header_name=token_header(server),
profile=profile,
subject_token_type=server.subject_token_type or DEFAULT_SUBJECT_TOKEN_TYPE,
token_exchange_endpoint=endpoint,
@ -206,7 +218,7 @@ def _shared_key_spec(
server_id=server.server_id,
resource=resource,
config=ApiKeyConfig(
header_name=header_name,
header_name=server.upstream_token_header or header_name,
value_prefix=value_prefix,
key_source=SharedKey(value=SecretStr(value)),
),
@ -231,6 +243,7 @@ def _id_jag_spec(server: MCPServer, resource: str) -> ServerSpec | None:
server_id=server.server_id,
resource=resource,
config=IdJagConfig(
header_name=token_header(server),
org_token_endpoint=org_token_endpoint,
resource_token_endpoint=resource_token_endpoint,
client_id=client_id,

View file

@ -50,6 +50,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ClientCredentialsConfig,
CredError,
HeaderCarrier,
)
@ -328,14 +329,21 @@ class ClientCredentialsBearerAuth(httpx.Auth):
refetch fails, or the retried request 401s again, the upstream's response stands.
"""
def __init__(self, access_token: str, refetch: Callable[[str], Awaitable[str | None]]) -> None:
self.header_name = "Authorization"
def __init__(
self,
access_token: str,
refetch: Callable[[str], Awaitable[str | None]],
carrier: HeaderCarrier,
) -> None:
self._carrier = carrier
self.header_name = carrier.header_name
self._access_token = SecretStr(access_token)
self._refetch = refetch
async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx.Request, httpx.Response]:
token: Final = self._access_token.get_secret_value()
request.headers[self.header_name] = f"Bearer {token}"
name, value = self._carrier.header(token)
request.headers[name] = value
response: Final = yield request
if response.status_code != 401:
return
@ -343,7 +351,8 @@ class ClientCredentialsBearerAuth(httpx.Auth):
if fresh is None:
return
self._access_token = SecretStr(fresh)
request.headers[self.header_name] = f"Bearer {fresh}"
fresh_name, fresh_value = self._carrier.header(fresh)
request.headers[fresh_name] = fresh_value
yield request
def sync_auth_flow(self, request: httpx.Request) -> Generator[httpx.Request, httpx.Response, None]:

View file

@ -145,8 +145,8 @@ class UpstreamCredentialProvider:
return await self._token_exchange(subject, server, config)
case IdJagConfig() as config:
return await self._id_jag(subject, server, config)
case AuthorizationCodeConfig():
return await self._authorization_code(subject, server)
case AuthorizationCodeConfig() as config:
return await self._authorization_code(subject, server, config)
case AwsSigV4Config():
return _not_implemented(AuthSpecKind.aws_sigv4)
assert_never(server.config)
@ -284,15 +284,19 @@ class UpstreamCredentialProvider:
match await self._exchanged_tokens.get_or_compute(slot, _exchange, fingerprint=fingerprint):
case Ok(access_token):
return Ok(StaticHeaderAuth(f"Bearer {access_token}"))
header_name, header_value = config.header(access_token)
return Ok(StaticHeaderAuth(header_value, header_name=header_name))
case Error(err):
return Error(err)
async def _authorization_code(self, subject: Subject, server: ServerSpec) -> Result[StaticHeaderAuth, CredError]:
async def _authorization_code(
self, subject: Subject, server: ServerSpec, config: AuthorizationCodeConfig
) -> Result[StaticHeaderAuth, CredError]:
token: Final = await self._authz_token(subject, server)
if token is None:
return Error(CredError.of_unauthorized("Authorization required: complete the OAuth flow for this server."))
return Ok(StaticHeaderAuth(f"Bearer {token.access_token}", header_name="Authorization"))
header_name, header_value = config.header(token.access_token)
return Ok(StaticHeaderAuth(header_value, header_name=header_name))
async def _client_credentials(
self, server_id: str, config: ClientCredentialsConfig
@ -307,7 +311,7 @@ class UpstreamCredentialProvider:
match await self._client_credentials_source.get(server_id, config):
case Ok(token):
refetch: Final = partial(self._client_credentials_source.refetch, server_id, config)
return Ok(ClientCredentialsBearerAuth(token.access_token, refetch))
return Ok(ClientCredentialsBearerAuth(token.access_token, refetch, config))
case Error(err):
return Error(err)
@ -332,7 +336,8 @@ class UpstreamCredentialProvider:
inbound.get_secret_value(), server, config, tenant_id=subject.tenant_id
):
case Ok(token):
return Ok(StaticHeaderAuth(f"Bearer {token.access_token}", header_name="Authorization"))
header_name, header_value = config.header(token.access_token)
return Ok(StaticHeaderAuth(header_value, header_name=header_name))
case Error(err):
return Error(err)

View file

@ -31,7 +31,7 @@ from enum import Enum
from typing import Annotated, Final, Literal
from expression import case, tag, tagged_union
from pydantic import BaseModel, ConfigDict, Field, SecretStr
from pydantic import BaseModel, ConfigDict, Field, SecretStr, field_validator
from typing_extensions import assert_never
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
@ -39,7 +39,11 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
Ok,
Result,
)
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE
from litellm.types.mcp import (
DEFAULT_CREDENTIAL_HEADER,
DEFAULT_SUBJECT_TOKEN_TYPE,
normalize_upstream_header_name,
)
class AuthSpecKind(str, Enum):
@ -161,7 +165,52 @@ class CredError:
assert_never(self.tag)
class AuthorizationCodeConfig(BaseModel):
def validate_header_name(raw: str) -> Result[str, CredError]:
"""``normalize_upstream_header_name`` with this package's error-as-value policy.
The grammar itself lives in ``litellm.types.mcp`` so the v1 model, the management endpoint and
this vocabulary all judge a header name the same way while each keeps its own failure shape.
"""
normalized: Final = normalize_upstream_header_name(raw)
if normalized is None:
return Error(CredError.of_misconfigured(f"invalid upstream header name: {raw!r}"))
return Ok(normalized)
class HeaderCarrier(BaseModel):
"""Where a resolved credential is written upstream, and how its value is formatted.
``Authorization: Bearer`` is only OAuth's *default* conveyance (RFC 6750 section 2.1), not its
only one: an ESB or API gateway commonly terminates its own credential in a private header while
a second credential passes through to the origin, so a credential has to be able to say which
slot it owns. Modeled like OpenAPI's apiKey scheme, so any upstream convention is expressible
(Authorization + Bearer, a raw value on X-API-Key, Ocp-Apim-Subscription-Key, esb-oauth, ...).
Every config whose credential the gateway mints or holds inherits this, so no resolver arm names
a header itself and the conflict rule in ``_resolve_v2_auth`` can always ask the auth object
which slot it is about to occupy. ``passthrough`` deliberately does not: it forwards the
caller's own credential into the slot the caller used, and mints nothing to place.
"""
model_config = ConfigDict(frozen=True)
header_name: str = DEFAULT_CREDENTIAL_HEADER
value_prefix: str = "Bearer"
@field_validator("header_name")
@classmethod
def _check_header_name(cls, value: str) -> str:
match validate_header_name(value):
case Ok(name):
return name
case Error(err):
raise ValueError(err.summary)
def header(self, value: str) -> tuple[str, str]:
formatted: Final = f"{self.value_prefix} {value}" if self.value_prefix else value
return self.header_name, formatted
class AuthorizationCodeConfig(HeaderCarrier):
"""Per-user 3LO; the gateway is the OAuth client and stores the user's token.
Endpoints are discovered (RFC 9728 -> RFC 8414) and the client is registered via DCR
@ -179,7 +228,7 @@ class AuthorizationCodeConfig(BaseModel):
token_url: str | None = None
class ClientCredentialsConfig(BaseModel):
class ClientCredentialsConfig(HeaderCarrier):
"""M2M service account; one upstream identity for every user.
Fields are optional so the config can be built incomplete: a value may be supplied at
@ -203,7 +252,7 @@ class ClientCredentialsConfig(BaseModel):
token_endpoint_auth_method: Literal["client_secret_post", "client_secret_basic"] | None = None
class TokenExchangeConfig(BaseModel):
class TokenExchangeConfig(HeaderCarrier):
"""OBO: swap the caller's live inbound token for a token bound to the upstream's audience. The
gateway authenticates to the exchange endpoint as an OAuth client (`client_id`/`client_secret`);
the inbound token is sent only to that endpoint, never to the upstream.
@ -255,7 +304,7 @@ class ClientSecretAuth(BaseModel):
ClientAuth = Annotated[PrivateKeyJwtAuth | ClientSecretAuth, Field(discriminator="source")]
class IdJagConfig(BaseModel):
class IdJagConfig(HeaderCarrier):
"""draft-ietf-oauth-identity-assertion-authz-grant (Okta "AI agent token exchange").
Two legs: leg 1 is an RFC 8693 token exchange at the IdP org AS (`org_token_endpoint`) that
@ -297,23 +346,16 @@ class Byok(BaseModel):
ApiKeySource = Annotated[SharedKey | Byok, Field(discriminator="source")]
class ApiKeyConfig(BaseModel):
class ApiKeyConfig(HeaderCarrier):
"""A fixed credential injected as a header. The value is shared (in config) or seeded
per-user (pulled from the store); `header_name` and `value_prefix` say where and how it is
written, modeled like OpenAPI's apiKey scheme so any upstream convention is expressible
(Authorization + Bearer, a raw value on X-API-Key, Ocp-Apim-Subscription-Key, etc.).
per-user (pulled from the store); the inherited `header_name` and `value_prefix` say where
and how it is written.
"""
model_config = ConfigDict(frozen=True)
kind: Literal[AuthSpecKind.api_key] = AuthSpecKind.api_key
header_name: str = "Authorization"
value_prefix: str = "Bearer"
key_source: ApiKeySource
def header(self, value: str) -> tuple[str, str]:
formatted: Final = f"{self.value_prefix} {value}" if self.value_prefix else value
return self.header_name, formatted
class PassthroughConfig(BaseModel):
"""Client-driven upstream OAuth; the gateway forwards the client's upstream token."""

View file

@ -433,7 +433,6 @@ if MCP_AVAILABLE:
_client_forwarded_authorization_headers,
_resolve_openapi_tool_auth,
_should_strip_caller_authorization,
_without_authorization,
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
@ -452,6 +451,7 @@ if MCP_AVAILABLE:
split_server_prefix_from_name,
strip_known_server_prefix,
)
from litellm.types.mcp import DEFAULT_CREDENTIAL_HEADER, without_header
######################################################
############ MCP Tools List REST API Response Object #
@ -1733,7 +1733,7 @@ if MCP_AVAILABLE:
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
):
extra_headers = _without_authorization(extra_headers)
extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER)
elif is_client_forwarded_mode:
if not withhold_forwarded_authorization:
extra_headers = _client_forwarded_authorization_headers(

View file

@ -15038,6 +15038,17 @@
}
],
"title": "Upstream Resource"
},
"upstream_token_header": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Upstream Token Header"
}
},
"title": "MCPCredentials",
@ -17518,6 +17529,17 @@
}
],
"title": "Upstream Resource"
},
"upstream_token_header": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Upstream Token Header"
}
},
"title": "MCPCredentials",
@ -20352,6 +20374,17 @@
}
],
"title": "Upstream Resource"
},
"upstream_token_header": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Upstream Token Header"
}
},
"title": "MCPCredentials",
@ -23699,6 +23732,17 @@
}
],
"title": "Upstream Resource"
},
"upstream_token_header": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Upstream Token Header"
}
},
"title": "MCPCredentials",

View file

@ -686,6 +686,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
aws_profile_name: Final = self.optional_params.get("aws_profile_name", None)
aws_web_identity_token: Final = self.optional_params.get("aws_web_identity_token", None)
aws_sts_endpoint: Final = self.optional_params.get("aws_sts_endpoint", None)
aws_external_id: Final = self.optional_params.get("aws_external_id", None)
### SET REGION NAME ###
aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(
@ -702,6 +703,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
return credentials, aws_region_name

View file

@ -34,6 +34,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail):
aws_role_name=litellm_params.aws_role_name,
aws_web_identity_token=litellm_params.aws_web_identity_token,
aws_sts_endpoint=litellm_params.aws_sts_endpoint,
aws_external_id=litellm_params.aws_external_id,
aws_bedrock_runtime_endpoint=litellm_params.aws_bedrock_runtime_endpoint,
experimental_use_latest_role_message_only=litellm_params.experimental_use_latest_role_message_only,
only_scan_new_messages=litellm_params.only_scan_new_messages or False,

View file

@ -204,6 +204,7 @@ if MCP_AVAILABLE:
MCP_ADMIN_CONFIG_CREDENTIAL_KEYS,
MCPAuth,
MCPCredentials,
normalize_upstream_header_name,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -239,9 +240,26 @@ if MCP_AVAILABLE:
detail={"error": error_messages_text},
)
def _validate_upstream_token_header(payload: McpServerPayloadLike) -> None:
credentials: Final = getattr(payload, "credentials", None)
raw: Final = credentials.get("upstream_token_header") if isinstance(credentials, dict) else None
if not isinstance(raw, str) or raw == "":
return
if normalize_upstream_header_name(raw) is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": (
f"Invalid upstream_token_header {raw!r}: must be a valid HTTP header name "
"(RFC 7230 token, e.g. 'esb-oauth')"
)
},
)
def validate_and_normalize_mcp_server_payload(payload: McpServerPayloadLike) -> None:
_base_validate_and_normalize_mcp_server_payload(payload)
_validate_mcp_server_name_fields(payload)
_validate_upstream_token_header(payload)
def stamp_omitted_oauth2_flow(payload: NewMCPServerRequest) -> None:
"""Fallback only: fill in oauth2_flow when an oauth2 create omits it.
@ -739,6 +757,7 @@ if MCP_AVAILABLE:
("aws_region_name", "aws_region_name"),
("aws_service_name", "aws_service_name"),
("upstream_resource", "upstream_resource"),
("upstream_token_header", "upstream_token_header"),
)
def _has_non_admin_config_credentials(credentials: "MCPCredentials | None") -> bool:

View file

@ -2661,6 +2661,7 @@ class LiteLLMCompletionResponsesConfig:
optional_output_details: Final[dict[str, int]] = {
field: value
for field, value in (
("audio_tokens", getattr(completion_details, "audio_tokens", None)),
("text_tokens", getattr(completion_details, "text_tokens", None)),
("image_tokens", getattr(completion_details, "image_tokens", None)),
)

View file

@ -1,10 +1,13 @@
"""Resolve which reasoning_effort values a deployment, and by intersection a model group, accepts.
The model map's supports_*_reasoning_effort flags are the only signal, and each level's polarity
mirrors how a request path reads that same flag. medium and high are unconditional for a reasoning
model. minimal and low are opt-out: openai/chat/gpt_5_transformation.py refuses them only when the
map says false. xhigh and max are opt-in. none is opt-out everywhere except the azure gpt-5 family,
whose config raises UnsupportedParamsError without an explicit true.
An entry that states its levels outright in reasoning_effort_levels is read first and wins
whole, for a model whose set the per-level flags cannot express: Kimi K3 takes low, high and max,
and no flag can drop medium because medium has none. Every other entry answers through the
supports_*_reasoning_effort flags below, whose polarity mirrors how a request path reads that same
flag. medium and high are unconditional for a reasoning model. minimal and low are opt-out:
openai/chat/gpt_5_transformation.py refuses them only when the map says false. xhigh and max are
opt-in. none is opt-out everywhere except the azure gpt-5 family, whose config raises
UnsupportedParamsError without an explicit true.
xhigh is gated on the request path by the openai and azure gpt-5 configs. max is not gated there at
all: every entry carrying supports_max_reasoning_effort is Claude-family, and
@ -41,6 +44,7 @@ _EFFORT_FLAGS: Final = (
("xhigh", "supports_xhigh_reasoning_effort"),
("max", "supports_max_reasoning_effort"),
)
_DECLARED_EFFORTS_KEY: Final = "reasoning_effort_levels"
_OPT_OUT_EFFORTS: Final = ("minimal", "low")
_OPT_IN_EFFORTS: Final = ("xhigh", "max")
_UNCONDITIONAL_EFFORTS: Final = frozenset(("medium", "high"))
@ -69,6 +73,20 @@ def _declared_effort_flags(model_info: Mapping[str, object]) -> Mapping[str, obj
)
def declared_reasoning_efforts(model_info: Mapping[str, object]) -> tuple[str, ...] | None:
"""The entry's own answer, read through the same bare twin as the flags so both spellings of one
model agree. Present-and-a-list IS the answer, so a declared [] correctly empties the group and
an unknown level is dropped rather than raised: the bundled map is enum-validated by
validate-model-prices-json, but an operator can put this key on a config.yaml model_info block
where that schema never runs, and one mistyped level must not fail every sibling on the proxy."""
own: Final = model_info.get(_DECLARED_EFFORTS_KEY)
raw: Final = own if own is not None else _bare_model_entry(model_info).get(_DECLARED_EFFORTS_KEY)
if not isinstance(raw, Sequence) or isinstance(raw, (str, bytes)):
return None
declared: Final = frozenset(effort for effort in raw if isinstance(effort, str))
return tuple(effort for effort in REASONING_EFFORT_ADVERTISEMENT_ORDER if effort in declared)
def _supports_none_reasoning_effort(model_info: Mapping[str, object], flag: object) -> bool:
"""Opt-in only where a request path refuses the level. AzureOpenAIGPT5Config raises
UnsupportedParamsError on reasoning_effort='none' without an explicit true, and it is selected
@ -119,6 +137,10 @@ def resolve_supported_reasoning_efforts(
if supports_reasoning is not True:
return () if supports_reasoning is False or deployment_is_mapped else None
declared: Final = declared_reasoning_efforts(model_info)
if declared is not None:
return declared
flags: Final = _declared_effort_flags(model_info)
if all(value is None for value in flags.values()):
return None

View file

@ -496,6 +496,9 @@ class BedrockGuardrailConfigModel(BaseModel):
aws_role_name: str | None = Field(default=None, description="AWS role name for assuming roles")
aws_web_identity_token: str | None = Field(default=None, description="Web identity token for AWS role assumption")
aws_sts_endpoint: str | None = Field(default=None, description="AWS STS endpoint URL")
aws_external_id: str | None = Field(
default=None, description="External ID required by the target role's trust policy on sts:AssumeRole"
)
aws_bedrock_runtime_endpoint: str | None = Field(default=None, description="AWS Bedrock runtime endpoint URL")
checks: BedrockChecksConfigModel | None = Field(
default=None,

View file

@ -324,7 +324,12 @@ class AnthropicMessagesToolResultParam(TypedDict, total=False):
is_error: bool
content: (
str
| Iterable[AnthropicMessagesToolResultContent | AnthropicMessagesImageParam | AnthropicMessagesDocumentParam]
| Iterable[
AnthropicMessagesToolResultContent
| AnthropicMessagesImageParam
| AnthropicMessagesDocumentParam
| ToolReference
]
)
cache_control: dict | ChatCompletionCachedContent | None

View file

@ -1,7 +1,7 @@
from collections.abc import Iterable, Mapping
from enum import Enum
from os import PathLike
from typing import IO, Any, Final, Literal, Optional, Union
from typing import IO, Any, Final, Literal, Optional, TypeAlias, Union
import httpx
from openai import Omit
@ -820,9 +820,21 @@ class ChatCompletionAssistantMessage(OpenAIChatCompletionAssistantMessage, total
reasoning_items: list[ChatCompletionReasoningItem] | None
class ChatCompletionToolReferenceObject(TypedDict):
"""Anthropic tool-search result block, carried through untouched so it survives a round trip."""
type: Literal["tool_reference"] # writable-ok: Pydantic warns on ReadOnly TypedDict fields
tool_name: str # writable-ok: Pydantic warns on ReadOnly TypedDict fields
ToolMessageContentPart: TypeAlias = (
ChatCompletionTextObject | ChatCompletionImageObject | ChatCompletionToolReferenceObject
)
class ChatCompletionToolMessage(TypedDict):
role: Literal["tool"]
content: str | Iterable[ChatCompletionTextObject | ChatCompletionImageObject]
content: str | Iterable[ToolMessageContentPart] # writable-ok: Pydantic warns on ReadOnly TypedDict fields
tool_call_id: str
@ -1258,6 +1270,8 @@ class ResponsesAPIRequestParams(ResponsesAPIOptionalRequestParams, total=False):
class OutputTokensDetails(BaseLiteLLMOpenAIResponseObject):
audio_tokens: int | None = None
reasoning_tokens: int | None = None
text_tokens: int | None = None

View file

@ -1,6 +1,11 @@
import enum
import re
from collections.abc import Awaitable, Callable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal
from urllib.parse import urlsplit
import httpx
from pydantic import BaseModel
from typing_extensions import TypedDict
@ -181,6 +186,15 @@ class MCPCredentials(TypedDict, total=False):
``audience``, which is the RFC 8693 token-exchange parameter.
"""
upstream_token_header: str | None # writable-ok: pydantic warns it cannot honour ReadOnly here
"""
Which upstream header carries the credential LiteLLM resolves for this server. Omitted when
unset, which keeps RFC 6750's default of ``Authorization``. Set it when the upstream expects the
gateway's token somewhere else (an ESB terminating its own credential on e.g. ``esb-oauth``), so
a separate operator-configured ``Authorization`` reaches the origin untouched. Non-secret, so it
is stored in plaintext and returned on admin reads.
"""
client_private_key: str | None
"""
PEM private key used to sign the private-key-JWT client_assertion (RFC 7523)
@ -223,7 +237,92 @@ class MCPCredentials(TypedDict, total=False):
"""
MCP_ADMIN_CONFIG_CREDENTIAL_KEYS: Final[tuple[str, ...]] = ("upstream_resource",)
DEFAULT_CREDENTIAL_HEADER: Final = "Authorization"
_HEADER_NAME_TOKEN: Final = re.compile(r"^[!#$%&'*+\-.^_`|~0-9A-Za-z]+$")
def normalize_upstream_header_name(raw: str) -> str | None:
"""The trimmed header name if it is a usable RFC 7230 ``token``, else None.
One owner for the grammar; each caller picks its own failure shape (a config-load raise, an
API 400, a typed CredError). An operator-supplied name reaches egress verbatim, so a value
carrying CR/LF, spaces or separators must never get that far.
"""
stripped: Final = raw.strip()
return stripped if stripped and _HEADER_NAME_TOKEN.match(stripped) else None
def same_header(name: str, other: str) -> bool:
"""Whether two HTTP header names are the same one. They are case-insensitive (RFC 7230 3.2)."""
return name.lower() == other.lower()
def has_header(headers: Mapping[str, str] | None, name: str) -> bool:
"""Whether ``headers`` carries ``name`` under any casing."""
return bool(headers) and any(same_header(key, name) for key in headers or {})
def without_header(headers: Mapping[str, str] | None, name: str) -> dict[str, str] | None:
"""A copy of ``headers`` with every casing of ``name`` removed, or None if nothing remains.
The one owner of "drop this credential's header". Both MCP stacks and the upstream-credential
resolver share it so a slot can never be dropped case-sensitively in one place and
case-insensitively in another, which is how an injected header came to shadow a resolved
credential on the v1 path.
"""
if not headers:
return None
filtered: Final = {key: value for key, value in headers.items() if not same_header(key, name)}
return filtered or None
_DEFAULT_PORTS: Final[Mapping[str, int]] = MappingProxyType({"http": 80, "https": 443})
def crosses_origin(configured: str, target: str) -> bool:
"""Whether ``target`` leaves ``configured``'s origin, by the rule HTTP clients use.
Origin is scheme, host and port, not host alone, so a same-host HTTPS downgrade or a port change
counts as crossing it. A plain http -> https upgrade of the same host is exempt, matching what
httpx exempts when it decides whether to keep ``Authorization`` across a redirect.
"""
a: Final = urlsplit(configured)
b: Final = urlsplit(target)
port_a: Final = a.port or _DEFAULT_PORTS.get(a.scheme)
port_b: Final = b.port or _DEFAULT_PORTS.get(b.scheme)
if a.scheme == b.scheme and a.hostname == b.hostname and port_a == port_b:
return False
return not (
a.hostname == b.hostname and a.scheme == "http" and port_a == 80 and b.scheme == "https" and port_b == 443
)
def custom_credential_slot(headers: Mapping[str, str] | None) -> str | None:
"""The first header carrying a credential somewhere other than ``Authorization``, if any."""
return next((name for name in headers or {} if not same_header(name, DEFAULT_CREDENTIAL_HEADER)), None)
def credential_redirect_hook(
configured_url: str, slot: str | None
) -> Callable[[httpx.Request], Awaitable[None]] | None:
"""An httpx request hook dropping ``slot`` once a redirect leaves ``configured_url``'s origin.
None when no guard is needed, so callers do not each repeat the exemption: HTTP clients already
strip ``Authorization`` across origins, but forward every other header, so only a credential an
operator moved to its own slot can be replayed to whatever host the upstream redirects to.
"""
if not configured_url or not slot or same_header(slot, DEFAULT_CREDENTIAL_HEADER):
return None
async def guard(request: httpx.Request) -> None:
if slot in request.headers and crosses_origin(configured_url, str(request.url)):
del request.headers[slot]
return guard
MCP_ADMIN_CONFIG_CREDENTIAL_KEYS: Final[tuple[str, ...]] = ("upstream_resource", "upstream_token_header")
"""Non-secret credential keys returned on read so the admin form can show and clear them. Mirrors
``ADMIN_CONFIG_CREDENTIAL_KEYS`` in ``ui/litellm-dashboard/src/components/mcp_tools/types.tsx``."""

View file

@ -1,7 +1,7 @@
from datetime import datetime
from typing import Any, Final, Literal
from pydantic import BaseModel, ConfigDict
from pydantic import BaseModel, ConfigDict, field_validator
from litellm.types.mcp import (
DEFAULT_SUBJECT_TOKEN_TYPE,
@ -9,6 +9,7 @@ from litellm.types.mcp import (
MCPAuthType,
MCPTokenEndpointAuthMethod,
MCPTransportType,
normalize_upstream_header_name,
)
# MCPInfo now allows arbitrary additional fields for custom metadata
@ -86,6 +87,22 @@ class MCPServer(BaseModel):
# today's behavior; "auto" derives the canonical URI from ``url``; any other value is sent
# verbatim. Resolved by ``oauth_utils.resolve_upstream_resource``.
upstream_resource: str | None = None
# Which upstream header carries the credential LiteLLM resolves for this server (the minted
# OAuth token, or the static key). None keeps RFC 6750's default, ``Authorization``. An ESB or
# API gateway that terminates its own credential in a private header needs this so a second,
# operator-configured ``Authorization`` can pass through to the origin untouched.
upstream_token_header: str | None = None
@field_validator("upstream_token_header")
@classmethod
def _check_upstream_token_header(cls, value: str | None) -> str | None:
if value is None or not value.strip():
return None
normalized: Final = normalize_upstream_header_name(value)
if normalized is None:
raise ValueError(f"upstream_token_header must be a valid HTTP header name (RFC 7230 token), got {value!r}")
return normalized
# AWS SigV4 fields
aws_access_key_id: str | None = None
aws_secret_access_key: str | None = None

View file

@ -164,6 +164,7 @@ class ProviderSpecificModelInfo(TypedDict, total=False):
supports_low_reasoning_effort: bool | None
supports_xhigh_reasoning_effort: bool | None
supports_max_reasoning_effort: bool | None
reasoning_effort_levels: ReadOnly[Sequence[str] | None]
supports_output_config: bool | None
supports_image_size: bool | None
bedrock_output_config_effort_ceiling: Literal["low", "medium", "high", "max", "xhigh"] | None

View file

@ -3025,8 +3025,7 @@ def register_model(
and value.get("cache_read_input_token_cost") is None
and value.get("tiered_pricing") is None
and (
value.get("input_cost_per_token") is not None
or value.get("output_cost_per_token") is not None
value.get("input_cost_per_token") is not None or value.get("output_cost_per_token") is not None
)
):
verbose_logger.warning(
@ -5890,6 +5889,7 @@ def _get_model_info_helper(
supports_low_reasoning_effort=_model_info.get("supports_low_reasoning_effort", None),
supports_xhigh_reasoning_effort=_model_info.get("supports_xhigh_reasoning_effort", None),
supports_max_reasoning_effort=_model_info.get("supports_max_reasoning_effort", None),
reasoning_effort_levels=_model_info.get("reasoning_effort_levels", None),
bedrock_output_config_effort_ceiling=_model_info.get("bedrock_output_config_effort_ceiling", None),
bedrock_converse_supports_strict_tools=_model_info.get("bedrock_converse_supports_strict_tools", None),
supports_computer_use=_model_info.get("supports_computer_use", None),

View file

@ -9315,6 +9315,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-kimi-k3-through-fireworks-ai-on-microsoft-foundry/4540187",
"supported_modalities": [
"text",
@ -14647,6 +14652,22 @@
"/v1/images/generations"
]
},
"dashscope/qwen-image-3.0": {
"litellm_provider": "dashscope",
"mode": "image_generation",
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supported_endpoints": [
"/v1/images/generations"
]
},
"dashscope/qwen-image-3.0-pro": {
"litellm_provider": "dashscope",
"mode": "image_generation",
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supported_endpoints": [
"/v1/images/generations"
]
},
"databricks/databricks-bge-large-en": {
"cache_creation_input_token_cost": 1.0003e-07,
"cache_read_input_token_cost": 1.0003e-07,
@ -31881,6 +31902,11 @@
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://platform.kimi.ai/docs/pricing/chat-k3",
"supports_function_calling": true,
"supports_reasoning": true,
@ -36489,6 +36515,14 @@
"litellm_provider": "perplexity",
"mode": "responses",
"output_cost_per_token": 1.5e-05,
"reasoning_effort_levels": [
"minimal",
"low",
"medium",
"high",
"xhigh",
"max"
],
"source": "https://docs.perplexity.ai/docs/agent-api/models",
"supports_web_search": true,
"supports_reasoning": true,
@ -38791,14 +38825,14 @@
"supports_reasoning": true
},
"together_ai/Qwen/Qwen3.7-Max": {
"cache_read_input_token_cost": 1.3e-07,
"input_cost_per_token": 1.25e-06,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 2.5e-06,
"litellm_provider": "together_ai",
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"max_tokens": 1000000,
"mode": "chat",
"output_cost_per_token": 3.75e-06,
"output_cost_per_token": 7.5e-06,
"source": "https://docs.together.ai/docs/serverless-models",
"supports_prompt_caching": true
},
@ -38970,6 +39004,11 @@
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.together.ai/docs/serverless-models",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
@ -39053,6 +39092,24 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"together_ai/zai-org/GLM-5.3-Flash": {
"cache_read_input_token_cost": 3e-08,
"input_cost_per_token": 1.5e-07,
"litellm_provider": "together_ai",
"max_input_tokens": 1048575,
"max_output_tokens": 1048575,
"max_tokens": 1048575,
"mode": "chat",
"output_cost_per_token": 5e-07,
"source": "https://docs.together.ai/docs/serverless-models",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"tts-1": {
"input_cost_per_character": 1.5e-05,
"litellm_provider": "openai",
@ -51740,6 +51797,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -51804,6 +51866,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -51820,6 +51887,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.25e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -51836,6 +51908,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -52008,6 +52085,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.25e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -52024,6 +52106,11 @@
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"reasoning_effort_levels": [
"low",
"high",
"max"
],
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,

View file

@ -532,6 +532,22 @@
"type": "object",
"description": "Provider-internal routing hints (e.g. bedrock_invocation_schema)."
},
"reasoning_effort_levels": {
"type": "array",
"description": "Exact reasoning_effort levels this deployment accepts; wins over supports_* flags.",
"items": {
"type": "string",
"enum": [
"none",
"minimal",
"low",
"medium",
"high",
"xhigh",
"max"
]
}
},
"regional_endpoint_uplift_multiplier": {
"type": "number",
"minimum": 1,

Binary file not shown.

After

Width:  |  Height:  |  Size: 30 KiB

View file

@ -16,7 +16,10 @@ via /model/new (Cohere, Gemini, hosted_vllm), each deleted on teardown.
from __future__ import annotations
import base64
import os
from pathlib import Path
from typing import Final
import pytest
from pydantic import BaseModel
@ -79,18 +82,22 @@ def _streamed_tool_call(events: list[str]) -> tuple[str, str]:
return name, arguments
CAT_IMAGE_URL = "https://upload.wikimedia.org/wikipedia/commons/3/3a/Cat03.jpg"
_FIXTURES_DIR: Final = Path(__file__).parent / "fixtures"
CAT_IMAGE: Final = _FIXTURES_DIR / "cat.jpg"
OPENAI_VISION_BACKEND = "openai/gpt-4o"
# OpenAI caches a shared prompt prefix once it exceeds ~1024 tokens; this is well
# past that, so a repeat call reports cached prompt tokens.
def _cat_image_data_url() -> str:
return "data:image/jpeg;base64," + base64.b64encode(CAT_IMAGE.read_bytes()).decode()
def _vision_messages() -> list[ChatMessage]:
return [
ChatMessage(
role="user",
content=[
TextContentPart(text="What animal is in this image? Answer in one word."),
ImageContentPart(image_url=ImageUrl(url=CAT_IMAGE_URL)),
ImageContentPart(image_url=ImageUrl(url=_cat_image_data_url())),
],
)
]

View file

@ -210,16 +210,24 @@ def _deltas(result: StreamingResponse) -> list[_StreamDelta]:
]
def _single_weather_call(message: OutMessage) -> ToolCall:
assert message.tool_calls, f"Together dropped the tool call: {message}"
assert len(message.tool_calls) == 1, f"expected one tool call, got {message.tool_calls}"
call = message.tool_calls[0]
def _validated_weather_call_id(call: ToolCall) -> str:
assert call.id, f"tool call carries no id, so a tool result cannot answer it: {call}"
assert call.function.name == "get_weather", f"wrong tool called: {call}"
assert call.function.arguments, f"tool call carries no arguments: {call}"
args = _WeatherArgs.model_validate_json(call.function.arguments)
assert "paris" in args.location.lower(), f"tool arguments lost the location: {args}"
return call
return call.id
def _weather_call_ids(message: OutMessage) -> tuple[str, ...]:
"""The id of every tool call the model made, each one checked for the fields a
caller needs to answer it. The backend is whichever together_ai row is cheapest
with tools and reasoning, and those rows carry supports_parallel_function_calling,
so one weather prompt can legitimately come back as several get_weather calls.
What the gateway owes us is that each call survives translation intact; how many
the model chose to make is the model's business."""
assert message.tool_calls, f"Together dropped the tool call: {message}"
return tuple(_validated_weather_call_id(call) for call in message.tool_calls)
def _weather_call(client: PassthroughClient, key: str, model: str) -> OutMessage:
@ -289,7 +297,7 @@ class TestTogetherChatCompletions:
self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str
) -> None:
model, key = _register(client, resources, reasoning_tool_backend)
_single_weather_call(_weather_call(client, key, model))
_ = _weather_call_ids(_weather_call(client, key, model))
@pytest.mark.covers("llm.chat_completions.together_ai.tool_use.stream.works")
def test_tool_call_is_streamed(
@ -328,8 +336,7 @@ class TestTogetherChatCompletions:
) -> None:
model, key = _register(client, resources, reasoning_tool_backend)
first = _weather_call(client, key, model)
call = _single_weather_call(first)
assert call.id is not None
call_ids = _weather_call_ids(first)
answer = _message(
unwrap(
@ -344,7 +351,10 @@ class TestTogetherChatCompletions:
reasoning_content=first.reasoning_content,
tool_calls=first.tool_calls,
),
ChatToolResultTurn(tool_call_id=call.id, content=WEATHER_REPORT),
*(
ChatToolResultTurn(tool_call_id=call_id, content=WEATHER_REPORT)
for call_id in call_ids
),
],
tools=[WEATHER_TOOL],
max_tokens=512,
@ -470,9 +480,21 @@ def _tool_use_blocks(content: list[AnthropicContentBlock] | None) -> list[Anthro
return [block for block in content if block.type == "tool_use"]
def _validated_tool_use_id(block: AnthropicContentBlock) -> str:
assert block.name == "get_weather", f"wrong tool called: {block}"
assert block.id, f"tool_use block carries no id, so a tool_result cannot answer it: {block}"
assert block.input is not None, f"tool_use block carries no input: {block}"
args = _WeatherArgs.model_validate(block.input)
assert "paris" in args.location.lower(), f"tool input lost the location: {args}"
return block.id
def _messages_weather_call(
client: PassthroughClient, key: str, model: str
) -> tuple[list[AnthropicContentBlock], AnthropicContentBlock]:
) -> tuple[list[AnthropicContentBlock], tuple[str, ...]]:
"""The blocks /v1/messages returned and the id of every tool_use among them. The
count is the model's choice (see _weather_call_ids); what this surface owes us is
that each tool_use arrives named and addressable."""
response = unwrap(
client.proxy.messages(
key,
@ -485,12 +507,9 @@ def _messages_weather_call(
)
)
tool_uses = _tool_use_blocks(response.content)
assert len(tool_uses) == 1, f"expected one tool_use block, got {response.content}"
block = tool_uses[0]
assert block.name == "get_weather", f"wrong tool called: {block}"
assert block.id, f"tool_use block carries no id, so a tool_result cannot answer it: {block}"
assert tool_uses, f"/v1/messages carried no tool_use block: {response.content}"
assert response.content is not None
return response.content, block
return response.content, tuple(_validated_tool_use_id(block) for block in tool_uses)
class TestTogetherMessages:
@ -506,8 +525,7 @@ class TestTogetherMessages:
self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str
) -> None:
model, key = _register(client, resources, reasoning_tool_backend)
first_content, block = _messages_weather_call(client, key, model)
assert block.id is not None
first_content, tool_use_ids = _messages_weather_call(client, key, model)
response = unwrap(
client.proxy.messages(
@ -520,7 +538,10 @@ class TestTogetherMessages:
ChatMessage(role="user", content=WEATHER_PROMPT),
AnthropicAssistantTurn(content=first_content),
AnthropicToolResultTurn(
content=[AnthropicToolResultBlock(tool_use_id=block.id, content=WEATHER_REPORT)]
content=[
AnthropicToolResultBlock(tool_use_id=tool_use_id, content=WEATHER_REPORT)
for tool_use_id in tool_use_ids
]
),
],
),

View file

@ -421,6 +421,7 @@ class AnthropicContentBlock(BaseModel):
text: str | None = None
id: str | None = None
name: str | None = None
input: dict[str, object] | None = None
class AnthropicToolResultBlock(BaseModel):

View file

@ -23,26 +23,18 @@ class TestTogetherAI(BaseLLMChatTest):
pass
@pytest.mark.parametrize(
"model, expected_bool",
"model",
[
("meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo", True),
("nvidia/Llama-3.1-Nemotron-70B-Instruct-HF", False),
"meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo",
"nvidia/Llama-3.1-Nemotron-70B-Instruct-HF",
],
)
def test_get_supported_response_format_together_ai(
self, model: str, expected_bool: bool
) -> None:
def test_get_supported_response_format_together_ai(self, model: str) -> None:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
optional_params = litellm.get_supported_openai_params(
model, custom_llm_provider="together_ai"
)
# Mapped provider
assert isinstance(optional_params, list)
if expected_bool:
assert "response_format" in optional_params
assert "tools" in optional_params
else:
assert "response_format" not in optional_params
assert "tools" not in optional_params
assert "response_format" in optional_params
assert "tools" in optional_params

View file

@ -3903,3 +3903,39 @@ def test_stored_reasoning_items_win_over_thinking_blocks():
reasoning_items = [item for item in input_items if item.get("type") == "reasoning"]
assert len(reasoning_items) == 1
assert reasoning_items[0]["id"] == "rs_real"
def test_convert_chat_completion_messages_to_responses_api_tool_result_with_tool_reference():
"""Tool-search tool_reference blocks have no Responses API equivalent: skip them, never stringify them."""
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
handler = LiteLLMResponsesTransformationHandler()
messages = [
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc123",
"type": "function",
"function": {"name": "ToolSearch", "arguments": '{"query": "web"}'},
}
],
},
{
"role": "tool",
"tool_call_id": "call_abc123",
"content": [
{"type": "tool_reference", "tool_name": "WebFetch"},
{"type": "text", "text": "1 tool found"},
],
},
]
response, _ = handler.convert_chat_completion_messages_to_responses_api(messages)
function_call_output = next(item for item in response if item.get("type") == "function_call_output")
assert function_call_output["output"] == [{"type": "input_text", "text": "1 tool found"}]

View file

@ -9,6 +9,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import anyio
import httpx
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
from mcp import McpError
from mcp.shared.message import SessionMessage
from mcp.types import (
@ -1095,3 +1096,188 @@ def test_mcp_extra_matches_proxy_extra_and_supports_streamable_http():
specifier = Requirement(mcp_extra[0]).specifier
assert not specifier.contains("1.23.0")
assert specifier.contains("1.28.1")
@pytest.mark.parametrize(
"auth_type, default_header",
[
(MCPAuth.oauth2, "Authorization"),
(MCPAuth.bearer_token, "Authorization"),
(MCPAuth.api_key, "X-API-Key"),
],
)
def test_v1_auth_headers_default_to_the_auth_type_slot(auth_type: MCPAuth, default_header: str) -> None:
client = MCPClient(server_url="http://up.example.com/mcp", auth_type=auth_type)
client.update_auth_value("tok")
assert default_header in client._get_auth_headers()
@pytest.mark.parametrize("auth_type", [MCPAuth.oauth2, MCPAuth.bearer_token, MCPAuth.api_key])
def test_v1_auth_headers_honor_the_configured_slot(auth_type: MCPAuth) -> None:
"""The v1 stack mints its own client_credentials token (oauth2_token_cache) and writes it here,
so leaving this table hardcoded makes the knob a silent no-op for every server that resolves
through v1 rather than the v2 resolver."""
client = MCPClient(
server_url="http://up.example.com/mcp",
auth_type=auth_type,
auth_header_name="esb-oauth",
)
client.update_auth_value("tok")
headers = client._get_auth_headers()
assert "esb-oauth" in headers
assert "Authorization" not in headers
assert "X-API-Key" not in headers
def test_v1_static_headers_still_win_their_own_slot():
# extra_headers (which carries static_headers) is applied last on the v1 path, so a static
# Authorization survives untouched while the resolved credential sits on its own header.
client = MCPClient(
server_url="http://up.example.com/mcp",
auth_type=MCPAuth.oauth2,
auth_header_name="esb-oauth",
extra_headers={"Authorization": "Bearer static-upstream-mcp-token"},
)
client.update_auth_value("minted")
headers = client._get_auth_headers()
assert headers["esb-oauth"] == "Bearer minted"
assert headers["Authorization"] == "Bearer static-upstream-mcp-token"
@pytest.mark.asyncio
async def test_a_custom_credential_header_is_stripped_when_a_redirect_crosses_origin():
"""httpx drops Authorization across origins but keeps every other header, so a credential the
operator moved to its own slot would be replayed to whatever host the upstream redirects to.
Verified against real httpx redirect handling, not a hand-built request.
"""
seen: "list[tuple[str, str]]" = []
def handler(request: httpx.Request) -> httpx.Response:
seen.append((request.url.host, request.headers.get("esb-oauth", "<stripped>")))
if request.url.host == "upstream.example.com":
return httpx.Response(302, headers={"Location": "https://attacker.example.com/collect"})
return httpx.Response(200)
client = MCPClient(
server_url="https://upstream.example.com/mcp",
auth_type=MCPAuth.oauth2,
auth_header_name="esb-oauth",
)
client.update_auth_value("minted-token")
factory = client._create_httpx_client_factory()
async with factory(headers=client._get_auth_headers(), timeout=None) as http_client:
http_client._transport = httpx.MockTransport(handler)
await http_client.get("https://upstream.example.com/mcp")
assert seen[0] == ("upstream.example.com", "Bearer minted-token")
assert seen[1] == ("attacker.example.com", "<stripped>")
@pytest.mark.asyncio
async def test_authorization_is_left_to_httpx_and_needs_no_guard():
# The default slot is already protected by httpx, so the client must not install a guard for it
# and must not interfere with the ordinary Authorization path.
url = "https://upstream.example.com/mcp"
from litellm.types.mcp import credential_redirect_hook
def guard_for(client: MCPClient):
return credential_redirect_hook(client.server_url, client._credential_slot)
assert guard_for(MCPClient(server_url=url, auth_type=MCPAuth.oauth2)) is None
assert guard_for(MCPClient(server_url=url, resolved_auth=StaticHeaderAuth("Bearer x"))) is None
# a v2 resolver slot is discovered from the auth object, without the caller naming it again
custom = MCPClient(server_url=url, resolved_auth=StaticHeaderAuth("Bearer x", header_name="esb-oauth"))
assert guard_for(custom) is not None
# and the same answer arrives via the v1 configured slot
assert guard_for(MCPClient(server_url=url, auth_header_name="ESB-OAuth")) is not None
def test_an_injected_header_cannot_shadow_the_configured_credential_slot():
"""The v2 path drops a colliding injected header so the resolved credential wins its slot. The
v1 path applies extra_headers last, so without this it silently sends the injected value and the
upstream rejects a credential the gateway thought it had sent.
"""
client = MCPClient(
server_url="https://upstream.example.com/mcp",
auth_type=MCPAuth.oauth2,
auth_header_name="esb-oauth",
extra_headers={"esb-oauth": "Bearer injected", "X-Trace": "keep"},
)
client.update_auth_value("minted-token")
headers = client._get_auth_headers()
assert headers["esb-oauth"] == "Bearer minted-token"
assert headers["X-Trace"] == "keep"
def test_without_a_configured_slot_the_existing_precedence_is_unchanged():
# extra_headers winning over authentication_token is long-standing v1 behavior; the fix above
# must apply only to the slot the operator explicitly named.
client = MCPClient(
server_url="https://upstream.example.com/mcp",
auth_type=MCPAuth.oauth2,
extra_headers={"Authorization": "Bearer injected"},
)
client.update_auth_value("minted-token")
assert client._get_auth_headers()["Authorization"] == "Bearer injected"
_REDIRECT_CASES = [
("https://upstream.example.com/mcp", "https://upstream.example.com/other"), # same origin
("https://upstream.example.com/mcp", "https://upstream.example.com:443/other"), # explicit default port
("https://upstream.example.com/mcp", "https://attacker.example.com/collect"), # different host
("https://upstream.example.com/mcp", "http://upstream.example.com/collect"), # scheme downgrade
("https://upstream.example.com/mcp", "https://upstream.example.com:8443/other"), # different port
("https://upstream.example.com/mcp", "https://sub.upstream.example.com/x"), # different host
("http://upstream.example.com/mcp", "https://upstream.example.com/other"), # http -> https upgrade
("http://upstream.example.com/mcp", "http://upstream.example.com/other"), # same origin, plain http
]
@pytest.mark.parametrize("start,target", _REDIRECT_CASES)
@pytest.mark.asyncio
async def test_the_guard_agrees_with_httpx_about_authorization(start: str, target: str) -> None:
"""Our custom slot must be dropped on exactly the redirects where httpx drops Authorization.
The rule is mirrored rather than imported, so this drives real httpx and compares the two
outcomes. A future httpx that changes its redirect rule reds here instead of silently leaving
the custom slot forwarded where Authorization is not (or stripped where it is not needed).
"""
seen: "list[tuple[str, str, str]]" = []
def handler(request: httpx.Request) -> httpx.Response:
seen.append(
(
str(request.url),
request.headers.get("authorization", "<stripped>"),
request.headers.get("esb-oauth", "<stripped>"),
)
)
if str(request.url) == start:
return httpx.Response(302, headers={"Location": target})
return httpx.Response(200)
client = MCPClient(server_url=start, auth_type=MCPAuth.oauth2, auth_header_name="esb-oauth")
factory = client._create_httpx_client_factory()
async with factory(headers={"Authorization": "Bearer AUTH", "esb-oauth": "Bearer ESB"}, timeout=None) as http:
http._transport = httpx.MockTransport(handler)
await http.get(start)
_url, authorization, esb = seen[-1]
assert (authorization == "<stripped>") == (esb == "<stripped>"), (
f"httpx and the guard disagree for {target}: authorization={authorization!r} esb-oauth={esb!r}"
)
def test_a_differently_cased_injected_header_cannot_shadow_the_slot() -> None:
# HTTP header names are case-insensitive and v2 drops the collision case-insensitively, so an
# exact-key check here would leave both spellings in the dict and let the injected value win.
client = MCPClient(
server_url="https://upstream.example.com/mcp",
auth_type=MCPAuth.oauth2,
auth_header_name="esb-oauth",
extra_headers={"ESB-OAuth": "Bearer injected", "X-Trace": "keep"},
)
client.update_auth_value("minted-token")
headers = client._get_auth_headers()
assert [v for k, v in headers.items() if k.lower() == "esb-oauth"] == ["Bearer minted-token"]
assert headers["X-Trace"] == "keep"

View file

@ -0,0 +1,134 @@
from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
SubtitleToken,
render_subtitle_tokens_as_srt,
render_subtitle_tokens_as_vtt,
synthesize_subtitle_document,
)
class TestRenderSubtitleTokensAsSrt:
def test_single_cue_full_document(self):
tokens = (
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
SubtitleToken(text="world.", start_ms=500, end_ms=1000),
)
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:01,000\nHello world.\n"
def test_speaker_change_starts_a_new_cue(self):
tokens = (
SubtitleToken(text="Hi.", start_ms=0, end_ms=1000, speaker="spk:0"),
SubtitleToken(text="Hey.", start_ms=1500, end_ms=2500, speaker="spk:1"),
)
assert render_subtitle_tokens_as_srt(tokens) == (
"1\n00:00:00,000 --> 00:00:01,000\nHi.\n\n2\n00:00:01,500 --> 00:00:02,500\nHey.\n"
)
def test_token_cap_starts_a_new_cue_after_15_tokens(self):
tokens = tuple(
SubtitleToken(text=f"{index} ", start_ms=index * 100, end_ms=index * 100 + 100) for index in range(16)
)
assert render_subtitle_tokens_as_srt(tokens) == (
"1\n00:00:00,000 --> 00:00:01,500\n0 1 2 3 4 5 6 7 8 9 10 11 12 13 14\n"
"\n2\n00:00:01,500 --> 00:00:01,600\n15\n"
)
def test_duration_cap_starts_a_new_cue_at_5000ms(self):
tokens = (
SubtitleToken(text="Alpha ", start_ms=0, end_ms=400),
SubtitleToken(text="beta ", start_ms=2000, end_ms=2400),
SubtitleToken(text="gamma.", start_ms=5000, end_ms=5400),
)
assert render_subtitle_tokens_as_srt(tokens) == (
"1\n00:00:00,000 --> 00:00:02,400\nAlpha beta\n\n2\n00:00:05,000 --> 00:00:05,400\ngamma.\n"
)
def test_timestampless_token_joins_the_current_cue(self):
tokens = (
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
SubtitleToken(text="there "),
SubtitleToken(text="world.", start_ms=900, end_ms=1300),
)
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:01,300\nHello there world.\n"
def test_only_timestampless_tokens_renders_empty(self):
assert render_subtitle_tokens_as_srt((SubtitleToken(text="no timestamps"),)) == ""
def test_empty_tokens_render_empty(self):
assert render_subtitle_tokens_as_srt(()) == ""
def test_timestamps_past_one_hour(self):
tokens = (SubtitleToken(text="Late.", start_ms=3_661_001, end_ms=3_662_002),)
assert render_subtitle_tokens_as_srt(tokens) == "1\n01:01:01,001 --> 01:01:02,002\nLate.\n"
def test_negative_timestamps_clamp_to_zero(self):
tokens = (SubtitleToken(text="Early.", start_ms=-100, end_ms=-50),)
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:00,000\nEarly.\n"
def test_missing_end_falls_back_to_cue_start(self):
tokens = (SubtitleToken(text="Open.", start_ms=1200),)
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:01,200 --> 00:00:01,200\nOpen.\n"
class TestRenderSubtitleTokensAsVtt:
def test_single_cue_full_document(self):
tokens = (
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
SubtitleToken(text="world.", start_ms=500, end_ms=1000),
)
assert render_subtitle_tokens_as_vtt(tokens) == "WEBVTT\n\n00:00:00.000 --> 00:00:01.000\nHello world.\n"
def test_empty_tokens_render_header_only(self):
assert render_subtitle_tokens_as_vtt(()) == "WEBVTT\n"
def test_timestamps_past_one_hour_use_dot_separator(self):
tokens = (SubtitleToken(text="Late.", start_ms=3_661_001, end_ms=3_662_002),)
assert render_subtitle_tokens_as_vtt(tokens) == "WEBVTT\n\n01:01:01.001 --> 01:01:02.002\nLate.\n"
def test_speaker_change_starts_a_new_cue(self):
tokens = (
SubtitleToken(text="Hi.", start_ms=0, end_ms=1000, speaker=1),
SubtitleToken(text="Hey.", start_ms=1500, end_ms=2500, speaker=2),
)
assert render_subtitle_tokens_as_vtt(tokens) == (
"WEBVTT\n\n00:00:00.000 --> 00:00:01.000\nHi.\n\n00:00:01.500 --> 00:00:02.500\nHey.\n"
)
class TestSynthesizeSubtitleDocument:
WORDS = [
{"word": "Four", "start": 0.4, "end": 0.7, "speaker": "spk:0"},
{"word": "score", "start": 0.7, "end": 1.1, "speaker": "spk:0"},
]
def test_srt_from_words_converts_seconds_to_milliseconds(self):
assert synthesize_subtitle_document(self.WORDS, "srt") == "1\n00:00:00,400 --> 00:00:01,100\nFour score\n"
def test_vtt_from_words_converts_seconds_to_milliseconds(self):
assert synthesize_subtitle_document(self.WORDS, "vtt") == (
"WEBVTT\n\n00:00:00.400 --> 00:00:01.100\nFour score\n"
)
def test_speaker_change_splits_cues(self):
words = [
{"word": "Hi", "start": 0.0, "end": 0.5, "speaker": "spk:0"},
{"word": "Hey", "start": 0.6, "end": 1.0, "speaker": "spk:1"},
]
assert synthesize_subtitle_document(words, "srt") == (
"1\n00:00:00,000 --> 00:00:00,500\nHi\n\n2\n00:00:00,600 --> 00:00:01,000\nHey\n"
)
def test_non_subtitle_format_returns_none(self):
assert synthesize_subtitle_document(self.WORDS, "verbose_json") is None
assert synthesize_subtitle_document(self.WORDS, "json") is None
def test_missing_words_returns_none(self):
assert synthesize_subtitle_document(None, "srt") is None
assert synthesize_subtitle_document([], "srt") is None
def test_words_without_timestamps_return_none(self):
assert synthesize_subtitle_document([{"word": "Hello"}], "srt") is None
assert synthesize_subtitle_document([{"word": "Hello"}], "vtt") is None
def test_malformed_words_return_none(self):
assert synthesize_subtitle_document("not words", "srt") is None
assert synthesize_subtitle_document([{"word": "ok", "start": "not-a-number"}], "srt") is None

View file

@ -1027,3 +1027,70 @@ def test_update_messages_xlitellm_decode_does_not_override_mapping():
updated = update_messages_with_model_file_ids(messages, "model-A", mapping)
assert updated[0]["content"][0]["file"]["file_id"] == "provider-explicit-id"
def test_drop_tool_reference_parts_keeps_text_parts():
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_tool_reference_parts_from_tool_messages,
)
messages = [
_assistant_tool_call_msg("call_1"),
_tool_msg(
[
{"type": "text", "text": "WebFetch tool loaded successfully."},
{"type": "tool_reference", "tool_name": "WebFetch"},
]
),
]
result = drop_tool_reference_parts_from_tool_messages(messages)
assert result[1]["content"] == [{"type": "text", "text": "WebFetch tool loaded successfully."}]
assert result[1]["tool_call_id"] == "call_1"
def test_drop_tool_reference_parts_reference_only_becomes_empty_text():
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_tool_reference_parts_from_tool_messages,
)
messages = [
_assistant_tool_call_msg("call_1"),
_tool_msg([{"type": "tool_reference", "tool_name": "WebFetch"}]),
]
result = drop_tool_reference_parts_from_tool_messages(messages)
assert result[1] == {"role": "tool", "tool_call_id": "call_1", "content": ""}
def test_drop_tool_reference_parts_without_references_passes_through():
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_tool_reference_parts_from_tool_messages,
)
messages = [
_assistant_tool_call_msg("call_1"),
_tool_msg([{"type": "text", "text": "plain result"}]),
]
assert drop_tool_reference_parts_from_tool_messages(messages) is messages
def test_drop_tool_reference_parts_leaves_non_tool_messages_alone():
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_tool_reference_parts_from_tool_messages,
)
user_message = {"role": "user", "content": [{"type": "tool_reference", "tool_name": "WebFetch"}]}
messages = [
user_message,
_assistant_tool_call_msg("call_1"),
_tool_msg([{"type": "tool_reference", "tool_name": "WebFetch"}]),
]
result = drop_tool_reference_parts_from_tool_messages(messages)
assert result[0] == user_message
assert result[2]["content"] == ""

View file

@ -3578,3 +3578,52 @@ async def test_bedrock_converse_pdf_only_user_message_gets_text_block_async():
assert len(result) == 1
assert any("document" in block for block in result[0]["content"])
assert _text_blocks(result[0]) == [BEDROCK_DOCUMENT_PLACEHOLDER_TEXT]
def test_convert_to_anthropic_tool_result_keeps_tool_reference_blocks():
from litellm.litellm_core_utils.prompt_templates.factory import convert_to_anthropic_tool_result
result = convert_to_anthropic_tool_result(
{
"role": "tool",
"tool_call_id": "toolu_01",
"content": [
{"type": "text", "text": "loaded"},
{"type": "tool_reference", "tool_name": "WebFetch"},
],
}
)
assert result == {
"type": "tool_result",
"tool_use_id": "toolu_01",
"content": [
{"type": "text", "text": "loaded"},
{"type": "tool_reference", "tool_name": "WebFetch"},
],
}
def test_convert_gemini_tool_call_result_answers_tool_reference_only_result():
"""Every Gemini function call needs a function response, even when the tool result carries no text.
Fixes: https://github.com/BerriAI/litellm/issues/37462
"""
result = convert_to_gemini_tool_call_result(
message=ChatCompletionToolMessage(
role="tool",
tool_call_id="toolu_01",
content=[{"type": "tool_reference", "tool_name": "WebFetch"}],
),
last_message_with_tool_calls={
"role": "assistant",
"tool_calls": [
{
"id": "toolu_01",
"type": "function",
"function": {"name": "ToolSearch", "arguments": '{"query": "select:WebFetch"}'},
}
],
},
)
assert result == {"function_response": {"name": "ToolSearch", "response": {"content": ""}}}

View file

@ -4460,3 +4460,51 @@ def test_handle_stream_fallback_error_restores_context_only_after_exception_mapp
finally:
trace_id_var.set("")
session_id_var.set("")
def test_chunk_creator_preserves_hidden_provider_specific_fields_from_parsed_chunk():
wrapper = CustomStreamWrapper(
completion_stream=None,
model="gemini-3.5-flash",
logging_obj=MagicMock(),
custom_llm_provider="vertex_ai",
)
parsed_chunk = ModelResponseStream(
choices=[StreamingChoices(index=0, delta=Delta(content="hello", role="assistant"), finish_reason=None)],
)
parsed_chunk._hidden_params["provider_specific_fields"] = {"traffic_type": "ON_DEMAND_FLEX"}
result = wrapper.chunk_creator(chunk=parsed_chunk)
assert result is not None
assert result._hidden_params["provider_specific_fields"] == {"traffic_type": "ON_DEMAND_FLEX"}
@pytest.mark.asyncio
async def test_async_stream_assembled_response_keeps_vertex_traffic_type(logging_obj: Logging):
content_chunk = ModelResponseStream(
choices=[StreamingChoices(index=0, delta=Delta(content="hello", role="assistant"), finish_reason=None)],
)
final_chunk = ModelResponseStream(
choices=[StreamingChoices(index=0, delta=Delta(content=""), finish_reason="stop")],
)
setattr(final_chunk, "usage", Usage(prompt_tokens=7, completion_tokens=5, total_tokens=12))
final_chunk._hidden_params["provider_specific_fields"] = {"traffic_type": "ON_DEMAND_FLEX"}
async def _stream():
yield content_chunk
yield final_chunk
wrapper = CustomStreamWrapper(
completion_stream=_stream(),
model="gemini-3.5-flash",
logging_obj=logging_obj,
custom_llm_provider="vertex_ai",
stream_options={"include_usage": True},
)
received = [chunk async for chunk in wrapper]
assembled = litellm.stream_chunk_builder(chunks=received, messages=[{"role": "user", "content": "hi"}])
assert assembled is not None
assert assembled._hidden_params["provider_specific_fields"]["traffic_type"] == "ON_DEMAND_FLEX"

View file

@ -290,6 +290,24 @@ class TestAnthropicMessagesHandlerInputProcessing:
assert data.get("litellm_metadata", {}).get("guardrails")
assert guardrail.dynamic_params == {"policy_id": "policy-123"}
@pytest.mark.asyncio
async def test_provider_native_tools_survive_guardrail_round_trip(self):
handler = AnthropicMessagesHandler()
guardrail = MockPassThroughGuardrail(guardrail_name="test")
data = {
"model": "gemini-2.5-flash",
"messages": [{"role": "user", "content": "coffee shops near Union Square?"}],
"tools": [
{"googleMaps": {"enable_widget": True}},
{"name": "get_weather", "input_schema": {"type": "object", "properties": {}}},
],
}
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert {"googleMaps": {"enable_widget": True}} in data["tools"]
assert [tool["name"] for tool in data["tools"] if "name" in tool] == ["get_weather"]
@pytest.mark.asyncio
async def test_midturn_system_correction_is_guardrailed_when_top_level_system_is_skipped(
self,
@ -1818,3 +1836,72 @@ class TestAnthropicMessagesScanOnlyToolResults:
assert guardrail.captured_inputs is not None
assert guardrail.captured_inputs.get("images") == ["TOOL_IMG"]
class TestStructuredWriteBackKeepsToolResults:
"""A guardrail rewrite must never leave a tool_use without its tool_result (Claude Code ToolSearch, LIT-6103)."""
@staticmethod
def _claude_code_tool_search_turns(tool_result_content):
return [
{"role": "user", "content": "load WebFetch for bob@example.com"},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_01",
"name": "ToolSearch",
"input": {"query": "select:WebFetch"},
}
],
},
{
"role": "user",
"content": [
{"type": "tool_result", "tool_use_id": "toolu_01", "content": tool_result_content},
{"type": "text", "text": "Now fetch the page."},
],
},
]
@staticmethod
def _blocks(message):
return message["content"] if isinstance(message["content"], list) else []
@pytest.mark.parametrize(
("tool_result_content", "expected_written_back_content"),
[
(
[{"type": "tool_reference", "tool_name": "WebFetch"}],
[{"type": "tool_reference", "tool_name": "WebFetch"}],
),
([], ""),
],
ids=["tool_reference", "empty"],
)
async def test_tool_result_stays_right_after_its_tool_use(
self, tool_result_content, expected_written_back_content
):
handler = AnthropicMessagesHandler()
data = {"model": "claude-fable-5", "messages": self._claude_code_tool_search_turns(tool_result_content)}
await handler.process_input_messages(data=data, guardrail_to_apply=MockStructuredMaskingGuardrail())
serialized = json.dumps(data["messages"])
assert "bob@example.com" not in serialized
assert "<EMAIL>" in serialized
messages = data["messages"]
tool_use_index = next(
i for i, m in enumerate(messages) if any(b.get("type") == "tool_use" for b in self._blocks(m))
)
answer = messages[tool_use_index + 1]
assert answer["role"] == "user"
assert answer["content"][0] == {
"type": "tool_result",
"tool_use_id": "toolu_01",
"content": expected_written_back_content,
}
later_blocks = [b for m in messages[tool_use_index + 1 :] for b in self._blocks(m)]
assert {"type": "text", "text": "Now fetch the page."} in later_blocks

View file

@ -16,6 +16,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
)
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
OPENAI_MAX_TOOL_NAME_LENGTH,
AnthropicAdapter,
LiteLLMAnthropicMessagesAdapter,
create_tool_name_mapping,
truncate_tool_name,
@ -1986,8 +1987,13 @@ def test_adaptive_thinking_output_config_effort_preserved_for_claude_model(model
backend. On Bedrock Converse, adaptive thinking without effort streams zero reasoning
blocks. The `format` subkey must still be excluded (it is translated to
`response_format` separately).
Bedrock keeps taking the tier as `output_config`, which attaches it without disturbing
`thinking`. Driving the translated request through the provider's own param mapping is what
makes the second half a claim about the wire rather than about an intermediate key.
"""
from litellm.types.llms.anthropic import AnthropicMessagesRequest
from litellm.utils import get_optional_params
anthropic_request = AnthropicMessagesRequest(
model=model,
@ -2005,8 +2011,18 @@ def test_adaptive_thinking_output_config_effort_preserved_for_claude_model(model
assert openai_request["thinking"] == {"type": "adaptive"}
assert openai_request["output_config"] == {"effort": "max"}
assert "reasoning_effort" not in openai_request
assert "response_format" in openai_request
on_the_wire = get_optional_params(
model=model,
custom_llm_provider="bedrock",
thinking=openai_request["thinking"],
output_config=openai_request["output_config"],
)
assert on_the_wire["output_config"] == {"effort": "max"}
def test_adaptive_thinking_format_only_output_config_not_forwarded_for_claude_model():
"""When `output_config` carries only `format`, nothing effort-bearing remains, so the
@ -2029,9 +2045,12 @@ def test_adaptive_thinking_format_only_output_config_not_forwarded_for_claude_mo
def test_adaptive_thinking_output_config_not_forwarded_for_non_bedrock_claude_model():
"""`output_config` is forwarded only for Bedrock-destined Claude models. Other
Claude-through-bridge providers (e.g. openrouter) accept `thinking` but reject a raw
`output_config` param with UnsupportedParamsError when drop_params is off."""
"""`output_config` is never forwarded raw to a bridged provider: openrouter and friends accept
`thinking` but reject that param with UnsupportedParamsError when drop_params is off.
Regression: the tier used to be dropped along with it, so an openrouter Claude deployment got a
bare adaptive `thinking` block and the caller's effort did nothing, byte-identical for `max` and
`minimal`. It now travels as `reasoning_effort`, which that provider does accept."""
from litellm.types.llms.anthropic import AnthropicMessagesRequest
anthropic_request = AnthropicMessagesRequest(
@ -2047,6 +2066,68 @@ def test_adaptive_thinking_output_config_not_forwarded_for_non_bedrock_claude_mo
assert openai_request["thinking"] == {"type": "adaptive"}
assert "output_config" not in openai_request
assert openai_request["reasoning_effort"] == "max"
@pytest.mark.parametrize("effort", ["minimal", "low", "medium", "high", "xhigh", "max"])
def test_every_adaptive_effort_tier_reaches_a_bridged_claude_target(effort):
"""The tier the caller asked for is the tier the bridge carries, for every level. The bug was
invisible per-request because each call returned 200; only comparing two tiers showed the
upstream body was the same either way."""
from litellm.types.llms.anthropic import AnthropicMessagesRequest
adapter = LiteLLMAnthropicMessagesAdapter()
openai_request, _ = adapter.translate_anthropic_to_openai(
anthropic_message_request=AnthropicMessagesRequest(
model="openrouter/anthropic/claude-opus-4-7",
max_tokens=1024,
messages=[{"role": "user", "content": "hi"}],
thinking={"type": "adaptive"},
output_config={"effort": effort},
)
)
assert openai_request["reasoning_effort"] == effort
def test_adaptive_thinking_without_a_tier_leaves_a_claude_target_on_its_own_default():
"""Adaptive with no `output_config.effort` must stay bare, so the provider's own adaptive
default still decides. Inventing a tier here would silently override it."""
from litellm.types.llms.anthropic import AnthropicMessagesRequest
adapter = LiteLLMAnthropicMessagesAdapter()
openai_request, _ = adapter.translate_anthropic_to_openai(
anthropic_message_request=AnthropicMessagesRequest(
model="openrouter/anthropic/claude-opus-4-7",
max_tokens=1024,
messages=[{"role": "user", "content": "hi"}],
thinking={"type": "adaptive"},
)
)
assert openai_request["thinking"] == {"type": "adaptive"}
assert "reasoning_effort" not in openai_request
assert "output_config" not in openai_request
def test_budgeted_thinking_on_a_claude_target_keeps_its_budget_and_gains_no_tier():
"""`enabled` + `budget_tokens` is more precise than any tier, so the bridge must forward it
untouched rather than coarsening it into a `reasoning_effort` bucket."""
from litellm.types.llms.anthropic import AnthropicMessagesRequest
adapter = LiteLLMAnthropicMessagesAdapter()
openai_request, _ = adapter.translate_anthropic_to_openai(
anthropic_message_request=AnthropicMessagesRequest(
model="openrouter/anthropic/claude-opus-4-7",
max_tokens=1024,
messages=[{"role": "user", "content": "hi"}],
thinking={"type": "enabled", "budget_tokens": 8000},
output_config={"effort": "max"},
)
)
assert openai_request["thinking"] == {"type": "enabled", "budget_tokens": 8000}
assert "reasoning_effort" not in openai_request
def test_stop_sequences_translated_to_stop_for_non_claude_model():
@ -2307,6 +2388,53 @@ def test_translate_anthropic_tools_to_openai_fills_missing_tool_name():
assert result[1]["function"]["name"] == "litellm_unnamed_tool_1"
def test_translate_anthropic_tools_to_openai_passes_provider_native_tool_dicts_through():
"""Deployment-level provider-native tools (e.g. Gemini googleMaps) must reach the provider transformation verbatim (LIT-6286)."""
tools = [
{"googleMaps": {}},
{"googleSearch": {}},
{
"name": "get_weather",
"input_schema": {"type": "object", "properties": {"location": {"type": "string"}}},
},
]
adapter = LiteLLMAnthropicMessagesAdapter()
result, tool_name_mapping = adapter.translate_anthropic_tools_to_openai(tools=tools, model=None)
assert result[0] == {"googleMaps": {}}
assert result[1] == {"googleSearch": {}}
assert result[2]["function"]["name"] == "get_weather"
assert tool_name_mapping == {}
def test_translate_anthropic_tools_to_openai_passes_openai_function_tools_through():
"""A tool already in OpenAI function format must pass through unchanged instead of becoming litellm_unnamed_tool_N."""
openai_tool = {
"type": "function",
"function": {
"name": "get_weather",
"parameters": {"type": "object", "properties": {"location": {"type": "string"}}},
},
}
adapter = LiteLLMAnthropicMessagesAdapter()
result, _ = adapter.translate_anthropic_tools_to_openai(tools=[openai_tool], model=None)
assert result == [openai_tool]
def test_translate_completion_input_params_keeps_provider_native_tools():
"""/v1/messages request translation must keep router-merged provider-native tools in kwargs['tools'] (LIT-6286)."""
adapter = AnthropicAdapter()
translated = adapter.translate_completion_input_params(
{
"model": "gemini/gemini-2.5-flash",
"max_tokens": 1024,
"messages": [{"role": "user", "content": "coffee shops near Union Square"}],
"tools": [{"googleMaps": {}}],
}
)
assert translated is not None
assert translated["tools"] == [{"googleMaps": {}}]
def test_translate_openai_content_to_anthropic_reasoning_content_without_thinking_blocks():
"""
Test that reasoning_content is converted to thinking block when thinking_blocks is not present.
@ -3999,6 +4127,75 @@ def test_translate_anthropic_messages_to_openai_carries_midturn_system_prompt_ca
]
def _tool_reference_block(tool_name="WebFetch"):
return {"type": "tool_reference", "tool_name": tool_name}
def test_tool_result_tool_reference_is_carried_through_untouched():
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_messages_to_openai(
messages=[
_anthropic_tool_use_turn("toolu_01"),
_anthropic_tool_result_turn({"toolu_01": [_tool_reference_block()]}),
]
)
assert [m["role"] for m in result] == ["assistant", "tool"]
assert result[1]["tool_call_id"] == "toolu_01"
assert result[1]["content"] == [{"type": "tool_reference", "tool_name": "WebFetch"}]
def test_tool_result_text_beside_tool_reference_keeps_both_parts_in_order():
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_messages_to_openai(
messages=[
_anthropic_tool_use_turn("toolu_01"),
_anthropic_tool_result_turn(
{"toolu_01": [{"type": "text", "text": "loaded"}, _tool_reference_block("Grep")]}
),
]
)
assert result[1]["content"] == [
{"type": "text", "text": "loaded"},
{"type": "tool_reference", "tool_name": "Grep"},
]
@pytest.mark.parametrize(
"tool_result_content",
[
[],
None,
"",
{"not": "a list"},
[{"type": "future_block", "payload": 1}],
[{"type": "search_result", "source": "https://example.com", "title": "t", "content": []}],
],
ids=["empty_list", "null", "empty_string", "non_list", "unknown_block", "search_result_only"],
)
def test_tool_result_without_translatable_content_still_answers_its_tool_use(tool_result_content):
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_messages_to_openai(
messages=[
_anthropic_tool_use_turn("toolu_01"),
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": tool_result_content}],
},
]
)
assert result == [
result[0],
{"role": "tool", "tool_call_id": "toolu_01", "content": ""},
]
assert result[0]["role"] == "assistant"
def _openai_response_with_usage(usage: Usage) -> ModelResponse:
return ModelResponse(
id="resp_web_search",
@ -4092,3 +4289,161 @@ def test_completion_cost_on_translated_anthropic_response_includes_web_search():
]
assert per_query_cost > 0
assert cost_with_search - cost_without_search == pytest.approx(2 * per_query_cost)
@pytest.mark.parametrize(
"model, provider, carried",
[
("databricks/databricks-claude-opus-4-7", "databricks", "max"),
("openrouter/anthropic/claude-opus-4-7", "openrouter", "xhigh"),
],
)
def test_a_summary_bearing_adaptive_request_still_delivers_its_tier(model, provider, carried):
"""The summary rides inside the forwarded `thinking` block for a Claude target, so the tier must
stay a plain string. Wrapping it into `{"effort": ..., "summary": ...}` made databricks raise
`Invalid reasoning_effort` and made bedrock drop `output_config` altogether, losing the tier on
exactly the path this translator exists to serve.
Each case names the exact tier that provider ends up sending, not merely that something arrived:
bedrock and databricks rebuild `output_config`, and openrouter applies its own max to xhigh
remap, so asserting presence alone would pass on a mapping that silently changed the tier."""
from litellm.types.llms.anthropic import AnthropicMessagesRequest
from litellm.utils import get_optional_params
adapter = LiteLLMAnthropicMessagesAdapter()
openai_request, _ = adapter.translate_anthropic_to_openai(
anthropic_message_request=AnthropicMessagesRequest(
model=model,
max_tokens=1024,
messages=[{"role": "user", "content": "hi"}],
thinking={"type": "adaptive", "summary": "detailed"},
output_config={"effort": "max"},
)
)
assert openai_request["reasoning_effort"] == "max"
on_the_wire = get_optional_params(
model=model,
custom_llm_provider=provider,
thinking=openai_request["thinking"],
reasoning_effort=openai_request["reasoning_effort"],
)
on_the_wire_tier = on_the_wire.get("output_config", {}).get("effort") or on_the_wire.get("reasoning_effort")
assert on_the_wire_tier == carried
ARN_MODEL = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123"
def test_an_inference_profile_arn_keeps_taking_its_tier_as_output_config():
"""Regression: an ARN contains neither `anthropic` nor `claude`, so it reaches this branch only
through `is_bedrock_arn_model`. Bedrock resolves no chat config for one, so `reasoning_effort`
is dropped there and the tier vanishes; `output_config` is what survives."""
from litellm.types.llms.anthropic import AnthropicMessagesRequest
from litellm.utils import get_optional_params
adapter = LiteLLMAnthropicMessagesAdapter()
openai_request, _ = adapter.translate_anthropic_to_openai(
anthropic_message_request=AnthropicMessagesRequest(
model=ARN_MODEL,
max_tokens=1024,
messages=[{"role": "user", "content": "hi"}],
thinking={"type": "adaptive"},
output_config={"effort": "max"},
)
)
assert openai_request["output_config"] == {"effort": "max"}
assert "reasoning_effort" not in openai_request
on_the_wire = get_optional_params(
model=ARN_MODEL,
custom_llm_provider="bedrock",
thinking=openai_request["thinking"],
output_config=openai_request["output_config"],
)
assert on_the_wire["output_config"] == {"effort": "max"}
def test_a_bedrock_target_keeps_a_caller_set_thinking_display():
"""`output_config` attaches the tier without touching `thinking`, so a caller who asked for
`display: omitted` still gets it. Carrying the tier as `reasoning_effort` instead lets the
provider mapping rewrite that block."""
from litellm.types.llms.anthropic import AnthropicMessagesRequest
from litellm.utils import get_optional_params
thinking = {"type": "adaptive", "display": "omitted"}
adapter = LiteLLMAnthropicMessagesAdapter()
openai_request, _ = adapter.translate_anthropic_to_openai(
anthropic_message_request=AnthropicMessagesRequest(
model="bedrock/converse/us.anthropic.claude-opus-4-7",
max_tokens=1024,
messages=[{"role": "user", "content": "hi"}],
thinking=thinking,
output_config={"effort": "max"},
)
)
on_the_wire = get_optional_params(
model="converse/us.anthropic.claude-opus-4-7",
custom_llm_provider="bedrock",
thinking=openai_request["thinking"],
output_config=openai_request["output_config"],
)
assert on_the_wire["thinking"] == thinking
assert on_the_wire["output_config"] == {"effort": "max"}
def test_a_non_claude_target_keeps_its_summary_wrapping():
"""The negative class: a target that gets no `thinking` block has nowhere else to put the
summary, so the wrapped dict is still the right shape there."""
from litellm.types.llms.anthropic import AnthropicMessagesRequest
adapter = LiteLLMAnthropicMessagesAdapter()
openai_request, _ = adapter.translate_anthropic_to_openai(
anthropic_message_request=AnthropicMessagesRequest(
model="gpt-5-mini",
max_tokens=1024,
messages=[{"role": "user", "content": "hi"}],
thinking={"type": "adaptive", "summary": "detailed"},
output_config={"effort": "max"},
)
)
assert openai_request["reasoning_effort"] == {"effort": "max", "summary": "detailed"}
assert "thinking" not in openai_request
def test_a_databricks_target_trades_its_thinking_display_for_the_tier():
"""The one accepted cost of carrying the tier as `reasoning_effort`: databricks rebuilds the
thinking block while mapping it, so a caller-set `display` is replaced. Pinned rather than left
silent. It only takes `output_config` when litellm sends one, which this bridge cannot do for a
provider whose own supported-params list omits it, so the tier is the thing worth keeping here.
Bedrock avoids this entirely by taking `output_config` directly."""
from litellm.types.llms.anthropic import AnthropicMessagesRequest
from litellm.utils import get_optional_params
adapter = LiteLLMAnthropicMessagesAdapter()
openai_request, _ = adapter.translate_anthropic_to_openai(
anthropic_message_request=AnthropicMessagesRequest(
model="databricks/databricks-claude-opus-4-7",
max_tokens=1024,
messages=[{"role": "user", "content": "hi"}],
thinking={"type": "adaptive", "display": "omitted"},
output_config={"effort": "max"},
)
)
on_the_wire = get_optional_params(
model="databricks-claude-opus-4-7",
custom_llm_provider="databricks",
thinking=openai_request["thinking"],
reasoning_effort=openai_request["reasoning_effort"],
)
assert on_the_wire["output_config"] == {"effort": "max"}
assert on_the_wire["thinking"]["display"] == "summarized"

View file

@ -14,6 +14,8 @@ from unittest.mock import patch
import pytest
import litellm
from litellm.llms.anthropic.experimental_pass_through.utils import (
normalize_reasoning_effort_value,
)
@ -291,3 +293,91 @@ class TestAdapterAdaptiveThinking:
)
assert result is not None
assert result["effort"] == "medium"
class TestDeclaredEffortsAnswerTheDegradationGate:
"""Without this the chain reads only the per-level booleans, so a kimi-k3 request asking for
max silently arrives as high."""
@pytest.mark.parametrize(
"model, provider",
[("kimi-k3", "moonshot"), ("kimi-k3", "fireworks_ai"), ("kimi-k3-us", "fireworks_ai")],
)
def test_a_declared_level_survives_instead_of_degrading(self, local_model_cost_map, model, provider):
assert normalize_reasoning_effort_value("max", model, provider) == "max"
def test_a_level_the_entry_does_not_declare_still_degrades(self, local_model_cost_map):
"""xhigh is not on kimi-k3's declaration, so it must keep degrading rather than be waved
past by the mere presence of one."""
assert normalize_reasoning_effort_value("xhigh", "kimi-k3", "moonshot") == "high"
assert normalize_reasoning_effort_value("minimal", "kimi-k3", "moonshot") == "low"
def test_the_wider_perplexity_entry_keeps_the_levels_it_declares(self, local_model_cost_map):
assert normalize_reasoning_effort_value("xhigh", "perplexity/kimi-k3", "perplexity") == "xhigh"
assert normalize_reasoning_effort_value("minimal", "perplexity/kimi-k3", "perplexity") == "minimal"
@pytest.mark.parametrize(
"model, provider, effort, expected",
[
("claude-opus-4-7", "anthropic", "max", "max"),
("claude-sonnet-4-6", "anthropic", "minimal", "low"),
("gpt-5-mini", "azure", "max", "high"),
],
)
def test_an_entry_on_the_per_level_flags_is_untouched(
self, local_model_cost_map, model, provider, effort, expected
):
"""The negative class that bounds this change to entries carrying the key."""
assert normalize_reasoning_effort_value(effort, model, provider) == expected
class TestDeclarationBeatsThePerLevelFlags:
"""An entry can carry both shapes. The declaration wins whole, or /model_group/info and this
path would disagree about the same deployment. Driven through the public entry point over a
seeded map entry rather than a patched get_model_info, so it pins behaviour and not wiring."""
MODEL = "declared-and-flagged"
@pytest.fixture
def seeded(self, local_model_cost_map, monkeypatch):
def _seed(**entry):
monkeypatch.setitem(
litellm.model_cost,
self.MODEL,
{"litellm_provider": "openai", "mode": "chat", "supports_reasoning": True, **entry},
)
litellm.get_model_info.cache_clear()
return _seed
@pytest.mark.parametrize("effort, expected", [("max", "max"), ("xhigh", "high"), ("minimal", "low")])
def test_a_flag_cannot_re_add_a_level_the_declaration_omits(self, seeded, effort, expected):
seeded(
reasoning_effort_levels=["low", "high", "max"],
supports_xhigh_reasoning_effort=True,
supports_minimal_reasoning_effort=True,
supports_max_reasoning_effort=False,
)
assert normalize_reasoning_effort_value(effort, self.MODEL, "openai") == expected
def test_a_flag_cannot_keep_max_when_the_declaration_drops_it(self, seeded):
seeded(
reasoning_effort_levels=["low", "high"],
supports_max_reasoning_effort=True,
supports_xhigh_reasoning_effort=True,
)
assert normalize_reasoning_effort_value("max", self.MODEL, "openai") == "high"
def test_a_false_flag_cannot_remove_a_level_the_declaration_names(self, seeded):
seeded(reasoning_effort_levels=["high", "xhigh"], supports_xhigh_reasoning_effort=False)
assert normalize_reasoning_effort_value("xhigh", self.MODEL, "openai") == "xhigh"
assert normalize_reasoning_effort_value("max", self.MODEL, "openai") == "xhigh"
def test_a_chain_the_declaration_omits_entirely_lands_on_its_terminal(self, seeded):
"""Documented residual: no strength ordering exists to pick a nearer declared level."""
seeded(reasoning_effort_levels=["high", "xhigh"])
assert normalize_reasoning_effort_value("minimal", self.MODEL, "openai") == "low"

View file

@ -102,6 +102,35 @@ def test_transform_request_hoists_tool_message_image():
]
def test_transform_request_drops_tool_reference_parts():
"""Azure's transform_request shares the tool-message sanitizing with OpenAI:
tool_reference parts are dropped, a reference-only result keeps its tool
message with empty text (#37462 round trip)."""
messages = [
{"role": "user", "content": "load the WebFetch tool"},
{
"role": "assistant",
"content": None,
"tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "ToolSearch", "arguments": "{}"}}],
},
{
"role": "tool",
"tool_call_id": "call_1",
"content": [{"type": "tool_reference", "tool_name": "WebFetch"}],
},
]
request = AzureOpenAIConfig().transform_request(
model="gpt-4o",
messages=messages,
optional_params={},
litellm_params={},
headers={},
)
assert request["messages"][2]["content"] == ""
@pytest.mark.parametrize(
"model, emitted_key, absent_key",
[

View file

@ -1901,6 +1901,76 @@ async def test_async_audio_transcriptions_sends_dict_data_as_json_body():
assert response.text == "transcribed"
class _WordTimestampAudioTranscriptionConfig(_JSONBodyAudioTranscriptionConfig):
def transform_audio_transcription_response(self, raw_response):
payload = raw_response.json()
response = TranscriptionResponse(text=payload["text"])
response["words"] = payload["words"]
return response
def test_transform_audio_transcription_response_without_subtitle_opt_in_keeps_text_and_words():
words = [
{"word": "hello", "start": 0.0, "end": 0.5},
{"word": "world", "start": 0.5, "end": 1.0},
]
raw_response = httpx.Response(200, json={"text": "hello world", "words": words})
response = BaseLLMHTTPHandler()._transform_audio_transcription_response(
provider_config=_WordTimestampAudioTranscriptionConfig(),
model="test-model",
response=raw_response,
model_response=TranscriptionResponse(),
logging_obj=Mock(),
optional_params={"response_format": "srt"},
api_key=None,
)
assert response.text == "hello world"
assert response["words"] == words
class _SubtitleSynthesisAudioTranscriptionConfig(_JSONBodyAudioTranscriptionConfig):
@property
def supports_subtitle_synthesis(self) -> bool:
return True
def transform_audio_transcription_response(self, raw_response):
payload = raw_response.json()
response = TranscriptionResponse(text=payload["text"])
if "words" in payload:
response["words"] = payload["words"]
return response
def _transform_subtitle_response(payload):
return BaseLLMHTTPHandler()._transform_audio_transcription_response(
provider_config=_SubtitleSynthesisAudioTranscriptionConfig(),
model="test-model",
response=httpx.Response(200, json=payload),
model_response=TranscriptionResponse(),
logging_obj=Mock(),
optional_params={"response_format": "srt"},
api_key=None,
)
def test_subtitle_synthesis_fallback_without_timings_drops_words():
response = _transform_subtitle_response(
{"text": "hello world", "words": [{"word": "hello"}, {"word": "world"}]}
)
assert response.text == "hello world"
assert "words" not in response
def test_subtitle_synthesis_without_words_keeps_plain_text():
response = _transform_subtitle_response({"text": "hello world"})
assert response.text == "hello world"
assert "words" not in response
@pytest.mark.asyncio
async def test_async_retrieve_file_content_raises_on_http_error():
"""

View file

@ -169,6 +169,42 @@ class TestTransformRequest:
}
}
@pytest.mark.parametrize("response_format", ["srt", "vtt"])
def test_subtitle_response_format_requests_word_timestamps(self, config, response_format):
request_data = config.transform_audio_transcription_request(
model="gemini-3.5-transcribe",
audio_file=("sample.wav", AUDIO_BYTES, "audio/wav"),
optional_params={"response_format": response_format},
litellm_params={},
)
transcription_config = request_data.data["generation_config"]["transcription_config"]
assert json.loads(json.dumps(transcription_config)) == {
"mode": {
"type": "verbatim",
"timestamp_granularities": ["word"],
"diarization_mode": "speaker",
}
}
@pytest.mark.parametrize("response_format", ["json", "text", "verbose_json"])
def test_non_subtitle_response_format_sends_no_mode(self, config, response_format):
request_data = config.transform_audio_transcription_request(
model="gemini-3.5-transcribe",
audio_file=("sample.wav", AUDIO_BYTES, "audio/wav"),
optional_params={"response_format": response_format},
litellm_params={},
)
assert "generation_config" not in request_data.data
def test_non_string_response_format_sends_no_mode(self, config):
request_data = config.transform_audio_transcription_request(
model="gemini-3.5-transcribe",
audio_file=("sample.wav", AUDIO_BYTES, "audio/wav"),
optional_params={"response_format": {"type": "json_object"}},
litellm_params={},
)
assert "generation_config" not in request_data.data
def test_segment_granularity_sends_no_mode(self, config):
request_data = config.transform_audio_transcription_request(
model="gemini-3.5-transcribe",
@ -214,6 +250,54 @@ class TestTransformResponse:
assert response.get("duration") is None
class TestSubtitleSynthesisThroughHandler:
def _transform(self, config, response_format):
from unittest.mock import Mock
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.utils import TranscriptionResponse
return BaseLLMHTTPHandler()._transform_audio_transcription_response(
provider_config=config,
model="gemini-3.5-transcribe",
response=make_response(COMPLETED_RESPONSE),
model_response=TranscriptionResponse(),
logging_obj=Mock(),
optional_params={"response_format": response_format},
api_key=None,
)
def test_supports_subtitle_synthesis(self, config):
assert config.supports_subtitle_synthesis is True
def test_srt_synthesizes_subtitle_document_and_drops_words(self, config):
response = self._transform(config, "srt")
assert response.text == (
"1\n00:00:00,100 --> 00:00:00,400\nHello\n\n2\n00:00:00,500 --> 00:00:00,900\nworld.\n"
)
assert "words" not in response
assert response["task"] == "transcribe"
assert response["duration"] == 0.9
assert response.usage.total_tokens == 200
def test_vtt_synthesizes_subtitle_document_and_drops_words(self, config):
response = self._transform(config, "vtt")
assert response.text == (
"WEBVTT\n\n00:00:00.100 --> 00:00:00.400\nHello\n\n00:00:00.500 --> 00:00:00.900\nworld.\n"
)
assert "words" not in response
assert response.usage.total_tokens == 200
@pytest.mark.parametrize("response_format", ["json", "verbose_json"])
def test_non_subtitle_formats_keep_plain_text_and_words(self, config, response_format):
response = self._transform(config, response_format)
assert response.text == "Hello world."
assert response["words"] == [
{"word": "Hello", "start": 0.1, "end": 0.4, "speaker": "spk:0"},
{"word": "world.", "start": 0.5, "end": 0.9, "speaker": "spk:1"},
]
class TestCostRegression:
@pytest.fixture
def local_cost_map(self, monkeypatch):

View file

@ -1866,6 +1866,54 @@ def test_map_openai_params_drops_stock_voice_case_insensitively():
assert passthrough["generationConfig"]["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore"
def test_gemini_response_done_bills_audio_output_tokens_at_audio_rate(monkeypatch):
"""Regression for the Gemini Live AUDIO output breakdown: responseTokensDetails
must survive into response.done usage and bill at output_cost_per_audio_token,
not the text rate."""
from litellm.cost_calculator import (
RealtimeAPITokenUsageProcessor,
handle_realtime_stream_cost_calculation,
)
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
config = GeminiRealtimeConfig()
done_event = config.transform_response_done_event(
message={
"serverContent": {"turnComplete": True},
"usageMetadata": {
"promptTokenCount": 377,
"responseTokenCount": 51,
"totalTokenCount": 428,
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 377}],
"responseTokensDetails": [{"modality": "AUDIO", "tokenCount": 51}],
"thoughtsTokenCount": 37,
},
},
current_response_id="resp_lit6277",
current_conversation_id="conv_lit6277",
output_items=None,
)
usage = done_event["response"]["usage"]
assert usage["output_tokens_details"]["audio_tokens"] == 51
assert usage["output_token_details"]["audio_tokens"] == 51
results = [done_event]
combined_usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(
results=results,
)
assert combined_usage.completion_tokens_details is not None
assert combined_usage.completion_tokens_details.audio_tokens == 51
cost = handle_realtime_stream_cost_calculation(
results=results,
combined_usage_object=combined_usage,
custom_llm_provider="gemini",
litellm_model_name="gemini-2.5-flash-native-audio-preview-12-2025",
)
assert cost == pytest.approx(377 * 5e-07 + 51 * 1.2e-05 + 37 * 2e-06)
@pytest.fixture(autouse=False)
def patch_gemini_transcribe_live_cost_map_entry(monkeypatch):
"""Inject the gemini-3.5-transcribe-live registry entry locally.

View file

@ -869,6 +869,69 @@ class TestToolMessageImageHoisting:
assert result[3]["content"] == self.HOISTED_USER_CONTENT
class TestToolReferenceStripping:
"""transform_request drops tool_reference parts from tool messages: OpenAI's
chat API rejects them, and the reference names an already-declared tool
rather than carrying content (#37462 round trip)."""
def setup_method(self):
self.config = OpenAIGPTConfig()
def _messages_with_tool_reference(self, extra_parts=()):
return [
{"role": "user", "content": "load the WebFetch tool"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{"id": "call_1", "type": "function", "function": {"name": "ToolSearch", "arguments": "{}"}}
],
},
{
"role": "tool",
"tool_call_id": "call_1",
"content": [*extra_parts, {"type": "tool_reference", "tool_name": "WebFetch"}],
},
]
def test_transform_request_keeps_text_and_drops_reference(self):
request = self.config.transform_request(
model="gpt-4.1",
messages=self._messages_with_tool_reference(extra_parts=({"type": "text", "text": "loaded"},)),
optional_params={},
litellm_params={},
headers={},
)
tool_message = request["messages"][2]
assert tool_message["content"] == [{"type": "text", "text": "loaded"}]
assert tool_message["tool_call_id"] == "call_1"
def test_transform_request_reference_only_keeps_tool_message_with_empty_text(self):
request = self.config.transform_request(
model="gpt-4.1",
messages=self._messages_with_tool_reference(),
optional_params={},
litellm_params={},
headers={},
)
assert [m.get("role") for m in request["messages"]] == ["user", "assistant", "tool"]
assert request["messages"][2]["content"] == ""
@pytest.mark.asyncio
async def test_async_transform_request_drops_reference(self):
request = await self.config.async_transform_request(
model="gpt-4.1",
messages=self._messages_with_tool_reference(),
optional_params={},
litellm_params={},
headers={},
)
assert request["messages"][2]["content"] == ""
class TestOpenAIPromptCacheBreakpointChatPath:
"""Chat-path shape for OpenAI explicit prompt caching (#37509)."""

View file

@ -477,12 +477,12 @@ class TestBuildResponseWithResponseFormat:
}
}
# SRT requested but tokens have no start_ms/end_ms -> empty SRT
# falls back gracefully since _group_tokens_into_cues skips them
# falls back gracefully since group_subtitle_tokens_into_cues skips them
resp = cfg._build_response_from_payload(payload, response_format="srt")
# With no timestamp data, SRT rendering produces empty string,
# but we still get output because the code checks `tokens` truthiness
# before choosing SRT path. Actually the tokens list is truthy but
# _group_tokens_into_cues will produce no cues -> empty SRT string.
# group_subtitle_tokens_into_cues will produce no cues -> empty SRT string.
# Let's verify it doesn't crash.
assert isinstance(resp.text, str)

View file

@ -18,8 +18,14 @@ from litellm.types.utils import LlmProviders, ModelResponse
TOOL_CALLING_MODEL = "openai/gpt-oss-20b"
REASONING_MODEL = "deepseek-ai/DeepSeek-V3.1"
PLAIN_MODEL = "Qwen/Qwen3-235B-A22B-fp8-tput"
UNMAPPED_MODEL = "example-org/brand-new-model"
NO_TOOLS_MODEL = "example-org/no-tools-model"
ADJUSTABLE_REASONING_MODEL = "openai/gpt-oss-120b"
HYBRID_REASONING_MODEL = "Qwen/Qwen3.5-9B"
HIGH_MAX_REASONING_MODEL = "deepseek-ai/DeepSeek-V4-Pro"
REGISTRY_FLAGGED_REASONING_MODEL = "zai-org/GLM-4.6"
NON_REASONING_MODEL = "meta-llama/Llama-3.3-70B-Instruct-Turbo"
NO_SCHEMA_MODEL = "example-org/no-schema-model"
TOOL_PARAMS = ("tools", "tool_choice", "function_call")
@ -39,6 +45,15 @@ JSON_SCHEMA_RESPONSE_FORMAT = {
REGEX_RESPONSE_FORMAT = {"type": "regex", "pattern": "(positive|neutral|negative)"}
def _map_reasoning_effort(model: str, effort: str) -> dict:
return TogetherAIChatConfig().map_openai_params(
non_default_params={"reasoning_effort": effort},
optional_params={},
model=model,
drop_params=False,
)
@pytest.fixture(autouse=True)
def force_local_model_cost(monkeypatch):
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
@ -191,6 +206,116 @@ def test_map_openai_params_schema_model_passes_response_format_through(response_
assert mapped["response_format"] == response_format
@pytest.mark.parametrize(
"model",
[ADJUSTABLE_REASONING_MODEL, HYBRID_REASONING_MODEL, HIGH_MAX_REASONING_MODEL, REGISTRY_FLAGGED_REASONING_MODEL],
)
def test_supported_params_includes_reasoning_effort_for_reasoning_models(model):
supported = TogetherAIChatConfig().get_supported_openai_params(model=model)
assert "reasoning_effort" in supported
@pytest.mark.parametrize("model", [NON_REASONING_MODEL, PLAIN_MODEL])
def test_supported_params_excludes_reasoning_effort_for_non_reasoning_models(model):
supported = TogetherAIChatConfig().get_supported_openai_params(model=model)
assert "reasoning_effort" not in supported
@pytest.mark.parametrize(
"effort, expected",
[("low", "low"), ("medium", "medium"), ("high", "high"), ("minimal", "low"), ("xhigh", "high"), ("max", "high")],
)
def test_adjustable_model_translates_reasoning_effort(effort, expected):
mapped = _map_reasoning_effort(ADJUSTABLE_REASONING_MODEL, effort)
assert mapped["reasoning_effort"] == expected
assert "reasoning" not in mapped
def test_adjustable_model_cannot_disable_reasoning_so_none_becomes_low():
mapped = _map_reasoning_effort(ADJUSTABLE_REASONING_MODEL, "none")
assert mapped["reasoning_effort"] == "low"
assert "reasoning" not in mapped
@pytest.mark.parametrize(
"effort, expected",
[("low", "low"), ("medium", "medium"), ("high", "high"), ("minimal", "low"), ("xhigh", "high"), ("max", "high")],
)
def test_hybrid_model_translates_reasoning_effort(effort, expected):
mapped = _map_reasoning_effort(HYBRID_REASONING_MODEL, effort)
assert mapped["reasoning_effort"] == expected
assert "reasoning" not in mapped
@pytest.mark.parametrize("model", [HYBRID_REASONING_MODEL, HIGH_MAX_REASONING_MODEL, REGISTRY_FLAGGED_REASONING_MODEL])
def test_reasoning_effort_none_becomes_reasoning_toggle(model):
mapped = _map_reasoning_effort(model, "none")
assert mapped["reasoning"] == {"enabled": False}
assert "reasoning_effort" not in mapped
def test_reasoning_effort_none_does_not_clobber_user_reasoning():
mapped = TogetherAIChatConfig().map_openai_params(
non_default_params={"reasoning_effort": "none"},
optional_params={"reasoning": {"enabled": True}},
model=HYBRID_REASONING_MODEL,
drop_params=False,
)
assert mapped["reasoning"] == {"enabled": True}
assert "reasoning_effort" not in mapped
@pytest.mark.parametrize(
"effort, expected",
[("minimal", "high"), ("low", "high"), ("medium", "high"), ("high", "high"), ("xhigh", "max"), ("max", "max")],
)
def test_deepseek_v4_pro_remaps_to_high_max(effort, expected):
mapped = _map_reasoning_effort(HIGH_MAX_REASONING_MODEL, effort)
assert mapped["reasoning_effort"] == expected
def test_deepseek_v4_pro_dated_variant_remaps_via_prefix():
mapped = _map_reasoning_effort(f"{HIGH_MAX_REASONING_MODEL}-0813", "low")
assert mapped["reasoning_effort"] == "high"
@pytest.mark.parametrize("model", [ADJUSTABLE_REASONING_MODEL, HYBRID_REASONING_MODEL, HIGH_MAX_REASONING_MODEL])
def test_reasoning_effort_default_is_dropped(model):
mapped = _map_reasoning_effort(model, "default")
assert "reasoning_effort" not in mapped
assert "reasoning" not in mapped
def test_get_optional_params_translates_reasoning_effort_for_together():
optional_params = litellm.get_optional_params(
model=ADJUSTABLE_REASONING_MODEL,
custom_llm_provider="together_ai",
reasoning_effort="max",
)
assert optional_params["reasoning_effort"] == "high"
def test_get_optional_params_rejects_reasoning_effort_for_non_reasoning_together_model():
with pytest.raises(litellm.UnsupportedParamsError):
litellm.get_optional_params(
model=NON_REASONING_MODEL,
custom_llm_provider="together_ai",
reasoning_effort="low",
drop_params=False,
)
@pytest.mark.parametrize("drop_params", [False, True])
def test_map_openai_params_unmapped_model_passes_response_format_through(drop_params, together_warning_log):
mapped = TogetherAIChatConfig().map_openai_params(

View file

@ -10,6 +10,7 @@ from types import SimpleNamespace
import pytest
from fastapi import HTTPException
from pydantic import ValidationError
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
oauth_protected_resource_path,
@ -598,3 +599,92 @@ def test_client_credentials_uses_admin_entered_token_url_when_issuer_yield_empti
assert spec is not None
assert isinstance(spec.config, ClientCredentialsConfig)
assert spec.config.token_url == "https://idp.example.com/token"
_M2M_FIELDS = dict(
auth_type=MCPAuth.oauth2,
oauth2_flow="client_credentials",
client_id="cid",
client_secret="csec",
token_url="https://idp.example.com/token",
)
_OBO_FIELDS = dict(
auth_type=MCPAuth.oauth2_token_exchange,
client_id="cid",
client_secret="csec",
token_exchange_endpoint="https://idp.example.com/token",
)
_ID_JAG_FIELDS = dict(
auth_type=MCPAuth.oauth2_id_jag,
client_id="cid",
client_secret="csec",
token_exchange_endpoint="https://idp.example.com/token",
id_jag_resource_token_endpoint="https://mcp-as.example.com/token",
audience="api://mcp",
)
_AUTHZ_CODE_FIELDS = dict(auth_type=MCPAuth.oauth2, url="https://up.example.com/mcp")
_STATIC_FIELDS = dict(auth_type=MCPAuth.bearer_token, authentication_token="static-tok")
_ARM_FIELDS = (
("client_credentials", _M2M_FIELDS),
("token_exchange", _OBO_FIELDS),
("id_jag", _ID_JAG_FIELDS),
("authorization_code", _AUTHZ_CODE_FIELDS),
("api_key", _STATIC_FIELDS),
)
@pytest.mark.parametrize("name,fields", _ARM_FIELDS, ids=[n for n, _ in _ARM_FIELDS])
def test_upstream_token_header_reaches_every_arms_config(name, fields):
# to_server_spec builds each arm's config from a hand-written kwargs list, so an arm that
# forgets to read the field fails silently: the server keeps writing to Authorization.
spec = to_server_spec(_server(upstream_token_header="esb-oauth", **fields))
assert spec is not None
assert spec.config.header_name == "esb-oauth"
@pytest.mark.parametrize("name,fields", _ARM_FIELDS, ids=[n for n, _ in _ARM_FIELDS])
def test_omitting_the_field_keeps_each_arms_shipped_default(name, fields):
spec = to_server_spec(_server(**fields))
assert spec is not None
assert spec.config.header_name == "Authorization"
def test_api_key_scheme_default_survives_when_the_field_is_unset():
spec = to_server_spec(_server(auth_type=MCPAuth.api_key, authentication_token="k"))
assert spec is not None
assert spec.config.header_name == "X-API-Key"
assert spec.config.value_prefix == ""
def test_the_field_overrides_the_api_key_scheme_default():
spec = to_server_spec(_server(auth_type=MCPAuth.api_key, authentication_token="k", upstream_token_header="X-Esb"))
assert spec is not None
assert spec.config.header_name == "X-Esb"
@pytest.mark.parametrize("bad", ["with space", "has:colon", "trailing\r\nX-Injected", 'quoted"name'])
def test_a_malformed_header_name_is_refused_when_the_server_is_built(bad):
"""Validation belongs at ingestion, not at spec building. Raising inside to_server_spec would
abort the whole aggregate tools/list, so one mistyped server would silently empty the tool list
for every other server too. Refusing at MCPServer construction fails the config load loudly
instead, and means no malformed value can ever reach an arm.
"""
with pytest.raises(ValidationError):
_server(upstream_token_header=bad, **_M2M_FIELDS)
def test_a_valid_header_name_is_trimmed_at_ingestion():
assert _server(upstream_token_header=" esb-oauth ", **_M2M_FIELDS).upstream_token_header == "esb-oauth"
@pytest.mark.parametrize("blank", ["", " ", "\t"])
def test_a_blank_header_name_means_unset_rather_than_an_error(blank):
"""The management API treats a blank as "not supplied" and stores it, so raising here made every
later rebuild of that server 500 instead of falling back to the default Authorization behavior.
"""
server = _server(upstream_token_header=blank, **_M2M_FIELDS)
assert server.upstream_token_header is None
spec = to_server_spec(server)
assert spec is not None
assert spec.config.header_name == "Authorization"

View file

@ -341,7 +341,7 @@ async def test_bearer_auth_sends_the_token_and_leaves_a_success_alone():
async def refetch(failed: str) -> "str | None":
raise AssertionError("must not refetch on success")
auth = ClientCredentialsBearerAuth("m2m-token", refetch)
auth = ClientCredentialsBearerAuth("m2m-token", refetch, ClientCredentialsConfig())
async with httpx.AsyncClient(transport=transport, auth=auth) as client:
response = await client.get("https://upstream.example.com/mcp")
assert response.status_code == 200
@ -357,7 +357,7 @@ async def test_bearer_auth_retries_a_401_once_with_a_fresh_token():
refetched.append(failed)
return "fresh-token"
auth = ClientCredentialsBearerAuth("stale-token", refetch)
auth = ClientCredentialsBearerAuth("stale-token", refetch, ClientCredentialsConfig())
async with httpx.AsyncClient(transport=transport, auth=auth) as client:
response = await client.get("https://upstream.example.com/mcp")
assert response.status_code == 200
@ -377,7 +377,7 @@ async def test_bearer_auth_remembers_the_rotated_token_for_later_requests():
refetched.append(failed)
return "fresh-token"
auth = ClientCredentialsBearerAuth("stale-token", refetch)
auth = ClientCredentialsBearerAuth("stale-token", refetch, ClientCredentialsConfig())
async with httpx.AsyncClient(transport=transport, auth=auth) as client:
first = await client.get("https://upstream.example.com/mcp")
second = await client.get("https://upstream.example.com/mcp")
@ -393,7 +393,7 @@ async def test_bearer_auth_surfaces_the_401_when_the_refetch_fails():
async def refetch(failed: str) -> "str | None":
return None
auth = ClientCredentialsBearerAuth("stale-token", refetch)
auth = ClientCredentialsBearerAuth("stale-token", refetch, ClientCredentialsConfig())
async with httpx.AsyncClient(transport=transport, auth=auth) as client:
response = await client.get("https://upstream.example.com/mcp")
assert response.status_code == 401
@ -409,7 +409,7 @@ async def test_bearer_auth_gives_up_after_a_second_401():
refetched.append(failed)
return "fresh-token"
auth = ClientCredentialsBearerAuth("stale-token", refetch)
auth = ClientCredentialsBearerAuth("stale-token", refetch, ClientCredentialsConfig())
async with httpx.AsyncClient(transport=transport, auth=auth) as client:
response = await client.get("https://upstream.example.com/mcp")
assert response.status_code == 401
@ -421,7 +421,60 @@ def test_bearer_auth_rejects_sync_clients():
async def refetch(failed: str) -> "str | None":
return None
auth = ClientCredentialsBearerAuth("token", refetch)
auth = ClientCredentialsBearerAuth("token", refetch, ClientCredentialsConfig())
with httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200)), auth=auth) as client:
with pytest.raises(RuntimeError):
client.get("https://upstream.example.com/mcp")
@pytest.mark.asyncio
async def test_bearer_auth_writes_the_minted_token_to_the_configured_header():
seen: "list[dict[str, str]]" = []
def handler(request: httpx.Request) -> httpx.Response:
seen.append(dict(request.headers))
return httpx.Response(200)
async def refetch(failed: str) -> "str | None":
raise AssertionError("must not refetch on success")
auth = ClientCredentialsBearerAuth("m2m-token", refetch, ClientCredentialsConfig(header_name="esb-oauth"))
async with httpx.AsyncClient(transport=httpx.MockTransport(handler), auth=auth) as client:
await client.get("https://upstream.example.com/mcp")
assert seen[0]["esb-oauth"] == "Bearer m2m-token"
assert "authorization" not in seen[0]
@pytest.mark.asyncio
async def test_the_401_refetch_retry_also_targets_the_configured_header():
# The retry is a SECOND write of the credential. Honoring the carrier only on the first write
# would silently send the fresh token to Authorization, so the ESB rejects every recovered
# request while the first attempt looked correct.
seen: "list[dict[str, str]]" = []
responses = [httpx.Response(401), httpx.Response(200)]
def handler(request: httpx.Request) -> httpx.Response:
seen.append(dict(request.headers))
return responses[min(len(seen) - 1, len(responses) - 1)]
async def refetch(failed: str) -> "str | None":
return "fresh-token"
auth = ClientCredentialsBearerAuth("stale-token", refetch, ClientCredentialsConfig(header_name="esb-oauth"))
async with httpx.AsyncClient(transport=httpx.MockTransport(handler), auth=auth) as client:
response = await client.get("https://upstream.example.com/mcp")
assert response.status_code == 200
assert [h["esb-oauth"] for h in seen] == ["Bearer stale-token", "Bearer fresh-token"]
assert all("authorization" not in h for h in seen)
@pytest.mark.asyncio
async def test_bearer_auth_advertises_the_header_it_will_occupy():
# _resolve_v2_auth reads header_name off the auth object to decide which injected header
# conflicts; an auth object that lies about its slot would drop the wrong one.
async def refetch(failed: str) -> "str | None":
return None
assert ClientCredentialsBearerAuth("t", refetch, ClientCredentialsConfig()).header_name == "Authorization"
default_carrier = ClientCredentialsConfig(header_name="esb-oauth")
assert ClientCredentialsBearerAuth("t", refetch, default_carrier).header_name == "esb-oauth"

View file

@ -1033,3 +1033,70 @@ async def test_invalidate_credentials_for_id_jag_is_a_noop_without_a_caller_toke
assert isinstance(first, Ok) and isinstance(second, Ok)
assert _emitted(second.ok)["Authorization"] == "Bearer cached-bearer"
assert len(endpoint.calls) == 2
async def _resolve_with_carrier(kind: str, header: str):
"""Resolve one minted-token arm whose config targets ``header``."""
if kind == "client_credentials":
source = _FakeM2MSource(Ok(OAuthToken(access_token="minted")))
config = _M2M.model_copy(update={"header_name": header})
provider = UpstreamCredentialProvider(client_credentials_source=source)
return await provider.resolve_credentials(_SUBJECT, _spec(config))
if kind == "token_exchange":
exchanger = _FakeExchanger(Ok(OAuthToken(access_token="minted")))
config = _OBO.model_copy(update={"header_name": header})
subject = Subject(tenant_id="acme", subject_id="alice", inbound_token=SecretStr("caller-jwt"))
provider = UpstreamCredentialProvider(token_exchanger=exchanger)
return await provider.resolve_credentials(subject, _spec(config))
if kind == "authorization_code":
store = _FakeTokenStore({("alice", "s"): OAuthToken(access_token="minted")})
provider = UpstreamCredentialProvider(oauth_token_store=store)
return await provider.resolve_credentials(
Subject(tenant_id="", subject_id="alice"),
_spec(AuthorizationCodeConfig(header_name=header)),
)
endpoint = _FakeTokenEndpoint(
[
Ok(ExchangedToken(access_token="id-jag-assertion", expires_in=300)),
Ok(ExchangedToken(access_token="minted", expires_in=300)),
]
)
config = _id_jag_config().model_copy(update={"header_name": header})
subject = Subject(tenant_id="acme", subject_id="alice", inbound_token=SecretStr("caller-id-token"))
provider = UpstreamCredentialProvider(token_endpoint=endpoint)
return await provider.resolve_credentials(subject, _spec(config))
_MINTED_ARMS = ("client_credentials", "token_exchange", "authorization_code", "id_jag")
@pytest.mark.parametrize("kind", _MINTED_ARMS)
@pytest.mark.asyncio
async def test_every_minted_arm_emits_its_configured_header(kind):
# One arm left on a hardcoded Authorization is a silent no-op for exactly the server that
# configured the knob, so this is asserted across all four rather than on the M2M arm alone.
result = await _resolve_with_carrier(kind, "esb-oauth")
assert isinstance(result, Ok)
headers, _ = await _emitted_async(result.ok)
assert headers["esb-oauth"] == "Bearer minted"
assert "authorization" not in headers
@pytest.mark.parametrize("kind", _MINTED_ARMS)
@pytest.mark.asyncio
async def test_every_minted_arm_still_defaults_to_authorization(kind):
result = await _resolve_with_carrier(kind, "Authorization")
assert isinstance(result, Ok)
headers, _ = await _emitted_async(result.ok)
assert headers["Authorization"] == "Bearer minted"
@pytest.mark.asyncio
async def test_passthrough_ignores_the_carrier_and_keeps_the_callers_slot():
# Passthrough mints nothing: it forwards the caller's own credential, so it has no carrier to
# configure and must keep using the header the caller aimed it at.
subject = Subject(tenant_id="", subject_id="", inbound_token=SecretStr("caller-token"))
result = await UpstreamCredentialProvider().resolve_credentials(subject, _spec(PassthroughConfig()))
assert isinstance(result, Ok)
headers, _ = await _emitted_async(result.ok)
assert headers["Authorization"] == "caller-token"

View file

@ -14,9 +14,11 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials import (
Ambient,
ApiKeyConfig,
AuthConfig,
AuthorizationCodeConfig,
AuthSpecKind,
AwsSigV4Config,
Byok,
ClientCredentialsConfig,
ClientSecretAuth,
CredError,
Error,
@ -27,7 +29,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials import (
ServerSpec,
SharedKey,
StaticKeys,
TokenExchangeConfig,
parse_auth_spec_kind,
validate_header_name,
)
_AUTH_CONFIG = TypeAdapter(AuthConfig)
@ -229,3 +233,61 @@ def test_id_jag_server_spec_derives_auth_spec_kind():
config=config,
)
assert spec.auth_spec_kind is AuthSpecKind.id_jag
_CARRIER_CONFIGS = (
("client_credentials", ClientCredentialsConfig),
("token_exchange", lambda **kw: TokenExchangeConfig(token_exchange_endpoint="https://idp/te", **kw)),
("authorization_code", AuthorizationCodeConfig),
(
"id_jag",
lambda **kw: IdJagConfig(
org_token_endpoint="https://idp.example.com/token",
resource_token_endpoint="https://mcp-as.example.com/token",
client_id="litellm",
client_auth=ClientSecretAuth(client_secret=SecretStr("s")),
**kw,
),
),
("api_key", lambda **kw: ApiKeyConfig(key_source=SharedKey(value=SecretStr("k")), **kw)),
)
@pytest.mark.parametrize("name,build", _CARRIER_CONFIGS, ids=[n for n, _ in _CARRIER_CONFIGS])
def test_every_resolved_credential_config_defaults_to_rfc6750_authorization(name, build):
# The default is what preserves today's wire behavior for every existing server.
assert build().header("tok") == ("Authorization", "Bearer tok")
@pytest.mark.parametrize("name,build", _CARRIER_CONFIGS, ids=[n for n, _ in _CARRIER_CONFIGS])
def test_every_resolved_credential_config_honors_a_custom_header(name, build):
assert build(header_name="esb-oauth").header("tok") == ("esb-oauth", "Bearer tok")
@pytest.mark.parametrize("name,build", _CARRIER_CONFIGS, ids=[n for n, _ in _CARRIER_CONFIGS])
def test_every_resolved_credential_config_can_send_a_raw_value(name, build):
assert build(header_name="esb-oauth", value_prefix="").header("tok") == ("esb-oauth", "tok")
@pytest.mark.parametrize(
"bad",
[
"with space",
"has:colon",
"trailing\r\nX-Injected",
"",
" ",
"quoted\"name",
],
)
def test_header_name_outside_the_rfc7230_token_grammar_is_rejected(bad):
# An operator-supplied name reaches egress verbatim, so anything that could split a
# header must fail closed at construction rather than be sanitized later.
with pytest.raises(ValidationError):
ClientCredentialsConfig(header_name=bad)
assert isinstance(validate_header_name(bad), Error)
def test_header_name_is_trimmed_by_the_one_validator():
assert validate_header_name(" esb-oauth ") == Ok("esb-oauth")
assert ClientCredentialsConfig(header_name=" esb-oauth ").header_name == "esb-oauth"

View file

@ -43,7 +43,6 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_obo_retry_applies,
_resolve_openapi_tool_auth,
_should_strip_caller_authorization,
_without_authorization,
)
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
@ -2405,6 +2404,104 @@ class TestMCPServerManager:
assert client._resolved_auth is not None
assert "authorization" not in {k.lower() for k in (client.extra_headers or {})}
@staticmethod
def _esb_server(header: "str | None") -> MCPServer:
return MCPServer(
server_id="esb",
name="esb-server",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="client_credentials",
client_id="cid",
client_secret="csec",
token_url="https://idp.example.com/token",
upstream_token_header=header,
static_headers={"Authorization": "Bearer static-upstream-mcp-token"},
)
@pytest.mark.asyncio
async def test_static_authorization_survives_a_minted_token_aimed_elsewhere(self):
"""The dual-credential case: an ESB wants the gateway-minted token on its own header while a
separate static Authorization passes through to the origin. Dropping Authorization here (the
old name-blind behavior) deletes the second credential and the upstream 401s."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
class _FakeProvider:
async def resolve_credentials(self, subject, server):
return Ok(StaticHeaderAuth("Bearer MINTED-M2M", header_name="esb-oauth"))
manager = MCPServerManager(cred_provider=_FakeProvider())
client = await manager._create_mcp_client(
self._esb_server("esb-oauth"),
extra_headers={"Authorization": "Bearer static-upstream-mcp-token"},
)
assert client._resolved_auth is not None
assert (client.extra_headers or {})["Authorization"] == "Bearer static-upstream-mcp-token"
@pytest.mark.asyncio
async def test_a_minted_token_aimed_at_the_static_header_still_wins_that_slot(self):
"""The negative class of the test above: when the two DO collide the resolver-owned
credential is still authoritative, so the knob cannot be used to smuggle a second
credential into the same slot."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
class _FakeProvider:
async def resolve_credentials(self, subject, server):
return Ok(StaticHeaderAuth("Bearer MINTED-M2M", header_name="esb-oauth"))
manager = MCPServerManager(cred_provider=_FakeProvider())
client = await manager._create_mcp_client(
self._esb_server("esb-oauth"),
extra_headers={"esb-oauth": "Bearer signer-jwt", "X-Trace": "keep-me"},
)
assert client._resolved_auth is not None
assert "esb-oauth" not in {k.lower() for k in (client.extra_headers or {})}
assert (client.extra_headers or {})["X-Trace"] == "keep-me"
@pytest.mark.asyncio
async def test_a_differently_cased_injected_header_is_still_recognised_as_the_collision(self):
"""HTTP header names are case-insensitive, so the conflict check must be too.
A case-sensitive check reports no conflict and hands the injected header back untouched, so
the returned extra_headers still carries a second copy of the credential slot for every
downstream consumer of that dict. httpx happens to collapse the two on the wire, which is
exactly why this needs pinning rather than being left to luck.
"""
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
class _FakeProvider:
async def resolve_credentials(self, subject, server):
return Ok(StaticHeaderAuth("Bearer MINTED", header_name="esb-oauth"))
manager = MCPServerManager(cred_provider=_FakeProvider())
client = await manager._create_mcp_client(
self._esb_server("esb-oauth"),
extra_headers={"ESB-OAuth": "Bearer injected", "X-Trace": "keep"},
)
assert client._resolved_auth is not None
assert not any(k.lower() == "esb-oauth" for k in (client.extra_headers or {}))
assert (client.extra_headers or {})["X-Trace"] == "keep"
def test_without_header_drops_only_the_named_header(self):
from litellm.types.mcp import DEFAULT_CREDENTIAL_HEADER, without_header
headers = {"Authorization": "Bearer a", "esb-oauth": "Bearer b", "X-Trace": "t"}
assert without_header(headers, "ESB-OAuth") == {"Authorization": "Bearer a", "X-Trace": "t"}
assert without_header(headers, DEFAULT_CREDENTIAL_HEADER) == {"esb-oauth": "Bearer b", "X-Trace": "t"}
@pytest.mark.asyncio
async def test_preflight_token_exchange_challenges_on_rejected_subject(self):
"""A subject the IdP rejects must raise the RFC 9728 401 challenge from the preflight, so a
@ -2624,14 +2721,16 @@ class TestMCPServerManager:
if captured_extra_headers:
assert "authorization" not in {k.lower() for k in captured_extra_headers}
def test_without_authorization_drops_only_the_credential(self):
def test_without_header_drops_only_the_credential(self):
from litellm.types.mcp import without_header
# None / empty -> None
assert _without_authorization(None) is None
assert _without_authorization({}) is None
assert without_header(None, "Authorization") is None
assert without_header({}, "Authorization") is None
# Only Authorization present -> nothing left -> None (case-insensitive)
assert _without_authorization({"authorization": "Bearer x"}) is None
assert without_header({"authorization": "Bearer x"}, "Authorization") is None
# Authorization dropped, other headers kept
assert _without_authorization({"Authorization": "Bearer x", "X-Trace-Id": "t"}) == {"X-Trace-Id": "t"}
assert without_header({"Authorization": "Bearer x", "X-Trace-Id": "t"}, "Authorization") == {"X-Trace-Id": "t"}
@pytest.mark.asyncio
async def test_call_regular_mcp_tool_passthrough_forwards_authorization_with_admission_header(
@ -9641,13 +9740,38 @@ class TestMaterializeAuthHeaders:
from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import (
ClientCredentialsBearerAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ClientCredentialsConfig,
)
async def _refetch(_stale: str):
return None
headers = await _materialize_auth_headers(ClientCredentialsBearerAuth("m2m-token", _refetch))
default_carrier = ClientCredentialsConfig()
headers = await _materialize_auth_headers(ClientCredentialsBearerAuth("m2m-token", _refetch, default_carrier))
assert headers == {"Authorization": "Bearer m2m-token"}
@pytest.mark.asyncio
async def test_materialize_follows_the_minted_token_to_a_custom_header(self):
# The OpenAPI arm reads header_name off the auth object rather than assuming Authorization,
# so it carries the knob with no per-arm change. This pins that it stays that way.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_materialize_auth_headers,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import (
ClientCredentialsBearerAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ClientCredentialsConfig,
)
async def _refetch(_stale: str):
return None
esb_carrier = ClientCredentialsConfig(header_name="esb-oauth")
headers = await _materialize_auth_headers(ClientCredentialsBearerAuth("m2m-token", _refetch, esb_carrier))
assert headers == {"esb-oauth": "Bearer m2m-token"}
@pytest.mark.asyncio
async def test_noop_and_none_materialize_to_none(self):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (

View file

@ -13,6 +13,7 @@ import pytest
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
MCPOAuth2TokenCache,
resolve_mcp_auth,
resolved_token_header,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
@ -411,3 +412,50 @@ async def test_m2m_mint_uses_admin_entered_token_url_when_issuer_yield_empties_r
assert result == "m2m-token-configured"
assert mock_client.post.call_args[0][0] == "https://auth.example.com/token"
def _m2m_server(**overrides):
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
fields = dict(
server_id="s",
name="n",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="client_credentials",
client_id="cid",
client_secret="csec",
token_url="https://idp.example.com/token",
)
fields.update(overrides)
return MCPServer(**fields)
def test_resolved_token_header_follows_the_configured_header_for_a_gateway_resolved_token():
# resolve_mcp_auth mints the M2M token on this branch, so the value is the gateway's own and
# follows upstream_token_header.
assert resolved_token_header(_m2m_server(upstream_token_header="esb-oauth")) == "esb-oauth"
def test_resolved_token_header_is_none_when_the_server_configures_nothing():
assert resolved_token_header(_m2m_server()) is None
def test_a_caller_supplied_credential_never_moves():
# The caller aimed their own token at the slot the upstream normally uses. Relocating it would
# break every existing x-mcp-auth caller on a server that sets the field for its own token.
server = _m2m_server(upstream_token_header="esb-oauth")
assert resolved_token_header(server, "Bearer caller-token") is None
assert resolved_token_header(server, {"Authorization": "Bearer caller-token"}) is None
def test_the_header_and_the_value_agree_on_which_branch_they_took():
# The two helpers are read as a pair at one call site, so they must never disagree about
# whether the credential came from the caller or from the server's own config.
import asyncio
server = _m2m_server(upstream_token_header="esb-oauth", authentication_token="static-tok")
caller = "Bearer caller-token"
assert asyncio.run(resolve_mcp_auth(server, caller)) == caller
assert resolved_token_header(server, caller) is None

View file

@ -729,3 +729,121 @@ async def test_local_dispatch_reports_the_outcome_instead_of_success(failure: st
# A non-auth upstream failure stays a 200 with isError, so REST does not report it as a gateway 500
assert result.isError is True
assert "upstream returned HTTP 429" in result.content[0].text
@pytest.mark.parametrize(
"resolved,expect_guard",
[
({"esb-oauth": "Bearer minted"}, True),
({"Authorization": "Bearer minted"}, False),
({}, False),
],
)
def test_only_a_custom_credential_slot_needs_the_redirect_guard(resolved, expect_guard):
"""The OpenAPI arm sends resolved credentials through a redirect-following client, so a custom
slot needs the same cross-origin guard the MCP client installs. Authorization does not: the HTTP
client already strips that one, and taking the guarded path would give up the shared client.
"""
from litellm.types.mcp import DEFAULT_CREDENTIAL_HEADER, same_header
guarded = next((n for n in resolved if not same_header(n, DEFAULT_CREDENTIAL_HEADER)), None)
assert (guarded is not None) is expect_guard
@pytest.mark.asyncio
async def test_the_openapi_arm_drops_a_custom_slot_across_origins():
"""End to end on the hook the OpenAPI arm installs: same origin keeps the credential, a redirect
to another host does not carry it.
"""
import httpx
from litellm.types.mcp import credential_redirect_hook
hook = credential_redirect_hook("https://api.example.com/v1/things", "esb-oauth")
same = httpx.Request("POST", "https://api.example.com/v1/other", headers={"esb-oauth": "Bearer m"})
await hook(same)
assert same.headers["esb-oauth"] == "Bearer m"
foreign = httpx.Request("POST", "https://attacker.example.com/collect", headers={"esb-oauth": "Bearer m"})
await hook(foreign)
assert "esb-oauth" not in foreign.headers
def test_the_openapi_arm_installs_the_guard_when_a_credential_rides_a_custom_slot():
"""Pins the wiring, not just the hook: the arm must actually build a guarded client. Testing the
hook alone passes even if this arm never installs it.
"""
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_resolved_auth_headers,
_upstream_client,
)
token = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"})
try:
client = _upstream_client()
assert client.client.event_hooks["request"], "custom slot must install a redirect guard"
finally:
_request_resolved_auth_headers.reset(token)
def test_the_guarded_client_is_reused_rather_than_built_per_call():
"""A fresh handler per guarded call is never closed, so every OpenAPI tool call on a server that
sets upstream_token_header would leak an httpx client and its connection pool. Both variants
have to come from the shared cache.
"""
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_resolved_auth_headers,
_upstream_client,
)
token = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"})
try:
assert _upstream_client() is _upstream_client()
finally:
_request_resolved_auth_headers.reset(token)
@pytest.mark.asyncio
async def test_the_shared_guard_reads_the_url_from_the_request_context():
"""The hook is one stable object so the client stays cacheable, which means the origin it guards
against has to arrive per request rather than being closed over.
"""
import httpx
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_drop_credential_across_origin,
_request_resolved_auth_headers,
_request_upstream_url,
)
creds = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"})
url = _request_upstream_url.set("https://api.example.com/v1/things")
try:
same = httpx.Request("POST", "https://api.example.com/v1/other", headers={"esb-oauth": "Bearer m"})
await _drop_credential_across_origin(same)
assert same.headers["esb-oauth"] == "Bearer m"
foreign = httpx.Request("POST", "https://attacker.example.com/x", headers={"esb-oauth": "Bearer m"})
await _drop_credential_across_origin(foreign)
assert "esb-oauth" not in foreign.headers
finally:
_request_upstream_url.reset(url)
_request_resolved_auth_headers.reset(creds)
@pytest.mark.parametrize("resolved", [{"Authorization": "Bearer minted"}, {}, None])
def test_the_openapi_arm_keeps_the_shared_client_when_no_guard_is_needed(resolved):
# Authorization is already stripped across origins by the HTTP client, so taking the guarded
# path for it would give up the shared connection pool for nothing.
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_resolved_auth_headers,
_upstream_client,
)
token = _request_resolved_auth_headers.set(resolved)
try:
client = _upstream_client()
assert not client.client.event_hooks.get("request")
finally:
_request_resolved_auth_headers.reset(token)

View file

@ -5274,3 +5274,74 @@ async def test_terminal_failure_logs_usage_and_cost_of_prior_passed_chunks(monke
assert logged["guardrail_cost"] == pytest.approx(0.0003)
assert logged["guardrail_response"]["usage"] == {"contentPolicyUnits": 2, "wordPolicyUnits": 1}
assert "error" in logged["guardrail_response"]
def test_load_credentials_assumes_role_with_external_id():
"""A trust policy requiring sts:ExternalId must be satisfied by the guardrail's aws_external_id."""
import datetime
import boto3
from botocore.exceptions import ClientError
class FakeSTSClient:
"""STS that mirrors a cross-account role whose trust policy requires an ExternalId."""
def get_caller_identity(self):
return {"Arn": "arn:aws:iam::111111111111:user/litellm-proxy-pod"}
def assume_role(self, **params):
if params.get("ExternalId") != "external-id-123":
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:AssumeRole"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": "ASIAASSUMEDROLEKEY",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session-token",
"Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
}
}
guardrail = BedrockGuardrail(
guardrail_name="bedrock-external-id",
event_hook=GuardrailEventHooks.pre_call,
guardrailIdentifier="gr-1",
guardrailVersion="DRAFT",
aws_region_name="us-east-1",
aws_access_key_id="AKIAPODCALLERKEY",
aws_secret_access_key="pod-caller-secret",
aws_role_name="arn:aws:iam::999999999999:role/litellm-guardrail-role",
aws_session_name="litellm-session",
aws_external_id="external-id-123",
)
with patch.object(boto3, "client", return_value=FakeSTSClient()):
credentials, aws_region_name = guardrail._load_credentials()
assert credentials.access_key == "ASIAASSUMEDROLEKEY"
assert credentials.token == "assumed-session-token"
assert aws_region_name == "us-east-1"
def test_initialize_bedrock_forwards_aws_external_id():
"""aws_external_id configured on the guardrail must survive LitellmParams and the initializer."""
from litellm.proxy.guardrails.guardrail_initializers import initialize_bedrock
from litellm.types.guardrails import LitellmParams
litellm_params = LitellmParams(
guardrail="bedrock",
mode="pre_call",
guardrailIdentifier="gr-1",
guardrailVersion="DRAFT",
aws_region_name="us-east-1",
aws_role_name="arn:aws:iam::999999999999:role/litellm-guardrail-role",
aws_external_id="external-id-123",
)
guardrail = initialize_bedrock(litellm_params, {"guardrail_name": "bedrock-external-id"})
try:
assert guardrail.optional_params["aws_external_id"] == "external-id-123"
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.callbacks, guardrail)

View file

@ -1573,6 +1573,7 @@ class TestTemporaryMCPSessionEndpoints:
existing_server.aws_region_name = None
existing_server.aws_service_name = None
existing_server.upstream_resource = None
existing_server.upstream_token_header = None
mock_manager = MagicMock()
mock_manager.get_mcp_server_by_id.return_value = existing_server
@ -1608,6 +1609,7 @@ class TestTemporaryMCPSessionEndpoints:
existing_server.aws_region_name = None
existing_server.aws_service_name = None
existing_server.upstream_resource = None
existing_server.upstream_token_header = None
for key, value in server_overrides.items():
setattr(existing_server, key, value)
@ -1639,6 +1641,23 @@ class TestTemporaryMCPSessionEndpoints:
assert updated.credentials["client_id"] == "client-123"
assert updated.credentials["client_secret"] == "secret-xyz"
def test_upstream_token_header_is_inherited_like_other_admin_config(self):
"""It is admin config rather than a credential, so a session server derived from an existing
one must carry it. Miss it and the derived server silently sends its token to Authorization
while the original sends it to the gateway's header."""
updated = self._inherit_with({}, upstream_token_header="esb-oauth")
assert updated.credentials["upstream_token_header"] == "esb-oauth"
def test_a_supplied_upstream_token_header_does_not_read_as_a_credential(self):
"""It is in the admin-config key set, so submitting only it must still inherit the declared
app rather than reading as "the caller supplied real credentials"."""
updated = self._inherit_with({"upstream_token_header": "esb-oauth"})
assert updated.credentials["client_id"] == "client-123"
assert updated.credentials["client_secret"] == "secret-xyz"
assert updated.credentials["upstream_token_header"] == "esb-oauth"
def test_supplied_credential_still_wins_over_inheritance(self):
"""A caller that supplies a real credential keeps it; inheritance must not overwrite it."""
updated = self._inherit_with({"auth_value": "caller-token"})
@ -2256,6 +2275,7 @@ class TestTemporaryMCPSessionEndpoints:
aws_region_name=None,
aws_service_name=None,
upstream_resource=None,
upstream_token_header=None,
)
built_server = generate_mock_mcp_server_config_record(server_id="temp-server")
mock_manager = MagicMock()

View file

@ -93,6 +93,32 @@ class TestExtractRequestToolNames:
"run_sql",
]
def test_anthropic_openai_format_tools_forwarded_by_bridge(self):
data = {
"tools": [
{"type": "function", "function": {"name": "get_weather"}},
{"name": "run_sql"},
{"googleSearch": {}},
]
}
assert extract_request_tool_names("/v1/messages", data) == [
"get_weather",
"run_sql",
]
def test_anthropic_hybrid_tool_yields_every_name(self):
data = {
"tools": [
{"type": "function", "name": "decoy", "function": {"name": "blocked_fn"}},
{"type": "function", "name": "", "function": {"name": "hidden_fn"}},
]
}
assert extract_request_tool_names("/v1/messages", data) == [
"decoy",
"blocked_fn",
"hidden_fn",
]
def test_generate_content_tools(self):
data = {
"tools": [
@ -159,6 +185,34 @@ class TestCheckToolsAllowlist:
assert exc_info.value.type == ProxyErrorTypes.tool_access_denied
assert "get_weather" in str(exc_info.value.message)
@pytest.mark.asyncio
async def test_disallowed_openai_format_tool_raises_on_messages_route(self):
token = _token(metadata={"allowed_tools": ["other_tool"]})
body = {"tools": [{"type": "function", "function": {"name": "get_weather"}}]}
with pytest.raises(ProxyException) as exc_info:
await check_tools_allowlist(
request_body=body,
valid_token=token,
team_object=None,
route="/v1/messages",
)
assert exc_info.value.type == ProxyErrorTypes.tool_access_denied
assert "get_weather" in str(exc_info.value.message)
@pytest.mark.asyncio
async def test_hybrid_tool_with_decoy_name_raises_on_messages_route(self):
token = _token(metadata={"allowed_tools": ["decoy"]})
body = {"tools": [{"type": "function", "name": "decoy", "function": {"name": "run_sql"}}]}
with pytest.raises(ProxyException) as exc_info:
await check_tools_allowlist(
request_body=body,
valid_token=token,
team_object=None,
route="/v1/messages",
)
assert exc_info.value.type == ProxyErrorTypes.tool_access_denied
assert "run_sql" in str(exc_info.value.message)
@pytest.mark.asyncio
async def test_disallowed_custom_tool_raises_on_responses_route(self):
token = _token(metadata={"allowed_tools": ["other_tool"]})

View file

@ -1,5 +1,6 @@
import pytest
import litellm
from litellm.router_utils.reasoning_effort_capability import (
deployment_is_catalog_mapped,
intersect_supported_reasoning_efforts,
@ -196,3 +197,158 @@ class TestIntersectSupportedReasoningEfforts:
def test_disjoint_sets_intersect_to_empty(self):
assert intersect_supported_reasoning_efforts(["max"], ["minimal"]) == ()
class TestDeclaredEffortList:
"""reasoning_effort_levels is what the catalog DECLARES per deployment;
ModelGroupInfo.supported_reasoning_efforts is what a group COMPUTED. test_router.py pins that
the computed one is never seeded from model_info, so the two names must stay apart."""
def test_a_declared_list_answers_where_no_flag_could(self):
"""No flag can drop medium, so before this key the entry could only stay silent or
over-advertise a level the model does not document."""
resolved = resolve_supported_reasoning_efforts(
{"supports_reasoning": True, "reasoning_effort_levels": ["low", "high", "max"]},
deployment_is_mapped=True,
)
assert resolved == ("low", "high", "max")
def test_a_declared_list_wins_whole_over_the_flags(self):
resolved = resolve_supported_reasoning_efforts(
{
"supports_reasoning": True,
"reasoning_effort_levels": ["low", "high", "max"],
"supports_none_reasoning_effort": True,
"supports_minimal_reasoning_effort": True,
"supports_xhigh_reasoning_effort": True,
"supports_max_reasoning_effort": False,
},
deployment_is_mapped=True,
)
assert resolved == ("low", "high", "max")
def test_a_declaration_is_reordered_into_the_advertisement_order(self):
resolved = resolve_supported_reasoning_efforts(
{"supports_reasoning": True, "reasoning_effort_levels": ["max", "low", "high"]},
deployment_is_mapped=True,
)
assert resolved == ("low", "high", "max")
def test_a_declared_empty_list_empties_the_group(self):
assert (
resolve_supported_reasoning_efforts(
{"supports_reasoning": True, "reasoning_effort_levels": []},
deployment_is_mapped=True,
)
== ()
)
@pytest.mark.parametrize("declared", [["low", "bogus"], ["bogus"], ["low", 7, None]])
def test_an_unknown_level_is_dropped_rather_than_raised(self, declared):
"""A config.yaml model_info block bypasses the map's enum schema, and one mistyped level
must not fail every sibling on the proxy."""
resolved = resolve_supported_reasoning_efforts(
{"supports_reasoning": True, "reasoning_effort_levels": declared},
deployment_is_mapped=True,
)
assert resolved == tuple(effort for effort in ("low",) if effort in declared)
@pytest.mark.parametrize("malformed", ["low,high,max", {"low": True}, 3, True])
def test_a_malformed_declaration_falls_through_to_the_flags(self, malformed):
resolved = resolve_supported_reasoning_efforts(
{
"supports_reasoning": True,
"reasoning_effort_levels": malformed,
"supports_max_reasoning_effort": True,
},
deployment_is_mapped=True,
)
assert resolved == ("none", "minimal", "low", "medium", "high", "max")
def test_a_model_the_map_calls_non_reasoning_ignores_its_declaration(self):
assert (
resolve_supported_reasoning_efforts(
{"supports_reasoning": False, "reasoning_effort_levels": ["low", "high", "max"]},
deployment_is_mapped=True,
)
== ()
)
def test_a_declaration_is_read_through_the_bare_twin(self, monkeypatch):
monkeypatch.setitem(
litellm.model_cost,
"some-declared-reasoner",
{"supports_reasoning": True, "reasoning_effort_levels": ["low", "max"]},
)
resolved = resolve_supported_reasoning_efforts(
{
"supports_reasoning": True,
"litellm_provider": "openai",
"key": "openai/some-declared-reasoner",
},
deployment_is_mapped=True,
)
assert resolved == ("low", "max")
KIMI_K3_PASSTHROUGH_KEYS = (
"azure_ai/FW-Kimi-K3",
"moonshot/kimi-k3",
"together_ai/moonshotai/Kimi-K3",
"fireworks_ai/kimi-k3",
"fireworks_ai/kimi-k3-fast",
"fireworks_ai/kimi-k3-us",
"fireworks_ai/accounts/fireworks/models/kimi-k3",
"fireworks_ai/accounts/fireworks/routers/kimi-k3-fast",
"fireworks_ai/accounts/fireworks/routers/kimi-k3-us",
)
KIMI_K3_PERPLEXITY_KEY = "perplexity/perplexity/kimi-k3"
class TestKimiK3AdvertisesItsDocumentedLevels:
@pytest.mark.parametrize("model_key", KIMI_K3_PASSTHROUGH_KEYS)
def test_a_passthrough_entry_advertises_the_models_own_levels(self, local_model_cost_map, model_key):
"""platform.kimi.ai documents exactly low, high and max, and these providers forward the
level unchanged. Undeclared, each entry resolves to unknown and the dashboard falls back to
a capability-blind list that omits max."""
entry = dict(litellm.model_cost[model_key], key=model_key)
assert resolve_supported_reasoning_efforts(entry, deployment_is_mapped=True) == ("low", "high", "max")
def test_the_perplexity_entry_advertises_the_wider_set_it_maps_down(self, local_model_cost_map):
"""Perplexity's Agent API takes a six-value enum and maps it down internally, so this
deployment is legitimately wider than a passthrough. One blanket list could not say both."""
entry = dict(litellm.model_cost[KIMI_K3_PERPLEXITY_KEY], key=KIMI_K3_PERPLEXITY_KEY)
assert resolve_supported_reasoning_efforts(entry, deployment_is_mapped=True) == (
"minimal",
"low",
"medium",
"high",
"xhigh",
"max",
)
@pytest.mark.parametrize("model, provider", [("kimi-k3", "moonshot"), ("kimi-k3", "fireworks_ai")])
def test_the_declaration_survives_model_info_hydration(self, local_model_cost_map, model, provider):
"""The hydration line is the load-bearing seam: without it the key the map carries never
reaches the resolver and reads as absent everywhere downstream."""
from litellm.utils import _get_model_info_helper
model_info = dict(_get_model_info_helper(model=model, custom_llm_provider=provider))
assert model_info["reasoning_effort_levels"] == ["low", "high", "max"]
assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == ("low", "high", "max")
def test_a_kimi_k3_deployment_now_narrows_a_mixed_group(self, local_model_cost_map):
"""kimi used to contribute unknown, which never narrows, so the group advertised whatever
its other deployments agreed on."""
kimi = resolve_supported_reasoning_efforts(
dict(litellm.model_cost["fireworks_ai/kimi-k3"], key="fireworks_ai/kimi-k3"),
deployment_is_mapped=True,
)
assert intersect_supported_reasoning_efforts(("none", "minimal", "low", "medium", "high", "xhigh"), kimi) == (
"low",
"high",
)

View file

@ -3977,6 +3977,74 @@ def test_completion_cost_prices_anthropic_shaped_cache_read_tokens(_local_model_
assert cost == pytest.approx(3 * 4e-6 + 4014 * 4e-7 + 5 * 2e-5, rel=1e-9)
def _together_chat_response(model: str, prompt_tokens: int, completion_tokens: int, cached_tokens: int) -> ModelResponse:
return ModelResponse(
id="chatcmpl-together-cache",
choices=[{"finish_reason": "stop", "index": 0, "message": {"content": "acknowledged", "role": "assistant"}}],
created=1756164000,
model=model,
object="chat.completion",
usage=Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens),
),
)
def test_completion_cost_prices_together_cached_tokens_at_cache_read_rate(_local_model_cost_map):
"""Regression: Together reports prompt_tokens_details.cached_tokens but no together_ai
registry entry carried cache_read_input_token_cost, so cache-hit tokens were priced at
0.0 and spend on cache-heavy workloads was understated."""
cost = completion_cost(
completion_response=_together_chat_response(
model="deepseek-ai/DeepSeek-V4-Flash-0731", prompt_tokens=7864, completion_tokens=16, cached_tokens=7863
),
custom_llm_provider="together_ai",
)
assert cost == pytest.approx(1 * 1.4e-07 + 7863 * 3e-08 + 16 * 2.8e-07, rel=1e-9)
def test_completion_cost_together_mapped_model_skips_size_bucket(_local_model_cost_map):
"""Regression: any together model whose name matches (\\d+b) was rewritten to a
together-ai-* size bucket before the registry lookup, so mapped models like
Muse-Glimmer-30B never used their per-model rates, cache fields included."""
cost = completion_cost(
completion_response=_together_chat_response(
model="meta-models/Muse-Glimmer-30B", prompt_tokens=63, completion_tokens=16, cached_tokens=0
),
custom_llm_provider="together_ai",
)
assert cost == pytest.approx(63 * 3.5e-07 + 16 * 1.5e-06, rel=1e-9)
def test_completion_cost_together_unmapped_model_still_uses_size_bucket(_local_model_cost_map):
cost = completion_cost(
completion_response=_together_chat_response(
model="qwen/Qwen2-72B-Instruct", prompt_tokens=23, completion_tokens=15, cached_tokens=0
),
custom_llm_provider="together_ai",
)
assert cost == pytest.approx((23 + 15) * 9e-07, rel=1e-9)
def test_completion_cost_together_metadata_only_model_still_uses_size_bucket(_local_model_cost_map):
assert "input_cost_per_token" not in litellm.model_cost["together_ai/togethercomputer/CodeLlama-34b-Instruct"]
cost = completion_cost(
completion_response=_together_chat_response(
model="togethercomputer/CodeLlama-34b-Instruct", prompt_tokens=23, completion_tokens=15, cached_tokens=0
),
custom_llm_provider="together_ai",
)
assert cost == pytest.approx((23 + 15) * 8e-07, rel=1e-9)
def test_select_model_name_strips_unregistered_alias_prefix(_local_model_cost_map):
"""A router-facing model_name alias containing "/" whose leading segment is NOT a
registered provider must not be double-prefixed into a non-existent cost key.

View file

@ -1,5 +1,6 @@
"""
Unit tests for DashScope image generation support (qwen-image-2.0, qwen-image-2.0-pro).
Unit tests for DashScope image generation support (qwen-image-2.0, qwen-image-2.0-pro,
qwen-image-3.0, qwen-image-3.0-pro).
Run in docker: pytest tests/test_litellm/test_dashscope_image_generation.py -v
"""
@ -30,6 +31,8 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException
[
"dashscope/qwen-image-2.0",
"dashscope/qwen-image-2.0-pro",
"dashscope/qwen-image-3.0",
"dashscope/qwen-image-3.0-pro",
],
)
def test_get_llm_provider_returns_dashscope(model_string: str):
@ -48,6 +51,8 @@ def test_get_llm_provider_returns_dashscope(model_string: str):
[
("dashscope/qwen-image-2.0", "dashscope"),
("dashscope/qwen-image-2.0-pro", "dashscope"),
("dashscope/qwen-image-3.0", "dashscope"),
("dashscope/qwen-image-3.0-pro", "dashscope"),
],
)
def test_get_model_info_mode_is_image_generation(
@ -93,6 +98,19 @@ class TestDashScopeImageGenerationConfig:
url = self.cfg.get_complete_url(custom, None, "qwen-image-2.0", {}, {})
assert url == custom
@pytest.mark.parametrize(
"chat_api_base",
[
"https://dashscope.aliyuncs.com/compatible-mode/v1",
"https://dashscope-intl.aliyuncs.com/compatible-mode/v1/",
],
)
def test_get_complete_url_ignores_chat_compatible_mode_base(
self, chat_api_base: str
):
url = self.cfg.get_complete_url(chat_api_base, None, "qwen-image-3.0", {}, {})
assert url == DEFAULT_API_BASE
def test_validate_environment_sets_auth_header(self):
headers = self.cfg.validate_environment(
headers={},
@ -135,6 +153,27 @@ class TestDashScopeImageGenerationConfig:
assert messages[0]["content"][0]["text"] == "a puppy on green grass"
assert req["parameters"]["size"] == "1024*1024"
@pytest.mark.parametrize("model", ["qwen-image-3.0", "qwen-image-3.0-pro"])
def test_transform_request_qwen_image_3(self, model: str):
req = self.cfg.transform_image_generation_request(
model=model,
prompt="a poster with small multilingual text",
optional_params=self.cfg.map_openai_params(
non_default_params={"size": "2048x2048", "n": 6},
optional_params={},
model=model,
drop_params=False,
),
litellm_params={},
headers={},
)
assert req["model"] == model
assert req["input"]["messages"][0]["content"][0]["text"] == (
"a poster with small multilingual text"
)
assert req["parameters"]["size"] == "2048*2048"
assert req["parameters"]["n"] == 6
def test_transform_request_empty_params(self):
req = self.cfg.transform_image_generation_request(
model="qwen-image-2.0-pro",
@ -238,6 +277,48 @@ class TestDashScopeImageGenerationConfig:
assert result.data[0].url == "https://example.com/img1.png"
assert result.data[1].url == "https://example.com/img2.png"
def test_transform_response_multiple_images_in_one_choice(self):
body = {
"output": {
"choices": [
{
"finish_reason": "stop",
"message": {
"role": "assistant",
"content": [
{"image": "https://example.com/img1.png", "type": "image"},
{"image": "https://example.com/img2.png", "type": "image"},
],
},
}
]
},
"usage": {
"output_width": 1024,
"output_height": 1024,
"output_image_count": 2,
},
}
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.headers = {}
mock_resp.json.return_value = body
result = self.cfg.transform_image_generation_response(
model="qwen-image-3.0",
raw_response=mock_resp,
model_response=ImageResponse(),
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert [image.url for image in result.data] == [
"https://example.com/img1.png",
"https://example.com/img2.png",
]
def test_transform_response_raises_on_non_200_status(self):
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 400
@ -294,14 +375,14 @@ class TestDashScopeImageGenerationConfig:
)
assert mapped["size"] == "1024*1024"
def test_map_openai_params_n_to_image_count(self):
def test_map_openai_params_n_passthrough(self):
mapped = self.cfg.map_openai_params(
non_default_params={"n": 2},
optional_params={},
model="qwen-image-2.0",
drop_params=False,
)
assert mapped["image_count"] == 2
assert mapped == {"n": 2}
def test_map_openai_params_unknown_size_uses_asterisk(self):
mapped = self.cfg.map_openai_params(
@ -338,7 +419,15 @@ class TestDashScopeImageGenerationConfig:
# ---------------------------------------------------------------------------
def test_litellm_image_generation_dashscope_end_to_end():
@pytest.mark.parametrize(
"model",
[
"dashscope/qwen-image-2.0",
"dashscope/qwen-image-3.0",
"dashscope/qwen-image-3.0-pro",
],
)
def test_litellm_image_generation_dashscope_end_to_end(model: str):
mock_response_body = {
"output": {
"choices": [
@ -374,7 +463,7 @@ def test_litellm_image_generation_dashscope_end_to_end():
mock_post.return_value = mock_http_response
response = litellm.image_generation(
model="dashscope/qwen-image-2.0",
model=model,
prompt="a puppy playing on green grass",
api_key="sk-test-key",
size="1024x1024",
@ -392,7 +481,7 @@ def test_litellm_image_generation_dashscope_end_to_end():
called_url = (
call_args[0][0] if call_args[0] else call_args.kwargs.get("url", "")
)
assert "dashscope" in called_url or "aliyuncs" in called_url
assert called_url == DEFAULT_API_BASE
# Verify request body contains DashScope format
call_kwargs = call_args[1] if call_args[1] else {}
@ -400,3 +489,4 @@ def test_litellm_image_generation_dashscope_end_to_end():
body = call_kwargs["json"]
assert "input" in body
assert "messages" in body["input"]
assert body["parameters"]["size"] == "1024*1024"

View file

@ -15,6 +15,7 @@ COST_MAP_ADAPTER: Final = TypeAdapter(CostMap)
SERVERLESS_CHAT_MODELS: Final = (
"together_ai/moonshotai/Kimi-K3",
"together_ai/zai-org/GLM-5.2",
"together_ai/zai-org/GLM-5.3-Flash",
"together_ai/deepseek-ai/DeepSeek-V4-Pro-0813",
"together_ai/deepseek-ai/DeepSeek-V4-Flash-0731",
"together_ai/MiniMaxAI/MiniMax-M3",
@ -110,6 +111,22 @@ def test_together_glm_52_pricing(cost_map: CostMap):
assert info["supports_reasoning"] is True
def test_together_glm_53_flash_pricing_and_capabilities(cost_map: CostMap):
info = cost_map["together_ai/zai-org/GLM-5.3-Flash"]
assert info["input_cost_per_token"] == 1.5e-07
assert info["output_cost_per_token"] == 5e-07
assert info["cache_read_input_token_cost"] == 3e-08
assert info["max_input_tokens"] == 1048575
assert info["max_output_tokens"] == 1048575
assert info["supports_function_calling"] is True
assert info["supports_parallel_function_calling"] is True
assert info["supports_prompt_caching"] is True
assert info["supports_tool_choice"] is True
assert info["supports_response_schema"] is True
assert info["supports_vision"] is True
assert info["supports_reasoning"] is True
def test_together_multilingual_e5_embedding_entry(cost_map: CostMap):
info = cost_map["together_ai/intfloat/multilingual-e5-large-instruct"]
assert info["mode"] == "embedding"
@ -159,3 +176,51 @@ def test_together_backup_cost_map_in_sync(cost_map: CostMap):
together_main = {k: v for k, v in cost_map.items() if k.startswith("together_ai/")}
together_backup = {k: v for k, v in backup.items() if k.startswith("together_ai/")}
assert together_backup == together_main
CACHED_INPUT_MODELS: Final = (
"together_ai/moonshotai/Kimi-K3",
"together_ai/zai-org/GLM-5.2",
"together_ai/meta-models/Muse-Glimmer-30B",
"together_ai/Qwen/Qwen3.8-2.4T-A95B",
"together_ai/deepseek-ai/DeepSeek-V4-Pro-0813",
"together_ai/deepseek-ai/DeepSeek-V4-Flash-0731",
"together_ai/thinkingmachines/Inkling",
"together_ai/MiniMaxAI/MiniMax-M3",
"together_ai/thinkingmachines/Inkling-Small",
"together_ai/moonshotai/Kimi-K2.7-Code",
"together_ai/deepseek-ai/DeepSeek-V4-Pro",
"together_ai/nvidia/nemotron-3-ultra-550b-a55b",
"together_ai/Qwen/Qwen3.7-Max",
)
@pytest.mark.parametrize("model", CACHED_INPUT_MODELS)
def test_together_cached_input_model_carries_cache_read_pricing(cost_map: CostMap, model: str):
info = cost_map.get(model)
assert info is not None, f"{model} missing from model_prices_and_context_window.json"
assert info.get("supports_prompt_caching") is True
cache_read = info.get("cache_read_input_token_cost")
assert isinstance(cache_read, float)
assert 0 < cache_read < info["input_cost_per_token"]
assert "cache_creation_input_token_cost" not in info
def test_together_prompt_caching_flag_implies_cache_read_rate(cost_map: CostMap):
for model, info in cost_map.items():
if model.startswith("together_ai/") and info.get("supports_prompt_caching"):
assert "cache_read_input_token_cost" in info, f"{model} flags caching without a cache read rate"
def test_together_deepseek_v4_flash_cache_read_rate(cost_map: CostMap):
info = cost_map["together_ai/deepseek-ai/DeepSeek-V4-Flash-0731"]
assert info["input_cost_per_token"] == 1.4e-07
assert info["cache_read_input_token_cost"] == 3e-08
assert info["output_cost_per_token"] == 2.8e-07
def test_together_qwen_37_max_repriced_to_current_together_rate(cost_map: CostMap):
info = cost_map["together_ai/Qwen/Qwen3.7-Max"]
assert info["input_cost_per_token"] == 2.5e-06
assert info["output_cost_per_token"] == 7.5e-06
assert info["cache_read_input_token_cost"] == 5e-07

View file

@ -1014,6 +1014,10 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"supports_none_reasoning_effort": {"type": "boolean"},
"supports_xhigh_reasoning_effort": {"type": "boolean"},
"supports_max_reasoning_effort": {"type": "boolean"},
"reasoning_effort_levels": {
"type": "array",
"items": {"type": "string", "enum": ["none", "minimal", "low", "medium", "high", "xhigh", "max"]},
},
"supports_adaptive_thinking": {"type": "boolean"},
"supports_legacy_thinking": {"type": "boolean"},
"thinking_always_on": {"type": "boolean"},

View file

@ -0,0 +1,87 @@
"""Tests for the shared MCP header primitives.
``same_header`` / ``has_header`` / ``without_header`` are the one owner of "is this the credential's
header", used by both MCP stacks and the upstream-credential resolver. They live here rather than in
either stack because a second implementation is exactly how an injected header came to shadow a
resolved credential on one path and not the other.
"""
import pytest
from litellm.types.mcp import (
credential_redirect_hook,
crosses_origin,
has_header,
same_header,
without_header,
)
@pytest.mark.parametrize(
"a,b,expected",
[
("Authorization", "authorization", True),
("ESB-OAuth", "esb-oauth", True),
("esb-oauth", "esb-oauth", True),
("esb-oauth", "esb_oauth", False),
("esb-oauth", "Authorization", False),
],
)
def test_header_names_compare_case_insensitively(a: str, b: str, expected: bool) -> None:
# RFC 7230 3.2. Every consumer of a credential slot routes through this, so a case-sensitive
# comparison anywhere would let an injected header shadow a resolved credential.
assert same_header(a, b) is expected
def test_without_header_drops_every_casing_and_keeps_the_rest() -> None:
headers = {"ESB-OAuth": "injected", "esb-oauth": "also injected", "X-Trace": "keep"}
assert without_header(headers, "esb-oauth") == {"X-Trace": "keep"}
def test_without_header_collapses_to_none_when_nothing_remains() -> None:
assert without_header({"Authorization": "Bearer x"}, "AUTHORIZATION") is None
assert without_header(None, "esb-oauth") is None
assert without_header({}, "esb-oauth") is None
def test_has_header_matches_any_casing() -> None:
assert has_header({"ESB-OAuth": "v"}, "esb-oauth") is True
assert has_header({"X-Other": "v"}, "esb-oauth") is False
assert has_header(None, "esb-oauth") is False
@pytest.mark.parametrize(
"target,expected",
[
("https://upstream.example.com/other", False), # same origin
("https://upstream.example.com:443/other", False), # explicit default port
("https://attacker.example.com/collect", True), # different host
("http://upstream.example.com/collect", True), # scheme downgrade, same host
("https://upstream.example.com:8443/other", True), # different port, same host
("https://sub.upstream.example.com/x", True), # different host
],
)
def test_origin_is_scheme_host_and_port_not_host_alone(target: str, expected: bool) -> None:
assert crosses_origin("https://upstream.example.com/mcp", target) is expected
def test_an_https_upgrade_of_the_same_host_is_not_crossing() -> None:
# HTTP clients exempt this when deciding to keep Authorization, so a credential slot that did
# not would lose the credential on every such redirect.
assert crosses_origin("http://upstream.example.com/mcp", "https://upstream.example.com/x") is False
assert crosses_origin("http://upstream.example.com/mcp", "http://upstream.example.com/x") is False
@pytest.mark.asyncio
async def test_the_hook_drops_the_slot_only_once_the_origin_changes() -> None:
import httpx
hook = credential_redirect_hook("https://upstream.example.com/mcp", "esb-oauth")
same = httpx.Request("GET", "https://upstream.example.com/other", headers={"esb-oauth": "Bearer x"})
await hook(same)
assert same.headers["esb-oauth"] == "Bearer x"
foreign = httpx.Request("GET", "https://attacker.example.com/x", headers={"esb-oauth": "Bearer x"})
await hook(foreign)
assert "esb-oauth" not in foreign.headers

View file

@ -3,7 +3,7 @@
"limit": 22733
},
"LIT002": {
"limit": 26863
"limit": 26860
},
"LIT003": {
"limit": 269
@ -33,6 +33,6 @@
"limit": 5583
},
"LIT012": {
"limit": 4510
"limit": 4509
}
}

View file

@ -3,6 +3,7 @@ import React from "react";
import { SimpleTooltip } from "@/components/ui/tooltip";
import { MountedFormField } from "@/components/common_components/MountedFormField";
import UpstreamTokenHeaderField from "./UpstreamTokenHeaderField";
import { requiredRule } from "@/components/common_components/formRules";
import { MultiSelect } from "@/components/shared/MultiSelect";
import { PasswordInput } from "@/components/shared/PasswordInput";
@ -205,6 +206,7 @@ const IdJagFormFields: React.FC<IdJagFormFieldsProps> = ({ isEditing = false })
>
{(control) => <MultiSelect {...tagsControl(control)} placeholder="Add scopes" className="rounded-lg" />}
</MountedFormField>
<UpstreamTokenHeaderField />
</>
);
};

View file

@ -286,4 +286,42 @@ describe("OAuthFormFields", () => {
});
});
});
describe("token header field", () => {
it("renders on the M2M flow", () => {
render(
<WithForm>
<OAuthFormFields isM2M={true} />
</WithForm>,
);
expect(screen.getByPlaceholderText("Authorization")).toBeInTheDocument();
});
it("renders on the interactive flow", () => {
render(
<WithForm>
<OAuthFormFields isM2M={false} />
</WithForm>,
);
expect(screen.getByPlaceholderText("Authorization")).toBeInTheDocument();
});
it("submits its value under credentials.upstream_token_header", async () => {
const onFinish = vi.fn();
render(
<WithForm onFinish={onFinish}>
<OAuthFormFields isM2M={true} />
</WithForm>,
);
fireEvent.change(screen.getByPlaceholderText("Authorization"), { target: { value: "esb-oauth" } });
fireEvent.click(screen.getByText("Submit"));
await waitFor(() => {
expect(onFinish).toHaveBeenCalledWith(
expect.objectContaining({
credentials: expect.objectContaining({ upstream_token_header: "esb-oauth" }),
}),
);
});
});
});
});

View file

@ -11,6 +11,7 @@ import { OAUTH_FLOW } from "@/components/mcp_tools/types";
import { MountedFormField } from "@/components/common_components/MountedFormField";
import { requiredRule } from "@/components/common_components/formRules";
import TokenEndpointAuthMethodField from "./TokenEndpointAuthMethodField";
import UpstreamTokenHeaderField from "./UpstreamTokenHeaderField";
import {
numberControl,
parsesAsJson,
@ -175,6 +176,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
{(control) => <MultiSelect {...tagsControl(control)} placeholder="Add scopes" className="rounded-lg" />}
</MountedFormField>
<UpstreamResourceField />
<UpstreamTokenHeaderField />
</>
) : (
<>
@ -237,6 +239,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
{(control) => <MultiSelect {...tagsControl(control)} placeholder="Add scopes" className="rounded-lg" />}
</MountedFormField>
<UpstreamResourceField />
<UpstreamTokenHeaderField />
<MountedFormField
label={
<FieldLabel

View file

@ -6,6 +6,7 @@ import { SimpleTooltip } from "@/components/ui/tooltip";
import { useWatch } from "react-hook-form";
import { MountedFormField } from "@/components/common_components/MountedFormField";
import UpstreamTokenHeaderField from "./UpstreamTokenHeaderField";
import { requiredRule } from "@/components/common_components/formRules";
import { PasswordInput } from "@/components/shared/PasswordInput";
import { Input } from "@/components/ui/input";
@ -184,6 +185,7 @@ const TokenExchangeFormFields: React.FC<TokenExchangeFormFieldsProps> = ({ isEdi
/>
)}
</MountedFormField>
<UpstreamTokenHeaderField />
</>
);
};

View file

@ -0,0 +1,31 @@
import { Info } from "lucide-react";
import React from "react";
import { SimpleTooltip } from "@/components/ui/tooltip";
import { Input } from "@/components/ui/input";
import { MountedFormField } from "@/components/common_components/MountedFormField";
import { textControl } from "./mcpFieldRules";
const UpstreamTokenHeaderField: React.FC = () => (
<MountedFormField
label={
<span className="text-sm font-medium text-foreground flex items-center">
Token Header (optional)
<SimpleTooltip content="Which upstream header carries the token LiteLLM resolves for this server. Leave blank to send it as 'Authorization: Bearer <token>', which is the default and what most servers expect. Set a header name when the upstream expects it elsewhere, for example an API gateway that terminates its own credential on 'esb-oauth' while a separate Authorization from Static Headers passes through to the server behind it.">
<Info className="ml-2 size-4 text-info hover:text-info/80 cursor-help" />
</SimpleTooltip>
</span>
}
name={["credentials", "upstream_token_header"]}
>
{(control) => (
<Input
{...textControl(control)}
placeholder="Authorization"
className="rounded-lg border-border focus:border-info focus:ring-ring"
/>
)}
</MountedFormField>
);
export default UpstreamTokenHeaderField;

View file

@ -255,13 +255,28 @@ export const CASES: readonly DifferentialCase[] = [
},
// --- credentials filtering ---
// ADMIN_CONFIG_CREDENTIAL_KEYS is exactly ["upstream_resource"], so only that key
// takes the blank-to-explicit-null branch. A blank client_id is dropped instead.
// Only a key in ADMIN_CONFIG_CREDENTIAL_KEYS takes the blank-to-explicit-null branch, which is
// what makes it clearable: the backend merge preserves an omitted key forever. A blank client_id
// is dropped instead.
{
label: "blank upstream_resource becomes an explicit null",
values: { ...ROOT, auth_type: "oauth2", credentials: { upstream_resource: "", client_secret: "keep" } },
ui: {},
},
{
label: "blank upstream_token_header becomes an explicit null",
values: { ...ROOT, auth_type: "oauth2", credentials: { upstream_token_header: "", client_secret: "keep" } },
ui: {},
},
{
label: "a set upstream_token_header rides the credentials blob",
values: {
...ROOT,
auth_type: "oauth2",
credentials: { upstream_token_header: "esb-oauth", client_secret: "keep" },
},
ui: {},
},
{
label: "blank non-admin credential is dropped, not nulled",
values: { ...ROOT, auth_type: "oauth2", credentials: { client_id: "", client_secret: "keep", scopes: [] } },

View file

@ -262,7 +262,14 @@ describe("edit root: exact mounted set per auth configuration", () => {
...PERMS,
"delegate_auth_to_upstream",
],
credentials: ["client_id", "client_secret", "token_endpoint_auth_method", "scopes", "upstream_resource"],
credentials: [
"client_id",
"client_secret",
"token_endpoint_auth_method",
"scopes",
"upstream_resource",
"upstream_token_header",
],
},
);
});
@ -286,7 +293,14 @@ describe("edit root: exact mounted set per auth configuration", () => {
...PERMS,
"delegate_auth_to_upstream",
],
credentials: ["client_id", "client_secret", "scopes", "upstream_resource", "token_endpoint_auth_method"],
credentials: [
"client_id",
"client_secret",
"scopes",
"upstream_resource",
"token_endpoint_auth_method",
"upstream_token_header",
],
},
);
});
@ -306,7 +320,7 @@ describe("edit root: exact mounted set per auth configuration", () => {
"env_vars",
...PERMS,
],
credentials: ["client_id", "client_secret", "scopes"],
credentials: ["client_id", "client_secret", "scopes", "upstream_token_header"],
},
);
});
@ -324,7 +338,7 @@ describe("edit root: exact mounted set per auth configuration", () => {
"env_vars",
...PERMS,
],
credentials: ["client_id", "client_secret", "scopes"],
credentials: ["client_id", "client_secret", "scopes", "upstream_token_header"],
},
);
});
@ -344,6 +358,7 @@ describe("edit root: exact mounted set per auth configuration", () => {
...PERMS,
],
credentials: [
"upstream_token_header",
"id_jag_resource_token_endpoint",
"client_id",
"client_secret",
@ -434,7 +449,14 @@ describe("create root: exact mounted set per configuration", () => {
...PERMS,
"delegate_auth_to_upstream",
],
credentials: ["client_id", "client_secret", "scopes", "upstream_resource", "token_endpoint_auth_method"],
credentials: [
"client_id",
"client_secret",
"scopes",
"upstream_resource",
"token_endpoint_auth_method",
"upstream_token_header",
],
},
);
});

View file

@ -24,6 +24,7 @@ const OAUTH_M2M_CREDENTIALS = [
"token_endpoint_auth_method",
"scopes",
"upstream_resource",
"upstream_token_header",
] as const;
const OAUTH_INTERACTIVE_CREDENTIALS = [
@ -32,6 +33,7 @@ const OAUTH_INTERACTIVE_CREDENTIALS = [
"scopes",
"upstream_resource",
"token_endpoint_auth_method",
"upstream_token_header",
] as const;
const OAUTH_INTERACTIVE_ROOT = [
@ -44,6 +46,7 @@ const OAUTH_INTERACTIVE_ROOT = [
] as const;
const ID_JAG_CREDENTIALS = [
"upstream_token_header",
"id_jag_resource_token_endpoint",
"client_id",
"client_secret",
@ -100,7 +103,7 @@ const authSubtreeCredentials = ({ authType, oauthFlowType }: AuthSubtreeGates):
];
}
if (authType === AUTH_TYPE.OAUTH2_TOKEN_EXCHANGE) {
return [...authValue, "client_id", "client_secret", "scopes"];
return [...authValue, "client_id", "client_secret", "scopes", "upstream_token_header"];
}
if (authType === AUTH_TYPE.OAUTH2_ID_JAG) {
return [...authValue, ...ID_JAG_CREDENTIALS];

View file

@ -1,5 +1,6 @@
import { render, screen } from "@testing-library/react";
import { render, screen, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { useState } from "react";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { useInfiniteTeams } from "@/app/(dashboard)/hooks/teams/useTeams";
import TeamMultiSelect from "./team_multi_select";
@ -58,9 +59,9 @@ describe("TeamMultiSelect", () => {
await user.click(combobox());
expect(screen.getByText("Alpha Team")).toBeInTheDocument();
expect(screen.getByText("(team-1)")).toBeInTheDocument();
expect(screen.getByText("team-1")).toBeInTheDocument();
expect(screen.getByText("Beta Team")).toBeInTheDocument();
expect(screen.getByText("(team-2)")).toBeInTheDocument();
expect(screen.getByText("team-2")).toBeInTheDocument();
});
it("deduplicates a team that appears on more than one page", async () => {
@ -116,6 +117,29 @@ describe("TeamMultiSelect", () => {
expect(screen.getByText(/No teams found/)).toBeInTheDocument();
});
it("keeps a picked team's alias on its chip once a later search drops it from the loaded page", async () => {
const user = userEvent.setup();
function Controlled() {
const [value, setValue] = useState<string[]>([]);
return <TeamMultiSelect value={value} onChange={setValue} />;
}
const { rerender } = render(<Controlled />);
await user.click(combobox());
const matches = screen.getAllByText("Beta Team");
await user.click(matches[matches.length - 1]);
mockUseInfiniteTeams.mockReturnValue(
mockTeamsResult({ pages: [{ teams: [team("team-3", "Gamma Team")] }] }) as never,
);
rerender(<Controlled />);
const chips = document.querySelector('[data-slot="combobox-chips"]') as HTMLElement;
expect(within(chips).getByText("Beta Team")).toBeInTheDocument();
expect(within(chips).queryByText("team-2")).not.toBeInTheDocument();
});
it("passes the page size and organization filter through to the teams query", () => {
render(<TeamMultiSelect pageSize={25} organizationId="org-7" />);

View file

@ -1,22 +1,7 @@
import React, { useMemo, useState, type UIEvent } from "react";
import { Loader2 } from "lucide-react";
import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer";
import {
Combobox,
ComboboxChip,
ComboboxChips,
ComboboxChipsInput,
ComboboxClear,
ComboboxContent,
ComboboxEmpty,
ComboboxItem,
ComboboxList,
ComboboxValue,
useComboboxAnchor,
} from "@/components/ui/combobox";
import React, { useMemo, useState } from "react";
import { PaginatedMultiSelect } from "@/components/shared/PaginatedMultiSelect";
import type { SearchSelectOption } from "@/components/shared/SearchSelect";
import { useInfiniteTeams } from "@/app/(dashboard)/hooks/teams/useTeams";
import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants";
import { Team } from "../key_team_helpers/key_list";
interface TeamMultiSelectProps {
value?: string[];
@ -27,8 +12,6 @@ interface TeamMultiSelectProps {
placeholder?: string;
}
const SCROLL_THRESHOLD = 0.8;
const TeamMultiSelect: React.FC<TeamMultiSelectProps> = ({
value = [],
onChange,
@ -37,9 +20,7 @@ const TeamMultiSelect: React.FC<TeamMultiSelectProps> = ({
pageSize = 20,
placeholder = "Search teams by alias...",
}) => {
const anchor = useComboboxAnchor();
const [search, setSearch] = useState("");
const debouncedSetSearch = useDebouncedCallback(setSearch, { wait: DEBOUNCE_WAIT_MS });
const { data, fetchNextPage, hasNextPage, isFetchingNextPage, isLoading } = useInfiniteTeams(
pageSize,
@ -47,68 +28,40 @@ const TeamMultiSelect: React.FC<TeamMultiSelectProps> = ({
organizationId,
);
const teamById = useMemo(
const options = useMemo<SearchSelectOption[]>(
() =>
new Map<string, Team>(
(data?.pages ?? []).flatMap((page) => page.teams).map((team) => [team.team_id, team] as const),
Array.from(
new Map(
(data?.pages ?? [])
.flatMap((page) => page.teams)
.map(
(team) =>
[
team.team_id,
{ label: team.team_alias || team.team_id, value: team.team_id, sublabel: team.team_id },
] as const,
),
).values(),
),
[data],
);
const teamIds = useMemo(() => Array.from(teamById.keys()), [teamById]);
const aliasOf = (teamId: string) => teamById.get(teamId)?.team_alias ?? teamId;
const handleScroll = (event: UIEvent<HTMLDivElement>) => {
const target = event.currentTarget;
if (target.scrollHeight === 0) return;
const scrollRatio = (target.scrollTop + target.clientHeight) / target.scrollHeight;
if (scrollRatio >= SCROLL_THRESHOLD && hasNextPage && !isFetchingNextPage) {
fetchNextPage();
}
};
return (
<Combobox
multiple
items={teamIds}
<PaginatedMultiSelect
options={options}
value={value}
onValueChange={(next: string[]) => onChange?.(next)}
filter={null}
onInputValueChange={debouncedSetSearch}
onSearchChange={setSearch}
onLoadMore={fetchNextPage}
hasNextPage={hasNextPage}
isLoading={isLoading}
isFetchingNextPage={isFetchingNextPage}
placeholder={placeholder}
emptyText="No teams found"
loadingText="Loading teams..."
clearAllLabel="Clear all teams"
disabled={disabled}
>
<ComboboxChips render={<div ref={anchor} />} className="w-full" aria-busy={isLoading}>
<ComboboxValue>
{(selected: string[]) =>
selected.map((teamId) => (
<ComboboxChip key={teamId} aria-label={aliasOf(teamId)}>
{aliasOf(teamId)}
</ComboboxChip>
))
}
</ComboboxValue>
<ComboboxChipsInput placeholder={placeholder} aria-label={placeholder} disabled={disabled} />
{value.length > 0 && <ComboboxClear aria-label="Clear all teams" disabled={disabled} />}
</ComboboxChips>
<ComboboxContent anchor={anchor}>
<ComboboxEmpty>
{isLoading ? <Loader2 className="size-4 animate-spin text-muted-foreground" /> : "No teams found"}
</ComboboxEmpty>
<ComboboxList onScroll={handleScroll}>
{(teamId: string) => (
<ComboboxItem key={teamId} value={teamId}>
<span className="font-medium">{aliasOf(teamId)}</span>{" "}
<span className="text-muted-foreground">({teamId})</span>
</ComboboxItem>
)}
</ComboboxList>
{isFetchingNextPage && (
<div className="flex justify-center py-2">
<Loader2 className="size-4 animate-spin text-muted-foreground" />
</div>
)}
</ComboboxContent>
</Combobox>
/>
);
};

View file

@ -199,6 +199,52 @@ describe("UserSearchModal submit payload", () => {
});
});
describe("UserSearchModal search lifecycle", () => {
const directory = [
{ user_id: "u-jones", user_email: "alice.jones@example.com" },
{ user_id: "u-smith", user_email: "alice.smith@example.com" },
{ user_id: "u-bob", user_email: "bob@example.com" },
];
beforeEach(() => {
vi.mocked(userFilterUICall).mockReset();
vi.mocked(userFilterUICall).mockImplementation((_accessToken, params) => {
const query = params.get("user_email") ?? "";
return Promise.resolve(directory.filter((user) => user.user_email.includes(query))) as never;
});
});
const searchedFor = (): string[] =>
vi.mocked(userFilterUICall).mock.calls.map((call) => {
const email = call[1].get("user_email");
return email === null ? `user_id=${call[1].get("user_id")}` : `user_email=${email}`;
});
const settleDebounce = () =>
act(async () => {
await new Promise((resolve) => setTimeout(resolve, DEBOUNCE_WAIT_MS + 100));
});
it("leaves the search unfiltered after a pick, so reopening searches the newly typed text", async () => {
const user = userEvent.setup();
render(<UserSearchModal isVisible onCancel={vi.fn()} onSubmit={vi.fn()} accessToken="sk-test" />);
const input = getEmailSearchInput();
await user.click(input);
await user.type(input, "ali");
await user.click(await screen.findByRole("option", { name: "alice.jones@example.com" }));
await settleDebounce();
expect(searchedFor()).toEqual(["user_email=ali"]);
await user.click(input);
await user.type(input, "bob");
expect(await screen.findByRole("option", { name: "bob@example.com" })).toBeInTheDocument();
expect(searchedFor()).toEqual(["user_email=ali", "user_email=bob"]);
});
});
describe("UserSearchModal out-of-order search results", () => {
const answers = new Map<string, (users: { user_id: string; user_email: string }[]) => void>();

Some files were not shown because too many files have changed in this diff Show more