Merge branch 'litellm_internal_staging' into litellm_oss_staging_04_21_2026_2

This commit is contained in:
mateo-berri 2026-05-21 00:32:59 +00:00
commit d87006881c
No known key found for this signature in database
107 changed files with 5177 additions and 641 deletions

View file

@ -158,6 +158,8 @@ jobs:
CHOCOLATEY_CONFIRM_ALL: "true"
- run:
name: Install Dependencies
environment:
UV_HTTP_TIMEOUT: "300"
command: |
$installer = Join-Path $env:TEMP "uv-install.ps1"
Invoke-WebRequest -Uri https://astral.sh/uv/0.10.9/install.ps1 -OutFile $installer

View file

@ -225,6 +225,10 @@ use_chat_completions_url_for_anthropic_messages: bool = bool(
route_all_chat_openai_to_responses: bool = (
os.getenv("LITELLM_ROUTE_ALL_CHAT_OPENAI_TO_RESPONSES", "false").lower() == "true"
) # When True, routes all OpenAI /chat/completions requests through the Responses API bridge
use_legacy_interactions_schema: bool = (
os.getenv("LITELLM_USE_LEGACY_INTERACTIONS_SCHEMA", "false").lower() == "true"
) # When True, sends Api-Revision: 2026-05-07 to Google so responses use the legacy `outputs`
# schema instead of the new `steps` schema. Remove this flag after June 8, 2026.
retry = True
### AUTH ###
api_key: Optional[str] = None

View file

@ -2,7 +2,17 @@
Streaming iterator for transforming Responses API stream to Interactions API stream.
"""
from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, cast
from collections import deque
from typing import (
Any,
AsyncIterator,
Deque,
Dict,
Iterator,
List,
Optional,
cast,
)
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
@ -15,7 +25,6 @@ from litellm.types.interactions import (
InteractionsAPIStreamingResponse,
)
from litellm.types.llms.openai import (
ContentPartAddedEvent,
OutputTextDeltaEvent,
ResponseCompletedEvent,
ResponseCreatedEvent,
@ -30,7 +39,13 @@ class LiteLLMResponsesInteractionsStreamingIterator:
This class handles both sync and async iteration, transforming Responses API
streaming events (output.text.delta, response.completed, etc.) to Interactions
API streaming events (content.delta, interaction.complete, etc.).
API streaming events.
Schema selection:
- New schema (default, use_legacy_interactions_schema=False):
interaction.created -> step.start -> step.delta ... -> step.stop -> interaction.completed
- Legacy schema (use_legacy_interactions_schema=True, remove after June 8 2026):
interaction.start -> content.start -> content.delta ... -> content.stop -> interaction.complete
"""
def __init__(
@ -42,6 +57,8 @@ class LiteLLMResponsesInteractionsStreamingIterator:
custom_llm_provider: Optional[str] = None,
litellm_metadata: Optional[Dict[str, Any]] = None,
):
import litellm
self.model = model
self.responses_stream_iterator = litellm_custom_stream_wrapper
self.request_input = request_input
@ -52,100 +69,156 @@ class LiteLLMResponsesInteractionsStreamingIterator:
self.collected_text = ""
self.sent_interaction_start = False
self.sent_content_start = False
self._pending_events: List[InteractionsAPIStreamingResponse] = []
# Capture the schema flag once at construction time so all events
# emitted by this stream use a consistent schema, even if the global
# flag is mutated mid-stream (e.g. by a config reload).
self._use_legacy: bool = litellm.use_legacy_interactions_schema
# Buffer of events that have been derived from upstream chunks but not
# yet returned to the caller. A single Responses API chunk may expand
# into multiple Interactions API events (e.g. the first text delta
# produces interaction.created + step.start + step.delta), and the
# terminal sequence on stream end may also span multiple events
# (step.stop + interaction.completed).
self._pending_events: Deque[InteractionsAPIStreamingResponse] = deque()
# Tracks whether we've already emitted a terminal completion event so
# the StopIteration fallback path doesn't double-emit.
self._sent_completion_event = False
# ID resolved from the first upstream chunk (item_id on a text delta or
# response.id on response.created). Persisted so the EOF terminal
# events stay correlated with the start events delivered earlier.
self._interaction_id: Optional[str] = None
def _transform_responses_chunk_to_interactions_chunk(
self,
responses_chunk: ResponsesAPIStreamingResponse,
) -> Optional[InteractionsAPIStreamingResponse]:
# ------------------------------------------------------------------
# Event builders
# ------------------------------------------------------------------
def _build_interaction_start_event(
self, interaction_id: str
) -> InteractionsAPIStreamingResponse:
event_type = "interaction.start" if self._use_legacy else "interaction.created"
return InteractionsAPIStreamingResponse(
event_type=event_type,
id=interaction_id,
object="interaction",
status="in_progress",
model=self.model,
)
def _build_content_start_event(
self, interaction_id: str
) -> InteractionsAPIStreamingResponse:
if self._use_legacy:
return InteractionsAPIStreamingResponse(
event_type="content.start",
id=interaction_id,
object="content",
delta={"type": "text", "text": ""},
)
return InteractionsAPIStreamingResponse(
event_type="step.start",
index=0,
step={"type": "model_output", "content": []},
)
def _build_text_delta_event(
self, interaction_id: str, delta_text: str
) -> InteractionsAPIStreamingResponse:
if self._use_legacy:
return InteractionsAPIStreamingResponse(
event_type="content.delta",
id=interaction_id,
object="content",
delta={"type": "text", "text": delta_text},
)
return InteractionsAPIStreamingResponse(
event_type="step.delta",
index=0,
delta={"type": "text", "text": delta_text},
)
def _build_content_stop_event(
self, interaction_id: Optional[str]
) -> InteractionsAPIStreamingResponse:
if self._use_legacy:
return InteractionsAPIStreamingResponse(
event_type="content.stop",
id=interaction_id,
object="content",
delta={"type": "text", "text": self.collected_text},
)
return InteractionsAPIStreamingResponse(
event_type="step.stop",
index=0,
)
def _build_completion_event(
self, response_id: str
) -> InteractionsAPIStreamingResponse:
if self._use_legacy:
return InteractionsAPIStreamingResponse(
event_type="interaction.complete",
id=response_id,
object="interaction",
status="completed",
model=self.model,
outputs=[{"type": "text", "text": self.collected_text}],
)
return InteractionsAPIStreamingResponse(
event_type="interaction.completed",
id=response_id,
object="interaction",
status="completed",
model=self.model,
steps=[
{
"type": "model_output",
"content": [{"type": "text", "text": self.collected_text}],
}
],
)
# ------------------------------------------------------------------
# Per-chunk transform (returns a list of events to enqueue)
# ------------------------------------------------------------------
def _events_for_chunk(
self, responses_chunk: ResponsesAPIStreamingResponse
) -> List[InteractionsAPIStreamingResponse]:
"""
Transform a Responses API streaming chunk to an Interactions API streaming chunk.
Translate a single upstream Responses API chunk into the list of
Interactions API events it should produce.
Responses API events:
- output.text.delta -> content.delta
- response.completed -> interaction.complete
Interactions API events:
- interaction.start
- content.start
- content.delta
- content.stop
- interaction.complete
Returning a list (rather than a single event) lets a chunk emit any
synthetic start events that haven't been sent yet *together with* the
actual delta event, so we never silently drop the chunk's payload.
"""
if not responses_chunk:
return None
return []
# Handle OutputTextDeltaEvent -> content.delta
# Text delta: emit any missing start events, then the delta itself.
if isinstance(responses_chunk, OutputTextDeltaEvent):
delta_text = (
responses_chunk.delta if isinstance(responses_chunk.delta, str) else ""
)
self.collected_text += delta_text
# Fallback: emit interaction.start, and queue content.start carrying this
# delta so the first token is preserved in the stream.
if not self.sent_interaction_start:
self.sent_interaction_start = True
self.sent_content_start = True
self._pending_events.append(
InteractionsAPIStreamingResponse(
event_type="content.start",
id=getattr(responses_chunk, "item_id", None),
object="content",
delta={"type": "text", "text": delta_text},
)
)
return InteractionsAPIStreamingResponse(
event_type="interaction.start",
id=getattr(responses_chunk, "item_id", None)
or f"interaction_{id(self)}",
object="interaction",
status="in_progress",
model=self.model,
)
# Fallback: emit content.start if ContentPartAddedEvent never arrived
if not self.sent_content_start:
self.sent_content_start = True
return InteractionsAPIStreamingResponse(
event_type="content.start",
id=getattr(responses_chunk, "item_id", None),
object="content",
delta={"type": "text", "text": delta_text},
)
# Normal path: emit content.delta with type field
return InteractionsAPIStreamingResponse(
event_type="content.delta",
id=getattr(responses_chunk, "item_id", None),
object="content",
delta={"type": "text", "text": delta_text},
interaction_id = (
getattr(responses_chunk, "item_id", None) or f"interaction_{id(self)}"
)
if self._interaction_id is None:
self._interaction_id = interaction_id
# Handle ContentPartAddedEvent -> content.start (arrives before text deltas)
if isinstance(responses_chunk, ContentPartAddedEvent):
# Fallback: emit interaction.start if ResponseCreatedEvent never arrived
events: List[InteractionsAPIStreamingResponse] = []
if not self.sent_interaction_start:
self.sent_interaction_start = True
return InteractionsAPIStreamingResponse(
event_type="interaction.start",
id=getattr(responses_chunk, "item_id", None)
or f"interaction_{id(self)}",
object="interaction",
status="in_progress",
model=self.model,
)
events.append(self._build_interaction_start_event(interaction_id))
if not self.sent_content_start:
self.sent_content_start = True
return InteractionsAPIStreamingResponse(
event_type="content.start",
id=getattr(responses_chunk, "item_id", None),
object="content",
delta={"type": "text", "text": ""},
)
return None
events.append(self._build_content_start_event(interaction_id))
events.append(self._build_text_delta_event(interaction_id, delta_text))
return events
# Handle ResponseCreatedEvent or ResponseInProgressEvent -> interaction.start
# Response created / in-progress: synthesize interaction start if we
# haven't already sent one.
if isinstance(responses_chunk, (ResponseCreatedEvent, ResponseInProgressEvent)):
if not self.sent_interaction_start:
self.sent_interaction_start = True
@ -153,177 +226,136 @@ class LiteLLMResponsesInteractionsStreamingIterator:
getattr(responses_chunk.response, "id", None)
if hasattr(responses_chunk, "response")
else None
)
return InteractionsAPIStreamingResponse(
event_type="interaction.start",
id=response_id or f"interaction_{id(self)}",
object="interaction",
status="in_progress",
model=self.model,
)
) or f"interaction_{id(self)}"
if self._interaction_id is None:
self._interaction_id = response_id
return [self._build_interaction_start_event(response_id)]
return []
# Handle ResponseCompletedEvent -> interaction.complete
# Response completed: emit step.stop (if content was started) followed
# by the terminal completion event. Prefer the interaction id already
# established by earlier events so consumers can correlate the start
# and completion events by id (response.id may differ from the item_id
# used to derive the initial id when the stream starts directly with a
# text delta).
if isinstance(responses_chunk, ResponseCompletedEvent):
self.finished = True
response = responses_chunk.response
# Send content.stop first if content was started
if self.sent_content_start:
# Note: We'll send this in the iterator, not here
pass
# Send interaction.complete
return InteractionsAPIStreamingResponse(
event_type="interaction.complete",
id=getattr(response, "id", None) or f"interaction_{id(self)}",
object="interaction",
status="completed",
model=self.model,
outputs=[
{
"type": "text",
"text": self.collected_text,
}
],
response_id = (
self._interaction_id
or getattr(response, "id", None)
or f"interaction_{id(self)}"
)
# For other event types, return None (skip)
return None
terminal: List[InteractionsAPIStreamingResponse] = []
if self.sent_content_start:
terminal.append(self._build_content_stop_event(response_id))
terminal.append(self._build_completion_event(response_id))
self._sent_completion_event = True
return terminal
return []
def _build_terminal_events_on_eof(
self,
) -> List[InteractionsAPIStreamingResponse]:
"""
Build the events to flush when the upstream stream ends without a
ResponseCompletedEvent. Ensures consumers always observe a terminal
interaction.completed/interaction.complete carrying the full text.
"""
if self._sent_completion_event:
return []
fallback_id = self._interaction_id or f"interaction_{id(self)}"
terminal: List[InteractionsAPIStreamingResponse] = []
if self.sent_content_start:
terminal.append(self._build_content_stop_event(fallback_id))
if self.sent_interaction_start or self.collected_text:
terminal.append(self._build_completion_event(fallback_id))
self._sent_completion_event = True
return terminal
# ------------------------------------------------------------------
# Iteration
# ------------------------------------------------------------------
def __iter__(self) -> Iterator[InteractionsAPIStreamingResponse]:
"""Sync iterator implementation."""
return self
def __next__(self) -> InteractionsAPIStreamingResponse:
"""Get next chunk in sync mode."""
if self._pending_events:
return self._pending_events.popleft()
if self.finished:
raise StopIteration
# Check if we have a pending interaction.complete to send
if hasattr(self, "_pending_interaction_complete"):
pending: InteractionsAPIStreamingResponse = getattr(
self, "_pending_interaction_complete"
)
delattr(self, "_pending_interaction_complete")
return pending
# Drain events queued from a prior chunk (e.g. content.start emitted alongside
# the interaction.start fallback for the first OutputTextDeltaEvent).
if self._pending_events:
return self._pending_events.pop(0)
# Use a loop instead of recursion to avoid stack overflow
sync_iterator = cast(
SyncResponsesAPIStreamingIterator, self.responses_stream_iterator
)
while True:
try:
# Get next chunk from responses API stream
chunk = next(sync_iterator)
# Transform chunk (chunk is already a ResponsesAPIStreamingResponse)
transformed = self._transform_responses_chunk_to_interactions_chunk(
chunk
)
if transformed:
# If we finished and content was started, send content.stop before interaction.complete
if (
self.finished
and self.sent_content_start
and transformed.event_type == "interaction.complete"
):
# Send content.stop first
content_stop = InteractionsAPIStreamingResponse(
event_type="content.stop",
id=transformed.id,
object="content",
delta={"type": "text", "text": self.collected_text},
)
# Store the interaction.complete to send next
self._pending_interaction_complete = transformed
return content_stop
return transformed
# If no transformation, continue to next chunk (loop continues)
except StopIteration:
self.finished = True
self._pending_events.extend(self._build_terminal_events_on_eof())
if self._pending_events:
return self._pending_events.popleft()
raise
# Send final events if needed
if self.sent_content_start:
return InteractionsAPIStreamingResponse(
event_type="content.stop",
object="content",
delta={"type": "text", "text": self.collected_text},
)
raise StopIteration
events = self._events_for_chunk(chunk)
if events:
self._pending_events.extend(events)
return self._pending_events.popleft()
def __aiter__(self) -> AsyncIterator[InteractionsAPIStreamingResponse]:
"""Async iterator implementation."""
return self
async def __anext__(self) -> InteractionsAPIStreamingResponse:
"""Get next chunk in async mode."""
if self._pending_events:
return self._pending_events.popleft()
if self.finished:
raise StopAsyncIteration
# Check if we have a pending interaction.complete to send
if hasattr(self, "_pending_interaction_complete"):
pending: InteractionsAPIStreamingResponse = getattr(
self, "_pending_interaction_complete"
)
delattr(self, "_pending_interaction_complete")
return pending
# Drain events queued from a prior chunk (e.g. content.start emitted alongside
# the interaction.start fallback for the first OutputTextDeltaEvent).
if self._pending_events:
return self._pending_events.pop(0)
# Use a loop instead of recursion to avoid stack overflow
async_iterator = cast(
ResponsesAPIStreamingIterator, self.responses_stream_iterator
)
while True:
try:
# Get next chunk from responses API stream
chunk = await async_iterator.__anext__()
# Transform chunk (chunk is already a ResponsesAPIStreamingResponse)
transformed = self._transform_responses_chunk_to_interactions_chunk(
chunk
)
if transformed:
# If we finished and content was started, send content.stop before interaction.complete
if (
self.finished
and self.sent_content_start
and transformed.event_type == "interaction.complete"
):
# Send content.stop first
content_stop = InteractionsAPIStreamingResponse(
event_type="content.stop",
id=transformed.id,
object="content",
delta={"type": "text", "text": self.collected_text},
)
# Store the interaction.complete to send next
self._pending_interaction_complete = transformed
return content_stop
return transformed
# If no transformation, continue to next chunk (loop continues)
except StopAsyncIteration:
self.finished = True
self._pending_events.extend(self._build_terminal_events_on_eof())
if self._pending_events:
return self._pending_events.popleft()
raise
# Send final events if needed
if self.sent_content_start:
return InteractionsAPIStreamingResponse(
event_type="content.stop",
object="content",
delta={"type": "text", "text": self.collected_text},
)
events = self._events_for_chunk(chunk)
if events:
self._pending_events.extend(events)
return self._pending_events.popleft()
raise StopAsyncIteration
# ------------------------------------------------------------------
# Backwards-compatible single-chunk transform (used by tests and any
# external callers that drove the iterator chunk-by-chunk pre-fix).
# ------------------------------------------------------------------
def _transform_responses_chunk_to_interactions_chunk(
self,
responses_chunk: ResponsesAPIStreamingResponse,
) -> Optional[InteractionsAPIStreamingResponse]:
"""
Compatibility shim: returns the *first* event produced for this chunk
and queues any remaining events on ``self._pending_events`` so they
are surfaced on subsequent calls/iterations.
Prefer ``_events_for_chunk`` in new code.
"""
events = self._events_for_chunk(responses_chunk)
if not events:
return None
first = events[0]
if len(events) > 1:
self._pending_events.extend(events[1:])
return first

View file

@ -226,29 +226,37 @@ class LiteLLMResponsesInteractionsConfig:
- Map status
- Extract usage
"""
# Extract text from outputs
outputs = []
# Extract text from outputs and build both `outputs` (legacy) and `steps` (new schema).
outputs: List[Dict[str, Any]] = []
steps: List[Dict[str, Any]] = []
if hasattr(responses_response, "output") and responses_response.output:
for output_item in responses_response.output:
# Use getattr with None default to safely access content
content = getattr(output_item, "content", None)
if content is not None:
content_items = content if isinstance(content, list) else [content]
model_output_contents: List[Dict[str, Any]] = []
for content_item in content_items:
# Check if content_item has text attribute
text = getattr(content_item, "text", None)
if text is not None:
outputs.append(
{
"type": "text",
"text": text,
}
)
# Use independent dict instances so mutations to one
# of `outputs` / `steps` don't leak into the other.
outputs.append({"type": "text", "text": text})
model_output_contents.append({"type": "text", "text": text})
elif (
isinstance(content_item, dict)
and content_item.get("type") == "text"
):
outputs.append(content_item)
outputs.append({**content_item})
model_output_contents.append({**content_item})
if model_output_contents:
steps.append(
{
"type": "model_output",
"content": model_output_contents,
}
)
# Convert created_at to ISO string
created_at = getattr(responses_response, "created_at", None)
@ -270,12 +278,14 @@ class LiteLLMResponsesInteractionsConfig:
else:
interactions_status = status
# Build interactions response
# Build interactions response — populate both `outputs` (legacy schema) and
# `steps` (new schema) so callers work regardless of which schema they expect.
interactions_response_dict: Dict[str, Any] = {
"id": getattr(responses_response, "id", ""),
"object": "interaction",
"status": interactions_status,
"outputs": outputs,
"steps": steps,
"model": model or getattr(responses_response, "model", ""),
"created": created,
}

View file

@ -101,10 +101,14 @@ class BaseInteractionsAPIStreamingIterator:
)
)
# Store the completed response (check for status=completed)
if (
streaming_response
and getattr(streaming_response, "status", None) == "completed"
# Store the completed response.
# Legacy schema signals completion via status="completed".
# New schema (Api-Revision: 2026-05-20) uses event_type="interaction.completed".
# Remove the legacy check after June 8, 2026.
if streaming_response and (
getattr(streaming_response, "status", None) == "completed"
or getattr(streaming_response, "event_type", None)
== "interaction.completed"
):
self.completed_response = streaming_response
self._handle_logging_completed_response()

View file

@ -1344,6 +1344,7 @@ def _get_dummy_thought_signature() -> str:
def convert_to_gemini_tool_call_invoke(
message: ChatCompletionAssistantMessage,
model: Optional[str] = None,
custom_llm_provider: Optional[str] = None,
) -> List[VertexPartType]:
"""
OpenAI tool invokes:
@ -1394,7 +1395,10 @@ def convert_to_gemini_tool_call_invoke(
)
forward_tool_call_id = bool(
model and VertexGeminiConfig._is_gemini_3_or_newer(model)
model
and VertexGeminiConfig._forward_gemini_function_call_id(
model, custom_llm_provider
)
)
if tool_calls is not None:
@ -1475,6 +1479,7 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage],
last_message_with_tool_calls: Optional[dict],
model: Optional[str] = None,
custom_llm_provider: Optional[str] = None,
) -> Union[VertexPartType, List[VertexPartType]]:
"""
OpenAI message with a tool result looks like:
@ -1616,14 +1621,16 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
name = tool.get("function", {}).get("name", "")
# Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix).
# Only Gemini 3+ accepts (and returns) an `id` on function_response parts;
# older Gemini models reject the field with a 400.
# Only Google AI Studio Gemini 3+ accepts `id` on function_response parts.
# Vertex AI and older Gemini models reject the field with HTTP 400.
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
gemini_call_id: Optional[str] = None
if model and VertexGeminiConfig._is_gemini_3_or_newer(model):
if model and VertexGeminiConfig._forward_gemini_function_call_id(
model, custom_llm_provider
):
raw_tool_call_id = message.get("tool_call_id")
if raw_tool_call_id and isinstance(raw_tool_call_id, str):
stripped_id = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0]

View file

@ -1506,9 +1506,21 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
optional_params["metadata"] = {"user_id": value}
elif param == "thinking":
optional_params["thinking"] = value
elif param == "reasoning_effort" and isinstance(value, str):
elif param == "reasoning_effort":
# Accept both string ("low") and dict ({"effort": "low",
# "summary": "concise"}). The Responses->Chat parser keeps the
# full dict when `summary` is set (see #25359), so a dict here
# is the standard shape Otto/OpenAI-Responses-Bridge callers
# send. Coerce to the effort string before mapping — same
# shape-tolerance the GPT-5 path already implements in
# `_normalize_reasoning_effort_for_chat_completion`.
effort_value = value
if isinstance(effort_value, dict):
effort_value = effort_value.get("effort")
if not isinstance(effort_value, str):
continue
mapped_thinking = AnthropicConfig._map_reasoning_effort(
reasoning_effort=value,
reasoning_effort=effort_value,
model=model,
llm_provider=self.custom_llm_provider or "anthropic",
)
@ -1519,12 +1531,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
optional_params["thinking"] = mapped_thinking
if AnthropicConfig._is_adaptive_thinking_model(model):
mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(
value
effort_value
)
if mapped_effort is None:
AnthropicConfig._raise_invalid_reasoning_effort(
model=model,
value=value,
value=effort_value,
llm_provider=self.custom_llm_provider or "anthropic",
)
optional_params["output_config"] = {"effort": mapped_effort}

View file

@ -1,9 +1,16 @@
from typing import Optional
from urllib.parse import parse_qs, urlparse, urlunparse
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
from litellm.types.router import GenericLiteLLMParams
# Endpoint-specific path suffixes that may appear in a deployment's api_base
# (e.g. the responses endpoint URL is stored as api_base for Azure models).
# Strip these before building the containers URL so we always start from the
# resource root (https://resource.cognitiveservices.azure.com).
_AZURE_ENDPOINT_PATHS = ("/openai/responses",)
class AzureContainerConfig(OpenAIContainerConfig):
"""
@ -27,6 +34,27 @@ class AzureContainerConfig(OpenAIContainerConfig):
litellm_params=GenericLiteLLMParams(api_key=api_key),
)
@staticmethod
def _normalize_api_base(api_base: Optional[str]) -> Optional[str]:
"""Strip endpoint-specific path suffixes from api_base to get the resource root."""
if not api_base:
return api_base
parsed = urlparse(api_base)
path = parsed.path.rstrip("/")
for ep in _AZURE_ENDPOINT_PATHS:
if path.endswith(ep):
return urlunparse(
(parsed.scheme, parsed.netloc, path[: -len(ep)], "", "", "")
)
return api_base
@staticmethod
def _extract_api_version(api_base: Optional[str]) -> Optional[str]:
"""Return the api-version query param from api_base if present."""
if not api_base:
return None
return parse_qs(urlparse(api_base).query).get("api-version", [None])[0]
def get_complete_url(
self,
api_base: Optional[str],
@ -39,10 +67,19 @@ class AzureContainerConfig(OpenAIContainerConfig):
{endpoint}/openai/v1/containers
when api_version is 'v1', 'latest', or 'preview'; otherwise:
{endpoint}/openai/containers
The deployment's api_base may be the responses endpoint URL
(e.g. .../openai/responses?api-version=2025-04-01-preview). We
prefer the api-version embedded there over the deployment's
api_version field, which may point to an older chat API version.
"""
effective_params = dict(litellm_params)
api_version_from_base = self._extract_api_version(api_base)
if api_version_from_base:
effective_params["api_version"] = api_version_from_base
return BaseAzureLLM._get_base_azure_url(
api_base=api_base,
litellm_params=litellm_params,
api_base=self._normalize_api_base(api_base),
litellm_params=effective_params,
route="/openai/containers",
default_api_version="v1",
)

View file

@ -257,14 +257,19 @@ class GenericContainerHandler:
returns_binary = endpoint_config.get("returns_binary", False)
is_multipart = endpoint_config.get("is_multipart", False)
# An empty dict passed as `params` to httpx strips any existing query
# string from the URL (e.g. ?api-version=...). Use None instead so
# httpx leaves the URL's own query string intact.
effective_params = query_params or None
try:
if method == "GET":
response = http_client.get(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
elif method == "DELETE":
response = http_client.delete(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
elif method == "POST":
if is_multipart and "file" in kwargs:
@ -272,11 +277,11 @@ class GenericContainerHandler:
kwargs["file"], headers
)
response = http_client.post(
url=url, headers=headers, params=query_params, files=files
url=url, headers=headers, params=effective_params, files=files
)
else:
response = http_client.post(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
else:
raise ValueError(f"Unsupported HTTP method: {method}")
@ -376,14 +381,19 @@ class GenericContainerHandler:
returns_binary = endpoint_config.get("returns_binary", False)
is_multipart = endpoint_config.get("is_multipart", False)
# An empty dict passed as `params` to httpx strips any existing query
# string from the URL (e.g. ?api-version=...). Use None instead so
# httpx leaves the URL's own query string intact.
effective_params = query_params or None
try:
if method == "GET":
response = await http_client.get(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
elif method == "DELETE":
response = await http_client.delete(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
elif method == "POST":
if is_multipart and "file" in kwargs:
@ -391,11 +401,11 @@ class GenericContainerHandler:
kwargs["file"], headers
)
response = await http_client.post(
url=url, headers=headers, params=query_params, files=files
url=url, headers=headers, params=effective_params, files=files
)
else:
response = await http_client.post(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
else:
raise ValueError(f"Unsupported HTTP method: {method}")

View file

@ -7838,7 +7838,7 @@ class BaseLLMHTTPHandler:
response = sync_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_list_response(
@ -7915,7 +7915,7 @@ class BaseLLMHTTPHandler:
response = await async_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_list_response(
@ -8005,7 +8005,7 @@ class BaseLLMHTTPHandler:
response = sync_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_retrieve_response(
@ -8082,7 +8082,7 @@ class BaseLLMHTTPHandler:
response = await async_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_retrieve_response(
@ -8172,7 +8172,7 @@ class BaseLLMHTTPHandler:
response = sync_httpx_client.delete(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_delete_response(
@ -8249,7 +8249,7 @@ class BaseLLMHTTPHandler:
response = await async_httpx_client.delete(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_delete_response(
@ -8345,7 +8345,7 @@ class BaseLLMHTTPHandler:
response = sync_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_file_list_response(
@ -8424,7 +8424,7 @@ class BaseLLMHTTPHandler:
response = await async_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_file_list_response(
@ -8512,7 +8512,7 @@ class BaseLLMHTTPHandler:
response = sync_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_file_content_response(
@ -8588,7 +8588,7 @@ class BaseLLMHTTPHandler:
response = await async_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_file_content_response(

View file

@ -164,5 +164,8 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
# If conversion fails, leave as is and let the API handle it
pass
return _gemini_convert_messages_with_history(
messages=messages, model=model, litellm_params=litellm_params
messages=messages,
model=model,
litellm_params=litellm_params,
custom_llm_provider="gemini",
)

View file

@ -6,13 +6,18 @@ Per OpenAPI spec (https://ai.google.dev/static/api/interactions.openapi.json):
- Get: GET https://generativelanguage.googleapis.com/{api_version}/interactions/{interaction_id}
- Delete: DELETE https://generativelanguage.googleapis.com/{api_version}/interactions/{interaction_id}
This is a thin wrapper - no transformation needed since we follow the spec directly.
Schema versioning:
- Default (Api-Revision: 2026-05-20): new `steps` schema.
- Legacy (Api-Revision: 2026-05-07): old `outputs` schema, controlled via
litellm.use_legacy_interactions_schema = True. Remove flag after June 8, 2026.
"""
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
@ -84,6 +89,15 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
api_key = GeminiModelInfo.get_api_key(litellm_params.get("api_key"))
if api_key:
headers["x-goog-api-key"] = api_key
# Inject the Api-Revision header to select the response schema.
# Default to the new `steps` schema unless the operator has opted out.
# Remove this conditional after June 8, 2026 and always use 2026-05-20.
if litellm.use_legacy_interactions_schema:
headers["Api-Revision"] = "2026-05-07"
else:
headers["Api-Revision"] = "2026-05-20"
return headers
def get_complete_url(
@ -119,8 +133,19 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
headers: dict,
) -> Dict:
"""
Build request body per OpenAPI spec - minimal transformation.
Build request body per OpenAPI spec.
When on the new schema (use_legacy_interactions_schema=False, the default):
- ``response_mime_type`` is folded into ``response_format`` and stripped from
the body (the field was removed in Api-Revision 2026-05-20).
- ``generation_config.image_config`` is moved to a ``response_format`` entry
with ``"type": "image"`` (also removed from generation_config in 2026-05-20).
When on the legacy schema (use_legacy_interactions_schema=True):
- All fields are forwarded as-is.
"""
use_legacy: bool = litellm.use_legacy_interactions_schema
request_body: Dict[str, Any] = {}
# Model or Agent (one required)
@ -135,24 +160,81 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
if input is not None:
request_body["input"] = input
# Pass through optional params directly (they match the spec)
# Pass through optional params — legacy schema keeps all fields as-is.
optional_keys = [
"tools",
"system_instruction",
"generation_config",
"stream",
"store",
"background",
"environment",
"response_modalities",
"response_format",
"response_mime_type",
"previous_interaction_id",
]
for key in optional_keys:
if optional_params.get(key) is not None:
request_body[key] = optional_params[key]
if use_legacy:
# Legacy schema: forward response_mime_type and response_format as-is.
for key in ("response_format", "response_mime_type", "generation_config"):
if optional_params.get(key) is not None:
request_body[key] = optional_params[key]
else:
# New schema (Api-Revision: 2026-05-20):
# response_mime_type is removed — fold it into response_format.
response_format = optional_params.get("response_format")
response_mime_type = optional_params.get("response_mime_type")
if (
response_mime_type
and not isinstance(response_format, list)
and (
not isinstance(response_format, dict)
or "mime_type" not in response_format
)
):
# Wrap the legacy schema into the new polymorphic format.
new_rf: Dict[str, Any] = {
"type": "text",
"mime_type": response_mime_type,
}
if response_format is not None:
new_rf["schema"] = response_format
response_format = new_rf
if response_format is not None:
request_body["response_format"] = response_format
# image_config moves out of generation_config into response_format.
generation_config: Optional[Dict[str, Any]] = optional_params.get(
"generation_config"
)
if generation_config is not None:
image_config = None
if isinstance(generation_config, dict):
generation_config = dict(
generation_config
) # avoid mutating the caller's dict
image_config = generation_config.pop("image_config", None)
if not generation_config:
generation_config = None
if generation_config is not None:
request_body["generation_config"] = generation_config
if image_config is not None:
# Move image_config to response_format with type=image.
image_rf: Dict[str, Any] = {"type": "image", **image_config}
existing_rf = request_body.get("response_format")
if existing_rf is None:
request_body["response_format"] = image_rf
elif isinstance(existing_rf, list):
request_body["response_format"] = [*existing_rf, image_rf]
else:
# Convert single entry to array for multimodal output.
request_body["response_format"] = [existing_rf, image_rf]
return request_body
def transform_response(

View file

@ -174,7 +174,9 @@ def transform_openai_messages_to_gemini_context_caching(
)
transformed_messages = _gemini_convert_messages_with_history(
messages=new_messages, model=model
messages=new_messages,
model=model,
custom_llm_provider=custom_llm_provider,
)
model_name = "models/{}".format(model)

View file

@ -682,6 +682,7 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
messages: List[AllMessageValues],
model: Optional[str] = None,
litellm_params: Optional[dict] = None,
custom_llm_provider: Optional[str] = None,
) -> List[ContentType]:
"""
Converts given messages from OpenAI format to Gemini format
@ -983,7 +984,9 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
or assistant_msg.get("function_call") is not None
): # support assistant tool invoke conversion
gemini_tool_call_parts = convert_to_gemini_tool_call_invoke(
assistant_msg, model=model
assistant_msg,
model=model,
custom_llm_provider=custom_llm_provider,
)
## check if gemini_tool_call already exists in assistant_content
for gemini_tool_call_part in gemini_tool_call_parts:
@ -1045,6 +1048,7 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
messages[msg_i], # type: ignore
last_message_with_tool_calls, # type: ignore
model=model,
custom_llm_provider=custom_llm_provider,
)
msg_i += 1
# Handle both single part and list of parts (for Computer Use with images)

View file

@ -289,6 +289,20 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return True
return False
@staticmethod
def _forward_gemini_function_call_id(
model: str, custom_llm_provider: Optional[str] = None
) -> bool:
"""
Whether to include `id` on function_call / function_response parts.
Gemini 3+ on Google AI Studio accepts (and returns) `id` for strict
tool-call matching. Vertex AI rejects the field with HTTP 400.
"""
if custom_llm_provider != "gemini":
return False
return VertexGeminiConfig._is_gemini_3_or_newer(model)
def _supports_penalty_parameters(self, model: str) -> bool:
# Gemini 3 models do not support penalty parameters
if VertexGeminiConfig._is_gemini_3_or_newer(model):
@ -2649,7 +2663,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
litellm_params: Optional[dict] = None,
) -> List[ContentType]:
return _gemini_convert_messages_with_history(
messages=messages, model=model, litellm_params=litellm_params
messages=messages,
model=model,
litellm_params=litellm_params,
custom_llm_provider="vertex_ai",
)
def get_error_class(

View file

@ -27326,6 +27326,58 @@
"supports_web_search": true,
"tpm": 800000
},
"openrouter/google/gemini-3.1-flash-lite": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_per_audio_token": 5e-08,
"input_cost_per_audio_token": 5e-07,
"input_cost_per_token": 2.5e-07,
"litellm_provider": "openrouter",
"max_audio_length_hours": 8.4,
"max_audio_per_prompt": 1,
"max_images_per_prompt": 3000,
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
"max_pdf_size_mb": 30,
"max_tokens": 65536,
"max_video_length": 1,
"max_videos_per_prompt": 10,
"mode": "chat",
"output_cost_per_reasoning_token": 1.5e-06,
"output_cost_per_token": 1.5e-06,
"rpm": 2000,
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-lite",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/completions",
"/v1/batch"
],
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_audio_input": true,
"supports_audio_output": false,
"supports_code_execution": true,
"supports_file_search": true,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_url_context": true,
"supports_video_input": true,
"supports_vision": true,
"supports_web_search": true,
"tpm": 800000
},
"openrouter/google/gemini-3.1-pro-preview": {
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,

View file

@ -1226,6 +1226,7 @@ class MCPServerManager:
tools = await self._get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
user_api_key_auth=user_api_key_auth,
)
return tools
except Exception as e:
@ -1406,6 +1407,7 @@ class MCPServerManager:
extra_headers: Optional[Dict[str, str]] = None,
add_prefix: bool = True,
raw_headers: Optional[Dict[str, str]] = None,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[MCPTool]:
"""
Helper method to get tools from a single MCP server with prefixed names.
@ -1432,6 +1434,46 @@ class MCPServerManager:
extra_headers = {}
extra_headers.update(server.static_headers)
# MCPJWTSigner: inject signed JWT for tools/list (list path skips pre_call_hook).
# Skip entirely when the signer is not configured (avoid an unnecessary
# dict copy on every list call), when the server has its own static
# Authorization header, when a per-user mcp_auth_header has already
# been resolved, or when the caller already supplied an Authorization
# entry in extra_headers (e.g. a per-user OAuth token resolved
# upstream) — admin-configured static auth and per-user OAuth must
# take precedence so the signer doesn't silently overwrite e.g. an
# upstream API key or a user's OAuth token (MCPClient._get_auth_headers
# applies extra_headers after writing Authorization from auth_value, so
# an injected JWT would otherwise clobber the per-user token).
if user_api_key_auth is not None and not server.spec_path:
from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
get_mcp_jwt_signer,
inject_mcp_jwt_headers_for_upstream,
)
static_headers = server.static_headers or {}
has_static_authorization = any(
isinstance(k, str) and k.lower() == "authorization"
for k in static_headers.keys()
)
has_extra_authorization = bool(extra_headers) and any(
isinstance(k, str) and k.lower() == "authorization"
for k in (extra_headers or {}).keys()
)
if (
get_mcp_jwt_signer() is not None
and not has_static_authorization
and not mcp_auth_header
and not has_extra_authorization
):
extra_headers = await inject_mcp_jwt_headers_for_upstream(
user_api_key_dict=user_api_key_auth,
extra_headers=extra_headers,
raw_headers=raw_headers,
for_list_tools=True,
)
stdio_env = self._build_stdio_env(server, raw_headers)
client = await self._create_mcp_client(
@ -2791,6 +2833,112 @@ class MCPServerManager:
return cast(CallToolResult, result)
def _resolve_mcp_server_for_tool_call(
self,
server_name: str,
name: str,
) -> MCPServer:
"""Resolve MCP server for call_tool (prefixed name, registry, fallback)."""
prefixed_tool_name = add_server_prefix_to_name(name, server_name)
mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name)
resolved_by_server_name_only = False
normalized_server_name = normalize_server_name(server_name)
def _candidate_matches_server_name(candidate: MCPServer) -> bool:
for identifier in (
candidate.alias,
candidate.server_name,
candidate.name,
):
if identifier and normalize_server_name(identifier) == (
normalized_server_name
):
return True
return False
if mcp_server is None:
for candidate in self.get_registry().values():
if _candidate_matches_server_name(candidate):
mcp_server = candidate
resolved_by_server_name_only = True
break
if mcp_server is None:
fallback = self._get_mcp_server_from_tool_name(name)
if fallback is not None and (
not server_name or _candidate_matches_server_name(fallback)
):
mcp_server = fallback
if mcp_server is None:
raise ValueError(f"Tool {name} not found")
if resolved_by_server_name_only:
tool_known = (
name in self.tool_name_to_mcp_server_name_mapping
or prefixed_tool_name in self.tool_name_to_mcp_server_name_mapping
)
if not tool_known:
raise ValueError(f"Tool {name} not found")
return mcp_server
async def _resolve_oauth2_headers_for_tool_call(
self,
mcp_server: MCPServer,
oauth2_headers: Optional[Dict[str, str]],
user_api_key_auth: Optional[UserAPIKeyAuth],
) -> Optional[Dict[str, str]]:
"""Look up per-user OAuth headers when the client did not supply a token."""
if (
not mcp_server.needs_user_oauth_token
or oauth2_headers
or user_api_key_auth is None
):
return oauth2_headers
user_id = getattr(user_api_key_auth, "user_id", None)
if not user_id:
return oauth2_headers
try:
from litellm.proxy._experimental.mcp_server.server import ( # noqa: PLC0415
_get_user_oauth_extra_headers_from_db,
)
stored_headers = await _get_user_oauth_extra_headers_from_db(
server=mcp_server,
user_api_key_auth=user_api_key_auth,
)
if stored_headers:
return stored_headers
except Exception as _lookup_exc:
verbose_logger.debug(
"call_tool: per-user token lookup failed for " "user=%s server=%s: %s",
user_id,
mcp_server.server_id,
_lookup_exc,
)
return oauth2_headers
async def _gather_openapi_tool_tasks(
self,
tasks: List[Any],
proxy_logging_obj: Optional[ProxyLogging],
) -> CallToolResult:
"""Await OpenAPI tool tasks and return the tool call result."""
try:
mcp_responses = await asyncio.gather(*tasks)
result_index = 1 if proxy_logging_obj else 0
return cast(CallToolResult, mcp_responses[result_index])
except (
BlockedPiiEntityError,
GuardrailRaisedException,
HTTPException,
) as e:
verbose_logger.error(
f"Guardrail blocked MCP tool call during result check: {str(e)}"
)
raise e
async def call_tool(
self,
server_name: str,
@ -2821,12 +2969,7 @@ class MCPServerManager:
CallToolResult from the MCP server
"""
start_time = datetime.datetime.now()
# Get the MCP server
prefixed_tool_name = add_server_prefix_to_name(name, server_name)
mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name)
if mcp_server is None:
raise ValueError(f"Tool {name} not found")
mcp_server = self._resolve_mcp_server_for_tool_call(server_name, name)
#########################################################
# Pre MCP Tool Call Hook
@ -2860,36 +3003,9 @@ class MCPServerManager:
)
tasks.append(during_hook_task)
# For per-user OAuth servers: if the client didn't supply a token in
# oauth2_headers, look up the stored token from Redis / DB. This is the
# call_tool equivalent of _get_user_oauth_extra_headers_from_db used in
# list_tools.
if (
mcp_server.needs_user_oauth_token
and not oauth2_headers
and user_api_key_auth is not None
):
user_id = getattr(user_api_key_auth, "user_id", None)
if user_id:
try:
from litellm.proxy._experimental.mcp_server.server import ( # noqa: PLC0415
_get_user_oauth_extra_headers_from_db,
)
stored_headers = await _get_user_oauth_extra_headers_from_db(
server=mcp_server,
user_api_key_auth=user_api_key_auth,
)
if stored_headers:
oauth2_headers = stored_headers
except Exception as _lookup_exc:
verbose_logger.debug(
"call_tool: per-user token lookup failed for "
"user=%s server=%s: %s",
user_id,
mcp_server.server_id,
_lookup_exc,
)
oauth2_headers = await self._resolve_oauth2_headers_for_tool_call(
mcp_server, oauth2_headers, user_api_key_auth
)
# For OpenAPI servers, call the tool handler directly instead of via MCP client
if mcp_server.spec_path:
@ -2925,26 +3041,7 @@ class MCPServerManager:
hook_extra_headers=hook_result.get("extra_headers"),
)
# For OpenAPI tools, await outside the client context
try:
mcp_responses = await asyncio.gather(*tasks)
# If proxy_logging_obj is None, the tool call result is at index 0
# If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task)
result_index = 1 if proxy_logging_obj else 0
result = mcp_responses[result_index]
return cast(CallToolResult, result)
except (
BlockedPiiEntityError,
GuardrailRaisedException,
HTTPException,
) as e:
# Re-raise guardrail exceptions to properly fail the MCP call
verbose_logger.error(
f"Guardrail blocked MCP tool call during result check: {str(e)}"
)
raise e
return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj)
#########################################################
# End of Methods that call the upstream MCP servers

View file

@ -29,6 +29,16 @@ _DEFAULT_PORTS = {"http": 80, "https": 443}
# subdomain. HTTPS only.
_TRUSTED_REDIRECT_ORIGINS_ENV = "MCP_TRUSTED_REDIRECT_ORIGINS"
# Comma-separated private-use URI allowlist for native MCP clients.
# A trailing ``*`` is a prefix match; end the prefix with ``/`` (e.g.
# ``myapp://host/oauth/*``) so ``.../oauth/callback*`` does not also
# match ``.../oauth/callback-2``.
_TRUSTED_NATIVE_REDIRECT_URIS_ENV = "MCP_TRUSTED_NATIVE_REDIRECT_URIS"
# Default allowlist for trusted native redirect URIs.
_DEFAULT_NATIVE_REDIRECT_URIS: List[str] = [
"cursor://anysphere.cursor-mcp/oauth/callback",
]
_warned_invalid_proxy_base_url: Optional[str] = None
@ -212,10 +222,82 @@ def _matches_trusted_origin_entry(netloc: str, entry: str) -> bool:
return netloc == entry
def _normalize_native_redirect_uri(
parsed,
) -> str:
"""Lowercase scheme, netloc, and path for allowlist comparison."""
return urlunparse(
(
(parsed.scheme or "").lower(),
(parsed.netloc or "").lower(),
(parsed.path or "").lower(),
"",
"",
"",
)
)
def _parse_trusted_native_redirect_uris() -> List[str]:
"""Built-in native MCP callbacks plus ``MCP_TRUSTED_NATIVE_REDIRECT_URIS``."""
entries: List[str] = [uri.lower() for uri in _DEFAULT_NATIVE_REDIRECT_URIS]
raw = os.environ.get(_TRUSTED_NATIVE_REDIRECT_URIS_ENV, "").strip()
if not raw:
return entries
for token in raw.split(","):
entry = token.strip().lower()
if entry and entry not in entries:
entries.append(entry)
return entries
def _native_wildcard_prefix_matches(normalized: str, prefix: str) -> bool:
"""Prefix match for ``entry*`` allowlist rows.
When the prefix does not end with ``/``, only exact matches or
deeper path segments (``prefix/...``) are accepted — not siblings
like ``prefix-2``.
"""
if not normalized.startswith(prefix):
return False
suffix = normalized[len(prefix) :]
if not suffix:
return True
if prefix.endswith("/"):
return True
return suffix[0] == "/"
def _matches_trusted_native_redirect_uri(parsed) -> bool:
"""Allowlisted private-use / custom-scheme OAuth callbacks for native MCP clients."""
if parsed.fragment:
return False
# Query strings are not part of registered redirect_uris (RFC 6749 §3.1.2).
# Rejecting them prevents allowlist bypass via ``.../callback?injected=...``.
if parsed.query:
return False
if not parsed.netloc:
return False
if parsed.username is not None or parsed.password is not None:
return False
if "\\" in parsed.netloc:
return False
normalized = _normalize_native_redirect_uri(parsed)
for entry in _parse_trusted_native_redirect_uris():
if entry.endswith("*"):
if _native_wildcard_prefix_matches(normalized, entry[:-1]):
return True
elif normalized == entry:
return True
return False
def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None:
"""Accept ``redirect_uri`` when it is (a) same-origin with the
proxy's own request origin, (b) loopback, or (c) listed in the
``MCP_TRUSTED_REDIRECT_ORIGINS`` ops allowlist.
proxy's own request origin, (b) loopback, (c) listed in the
``MCP_TRUSTED_REDIRECT_ORIGINS`` ops allowlist, or (d) a built-in /
env-configured native MCP client callback (e.g. ``cursor://``).
Same-origin is VERIA-57's threat-model-safe equivalent of loopback:
an attacker who can host content on the proxy's own HTTPS origin
@ -239,6 +321,8 @@ def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None:
except ValueError:
raise HTTPException(status_code=400, detail="invalid_request")
if parsed.scheme not in ("http", "https"):
if _matches_trusted_native_redirect_uri(parsed):
return
raise HTTPException(status_code=400, detail="invalid_request")
if parsed.fragment:
raise HTTPException(status_code=400, detail="invalid_request")
@ -310,9 +394,12 @@ def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None:
"Inbound headers: X-Forwarded-Proto=%r X-Forwarded-Host=%r "
"X-Forwarded-Port=%r Host=%r. "
"Trusted-redirect-origins env=%r. "
"Trusted-native-redirect-uris env=%r. "
"If this should be accepted, either align ingress X-Forwarded-* "
"with the browser URL, set PROXY_BASE_URL to your public origin, "
"or add the redirect_uri host to MCP_TRUSTED_REDIRECT_ORIGINS.",
"add the redirect_uri host to MCP_TRUSTED_REDIRECT_ORIGINS, or "
"for native MCP clients (cursor://, etc.) add the full redirect_uri "
"to MCP_TRUSTED_NATIVE_REDIRECT_URIS.",
redirect_uri,
proxy_base,
os.environ.get("PROXY_BASE_URL"),
@ -321,5 +408,6 @@ def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None:
request.headers.get("X-Forwarded-Port"),
request.headers.get("Host"),
os.environ.get(_TRUSTED_REDIRECT_ORIGINS_ENV),
os.environ.get(_TRUSTED_NATIVE_REDIRECT_URIS_ENV),
)
raise HTTPException(status_code=400, detail="invalid_request")

View file

@ -1,6 +1,17 @@
import importlib
from datetime import datetime
from typing import Any, Awaitable, Callable, Dict, List, Literal, Optional, Set, Union
from typing import (
Any,
Awaitable,
Callable,
Dict,
List,
Literal,
Optional,
Set,
Tuple,
Union,
)
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
@ -231,11 +242,32 @@ if MCP_AVAILABLE:
)
return mcp_auth_header, mcp_server_auth_headers, raw_headers
def _resolve_mcp_server_id_for_rest(
server_id: str,
allowed_server_ids: Union[Set[str], List[str]],
client_ip: Optional[str] = None,
) -> str:
"""
Map REST ``server_id`` (UUID, server_name, or alias) to canonical server_id.
tools/list already did this; tools/call must match so clients can pass
server names like ``order_status_mcp`` instead of only UUIDs.
"""
allowed = set(allowed_server_ids)
if server_id in allowed:
return server_id
by_name = global_mcp_server_manager.get_mcp_server_by_name(
server_id, client_ip=client_ip
)
if by_name is not None and by_name.server_id in allowed:
return by_name.server_id
return server_id
async def _resolve_allowed_mcp_servers_with_ip_filter(
request: Request,
user_api_key_dict: UserAPIKeyAuth,
server_id: str,
) -> List[MCPServer]:
) -> Tuple[List[MCPServer], str]:
"""
Resolve allowed MCP servers for a tool call with IP filtering.
@ -245,10 +277,10 @@ if MCP_AVAILABLE:
server_id: The server ID to validate access for
Returns:
List of allowed MCPServer objects
Tuple of (allowed MCPServer objects, canonical server_id)
Raises:
HTTPException: If the server_id is not allowed
HTTPException: If the server_id is not allowed or not found
"""
# Get all auth contexts
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
@ -268,8 +300,41 @@ if MCP_AVAILABLE:
)
)
# Check if the specified server_id is allowed
if server_id not in allowed_server_ids_set:
canonical_server_id = _resolve_mcp_server_id_for_rest(
server_id, allowed_server_ids_set, _rest_client_ip
)
if canonical_server_id not in allowed_server_ids_set:
_server = global_mcp_server_manager.get_mcp_server_by_id(
server_id
) or global_mcp_server_manager.get_mcp_server_by_name(server_id)
if (
_server is not None
and _rest_client_ip is not None
and not global_mcp_server_manager._is_server_accessible_from_ip(
_server, _rest_client_ip
)
):
raise HTTPException(
status_code=403,
detail={
"error": "ip_filtering",
"message": (
f"MCP server '{server_id}' is not accessible from your IP address "
f"({_rest_client_ip}). This server is restricted to internal "
"networks only. To make it externally accessible, set "
"'available_on_public_internet: true' in the server configuration."
),
},
)
if _server is None:
raise HTTPException(
status_code=404,
detail={
"error": "server_not_found",
"message": f"MCP server '{server_id}' was not found",
},
)
raise HTTPException(
status_code=403,
detail={
@ -285,7 +350,7 @@ if MCP_AVAILABLE:
if server is not None:
allowed_mcp_servers.append(server)
return allowed_mcp_servers
return allowed_mcp_servers, canonical_server_id
async def _get_tools_for_single_server(
server,
@ -301,6 +366,7 @@ if MCP_AVAILABLE:
extra_headers=extra_headers,
add_prefix=False,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
# Filter tools based on allowed_tools configuration
@ -753,7 +819,7 @@ if MCP_AVAILABLE:
},
)
tool_arguments = data.get("arguments")
tool_arguments = data.get("arguments") or {}
proxy_base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
(
@ -786,14 +852,18 @@ if MCP_AVAILABLE:
data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"]
# Resolve allowed MCP servers with IP filtering
allowed_mcp_servers = await _resolve_allowed_mcp_servers_with_ip_filter(
(
allowed_mcp_servers,
canonical_server_id,
) = await _resolve_allowed_mcp_servers_with_ip_filter(
request, user_api_key_dict, server_id
)
# Look up per-user OAuth headers for this server (mirrors list_tool_rest_api).
user_oauth_extra_headers: Optional[Dict[str, str]] = None
target_server = next(
(s for s in allowed_mcp_servers if s.server_id == server_id), None
(s for s in allowed_mcp_servers if s.server_id == canonical_server_id),
None,
)
if target_server is not None:
user_oauth_extra_headers = await _get_user_oauth_extra_headers(
@ -812,6 +882,7 @@ if MCP_AVAILABLE:
oauth2_headers=user_oauth_extra_headers or data.get("oauth2_headers"),
raw_headers=data.get("raw_headers"),
litellm_logging_obj=data.get("litellm_logging_obj"),
requested_server_id=canonical_server_id,
)
return result
except BlockedPiiEntityError as e:

View file

@ -1368,6 +1368,7 @@ if MCP_AVAILABLE:
extra_headers=extra_headers,
add_prefix=True, # Always add server prefix
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
filtered_tools = filter_tools_by_allowed_tools(tools, server)
@ -2074,6 +2075,7 @@ if MCP_AVAILABLE:
"""
# Track resolved MCP server for both permission checks and dispatch
mcp_server: Optional[MCPServer] = None
requested_server_id: Optional[str] = kwargs.get("requested_server_id")
# If the client called with a display-name override (e.g. "Get Pet"),
# translate it back to the original prefixed name before any routing.
@ -2082,14 +2084,55 @@ if MCP_AVAILABLE:
# Remove prefix from tool name for logging and processing
original_tool_name, server_name = split_server_prefix_from_name(name)
requested_server: Optional[MCPServer] = None
if requested_server_id:
requested_server = next(
(s for s in allowed_mcp_servers if s.server_id == requested_server_id),
None,
)
# Resolve the actual MCP server up-front so the permission check uses
# the canonical server.name even when the tool name is prefixed with a
# short ID (LITELLM_USE_SHORT_MCP_TOOL_PREFIX) that doesn't match the
# server's display name directly.
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
if mcp_server is None and requested_server is not None:
# REST callers may pass the raw tool name (no prefix) plus a
# ``requested_server_id``. The mapping might only contain the
# prefixed form, so retry the lookup with every known prefix of
# the requested server before treating the tool as unresolved —
# otherwise the tool_server_mismatch guard below is silently
# bypassed.
for known_prefix in iter_known_server_prefixes(requested_server):
candidate = global_mcp_server_manager._get_mcp_server_from_tool_name(
add_server_prefix_to_name(name, known_prefix)
)
if candidate is not None:
mcp_server = candidate
break
if mcp_server is not None:
server_name = mcp_server.name
# REST /mcp-rest/tools/call passes server_id — tool must belong to that server
if requested_server is not None:
if (
mcp_server is not None
and mcp_server.server_id != requested_server.server_id
):
raise HTTPException(
status_code=403,
detail={
"error": "tool_server_mismatch",
"message": (
f"Tool '{name}' belongs to MCP server '{mcp_server.name}' "
f"but request specified server_id for '{requested_server.name}'."
),
},
)
if mcp_server is None:
mcp_server = requested_server
server_name = requested_server.name
# Only enforce server-level permissions when we can resolve a server
if server_name:
if not MCPRequestHandler.is_tool_allowed(

File diff suppressed because one or more lines are too long

View file

@ -3171,7 +3171,7 @@
]
},
"post": {
"description": "Create a new agent\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/agents\" \\\n -H \"Authorization: Bearer <your_api_key>\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"agent\": {\n \"agent_name\": \"my-custom-agent\",\n \"agent_card_params\": {\n \"protocolVersion\": \"1.0\",\n \"name\": \"Hello World Agent\",\n \"description\": \"Just a hello world agent\",\n \"url\": \"http://localhost:9999/\",\n \"version\": \"1.0.0\",\n \"defaultInputModes\": [\"text\"],\n \"defaultOutputModes\": [\"text\"],\n \"capabilities\": {\n \"streaming\": true\n },\n \"skills\": [\n {\n \"id\": \"hello_world\",\n \"name\": \"Returns hello world\",\n \"description\": \"just returns hello world\",\n \"tags\": [\"hello world\"],\n \"examples\": [\"hi\", \"hello world\"]\n }\n ]\n },\n \"litellm_params\": {\n \"make_public\": true\n }\n }\n }'\n```",
"description": "Create a new agent\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/v1/agents\" \\\n -H \"Authorization: Bearer <your_api_key>\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"agent_name\": \"my-custom-agent\",\n \"agent_card_params\": {\n \"protocolVersion\": \"1.0\",\n \"name\": \"Hello World Agent\",\n \"description\": \"Just a hello world agent\",\n \"url\": \"http://localhost:9999/\",\n \"version\": \"1.0.0\",\n \"defaultInputModes\": [\"text\"],\n \"defaultOutputModes\": [\"text\"],\n \"capabilities\": {\n \"streaming\": true\n },\n \"skills\": [\n {\n \"id\": \"hello_world\",\n \"name\": \"Returns hello world\",\n \"description\": \"just returns hello world\",\n \"tags\": [\"hello world\"],\n \"examples\": [\"hi\", \"hello world\"]\n }\n ]\n },\n \"litellm_params\": {\n \"make_public\": true\n }\n }'\n```",
"operationId": "create_agent_v1_agents_post",
"requestBody": {
"content": {

View file

@ -2361,6 +2361,30 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
database_connection_timeout: Optional[float] = Field(
60, description="default timeout for a connection to the database"
)
database_connect_timeout: Optional[float] = Field(
None,
description=(
"Prisma `connect_timeout` URL param (seconds). Bounds how long the "
"engine waits to establish a new connection before failing. Defaults "
"to Prisma's built-in value when unset."
),
)
database_socket_timeout: Optional[float] = Field(
None,
description=(
"Prisma `socket_timeout` URL param (seconds). When set, an idle/slow "
"connection that has not produced data within this window is closed. "
"This is the main knob for capping idle DB connections from LiteLLM."
),
)
database_extra_connection_params: Optional[Dict[str, Any]] = Field(
None,
description=(
"Escape hatch: extra key/value pairs appended verbatim to the Prisma "
"DATABASE_URL / DIRECT_URL query string (e.g. `sslmode`, `pgbouncer`, "
"`statement_cache_size`). Keys here override any default LiteLLM sets."
),
)
database_type: Optional[Literal["dynamo_db"]] = Field(
None, description="to use dynamodb instead of postgres db"
)

View file

@ -523,6 +523,10 @@ async def retrieve_batch( # noqa: PLR0915
custom_llm_provider=custom_llm_provider, **data # type: ignore
)
response = await proxy_logging_obj.post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
# FIX: Update the database with the latest state from provider
await update_batch_in_database(
batch_id=batch_id,
@ -533,19 +537,9 @@ async def retrieve_batch( # noqa: PLR0915
verbose_proxy_logger=verbose_proxy_logger,
db_batch_object=db_batch_object,
operation="retrieve",
user_api_key_dict=user_api_key_dict,
)
### CALL HOOKS ### - modify outgoing data
response = await proxy_logging_obj.post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
# Fix: bug_feb14_batch_retrieve_returns_raw_input_file_id
# Resolve raw provider file IDs (input, output, error) to unified IDs.
if unified_batch_id:
await resolve_input_file_id_to_unified(response, prisma_client)
await resolve_output_file_ids_to_unified(response, prisma_client)
### ALERTING ###
asyncio.create_task(
proxy_logging_obj.update_request_status(
@ -917,10 +911,14 @@ async def cancel_batch(
**_cancel_batch_data,
)
# FIX: Update the database with the new cancelled state
managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files")
from litellm.proxy.proxy_server import prisma_client
response = await proxy_logging_obj.post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
# FIX: Update the database with the new cancelled state
await update_batch_in_database(
batch_id=batch_id,
unified_batch_id=unified_batch_id,
@ -929,11 +927,7 @@ async def cancel_batch(
prisma_client=prisma_client,
verbose_proxy_logger=verbose_proxy_logger,
operation="cancel",
)
### CALL HOOKS ### - modify outgoing data
response = await proxy_logging_obj.post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
user_api_key_dict=user_api_key_dict,
)
### ALERTING ###

View file

@ -328,7 +328,7 @@ async def retrieve_container(
custom_llm_provider=custom_llm_provider,
)
data.update(
get_container_forwarding_params(
await get_container_forwarding_params(
container_id,
original_container_id,
custom_llm_provider,
@ -433,7 +433,7 @@ async def delete_container(
custom_llm_provider=custom_llm_provider,
)
data.update(
get_container_forwarding_params(
await get_container_forwarding_params(
container_id,
original_container_id,
custom_llm_provider,

View file

@ -196,10 +196,12 @@ async def _process_binary_request(
)
data: Dict[str, Any] = {
"file_id": file_id,
**get_container_forwarding_params(
container_id=container_id,
original_container_id=original_container_id,
custom_llm_provider=resolved_provider,
**(
await get_container_forwarding_params(
container_id=container_id,
original_container_id=original_container_id,
custom_llm_provider=resolved_provider,
)
),
}
processor = ProxyBaseLLMRequestProcessing(data=data)
@ -316,7 +318,7 @@ async def _process_multipart_upload_request(
)
data.update(
get_container_forwarding_params(
await get_container_forwarding_params(
container_id=container_id,
original_container_id=original_container_id,
custom_llm_provider=resolved_provider,
@ -396,7 +398,7 @@ async def _process_request(
)
)
data.update(
get_container_forwarding_params(
await get_container_forwarding_params(
container_id=path_params["container_id"],
original_container_id=original_container_id,
custom_llm_provider=resolved_provider,

View file

@ -23,6 +23,13 @@ CONTAINER_OBJECT_PURPOSE = "container"
_NEGATIVE_OWNER_SENTINEL = "__litellm_container_no_owner__"
_CONTAINER_OWNER_CACHE = InMemoryCache(max_size_in_memory=10000, default_ttl=60)
# Caches the stored ``unified_object_id`` (the encoded container ID
# captured at create time) so ``get_container_forwarding_params`` can
# recover the deployment ``model_id`` for native upstream IDs without
# re-hitting Prisma on every retrieve/delete.
_NEGATIVE_STORED_ID_SENTINEL = "__litellm_container_no_stored_id__"
_CONTAINER_STORED_ID_CACHE = InMemoryCache(max_size_in_memory=10000, default_ttl=60)
# Per-caller-scope cache for ``GET /v1/containers`` list filtering. Without
# this, every list call issues a fresh ``find_many`` against
# ``litellm_managedobjecttable``. The cache key is the sorted owner-scope
@ -56,7 +63,7 @@ def decode_container_id_for_ownership(
return original_container_id, custom_llm_provider
def get_container_forwarding_params(
async def get_container_forwarding_params(
container_id: str, original_container_id: str, custom_llm_provider: str
) -> Dict[str, str]:
params = {
@ -65,6 +72,20 @@ def get_container_forwarding_params(
}
decoded = ResponsesAPIRequestUtils._decode_container_id(container_id)
model_id = decoded.get("model_id")
if not (isinstance(model_id, str) and model_id):
# Native upstream IDs (e.g. Azure ``cntr_<hex>``) carry no LiteLLM
# routing payload, so decoding the user-supplied id yields no
# ``model_id``. Recover it from the encoded ``unified_object_id``
# captured on the ownership row at create time — when the router
# selected a specific deployment that ID embeds the model_id.
stored_id = await _get_stored_container_id(
original_container_id, custom_llm_provider
)
if stored_id and stored_id != container_id:
stored_decoded = ResponsesAPIRequestUtils._decode_container_id(stored_id)
stored_model_id = stored_decoded.get("model_id")
if isinstance(stored_model_id, str) and stored_model_id:
model_id = stored_model_id
if isinstance(model_id, str) and model_id:
params["model_id"] = model_id
return params
@ -168,6 +189,7 @@ async def record_container_owner(
)
_CONTAINER_OWNER_CACHE.set_cache(model_object_id, owner)
_CONTAINER_STORED_ID_CACHE.set_cache(model_object_id, container_id)
# Drop the caller's own list-cache entry so the just-created container
# shows up on their next ``GET /v1/containers``. Other callers with
# disjoint scope tuples have their own entries; intersecting-scope
@ -207,9 +229,60 @@ async def _get_container_owner(
_CONTAINER_OWNER_CACHE.set_cache(
model_object_id, owner if owner is not None else _NEGATIVE_OWNER_SENTINEL
)
stored_id = getattr(row, "unified_object_id", None) if row is not None else None
_CONTAINER_STORED_ID_CACHE.set_cache(
model_object_id,
(
stored_id
if isinstance(stored_id, str) and stored_id
else _NEGATIVE_STORED_ID_SENTINEL
),
)
return owner
async def _get_stored_container_id(
original_container_id: str, custom_llm_provider: str
) -> Optional[str]:
"""Return the ``unified_object_id`` stored at create time, if any.
Used by :func:`get_container_forwarding_params` to recover the
deployment ``model_id`` for native upstream container IDs: the stored
value is the encoded form produced by ``encode_container_id_in_response``
when the router selected a specific deployment.
"""
model_object_id = _container_model_object_id(
original_container_id, custom_llm_provider
)
cached = _CONTAINER_STORED_ID_CACHE.get_cache(model_object_id)
if cached == _NEGATIVE_STORED_ID_SENTINEL:
return None
if isinstance(cached, str) and cached:
return cached
prisma_client = await _get_prisma_client()
if prisma_client is None:
return None
row = await prisma_client.db.litellm_managedobjecttable.find_first(
where={
"model_object_id": model_object_id,
"file_purpose": CONTAINER_OBJECT_PURPOSE,
}
)
stored_id = getattr(row, "unified_object_id", None) if row is not None else None
_CONTAINER_STORED_ID_CACHE.set_cache(
model_object_id,
(
stored_id
if isinstance(stored_id, str) and stored_id
else _NEGATIVE_STORED_ID_SENTINEL
),
)
return stored_id if isinstance(stored_id, str) and stored_id else None
async def assert_user_can_access_container(
container_id: str,
user_api_key_dict: UserAPIKeyAuth,

View file

@ -178,15 +178,28 @@ class SpendCounterReseed:
if db_spend is None:
return None
# Warm even when 0 so subsequent reads hit cache, not DB.
#
# Seed via SET NX (cross-pod safe): only one pod initializes the
# Redis key with db_spend; concurrent seeders read the winner's
# value. INCRBYFLOAT-of-db_spend from N pods would multiply the
# counter (N x db_spend) and trigger spurious budget alerts.
current_value: float = float(db_spend)
try:
if spend_counter_cache.redis_cache is not None:
current_value = (
await spend_counter_cache.redis_cache.async_increment(
key=counter_key,
value=db_spend,
refresh_ttl=True,
)
seeded = await spend_counter_cache.redis_cache.async_set_cache(
key=counter_key,
value=db_spend,
nx=True,
)
if seeded:
current_value = float(db_spend)
else:
cached = await spend_counter_cache.redis_cache.async_get_cache(
key=counter_key
)
current_value = (
float(cached) if cached is not None else float(db_spend)
)
spend_counter_cache.in_memory_cache.set_cache(
key=counter_key,
value=current_value,
@ -202,7 +215,7 @@ class SpendCounterReseed:
)
if require_cache_warm:
raise
return db_spend
return current_value
@staticmethod
async def window_from_spend_logs(

View file

@ -92,6 +92,8 @@ from litellm.types.utils import CallTypesLiteral
# Module-level singleton for the JWKS discovery endpoint to access.
_mcp_jwt_signer_instance: Optional["MCPJWTSigner"] = None
_MCP_JWT_CALL_TYPES = frozenset({"call_mcp_tool", "list_mcp_tools"})
# Simple in-memory JWKS cache: keyed by JWKS URI → (keys_list, fetched_at).
_jwks_cache: Dict[str, tuple] = {}
_JWKS_CACHE_TTL = 3600 # 1 hour
@ -603,17 +605,23 @@ class MCPJWTSigner(CustomGuardrail):
# FR-10: Scope building
# ------------------------------------------------------------------
def _build_scope(self, raw_tool_name: str) -> str:
def _build_scope(
self,
raw_tool_name: str,
call_type: Optional[CallTypesLiteral] = None,
) -> str:
"""
Build the JWT scope string.
When allowed_scopes is configured: join them verbatim.
Otherwise auto-generate minimal, least-privilege scopes:
- Tool call → mcp:tools/call mcp:tools/<name>:call
- No tool → mcp:tools/call mcp:tools/list
- No tool → mcp:tools/list
NOTE: tools/list is intentionally NOT granted on tool-call JWTs to
prevent callers from enumerating tools they didn't ask to use.
Conversely, tools/call is NOT granted on tools/list-only JWTs so an
intercepted list token cannot be replayed to invoke tools.
"""
if self.allowed_scopes is not None:
return " ".join(self.allowed_scopes)
@ -623,8 +631,14 @@ class MCPJWTSigner(CustomGuardrail):
)
if tool_name:
scopes = ["mcp:tools/call", f"mcp:tools/{tool_name}:call"]
elif call_type == "call_mcp_tool":
# Tool-call request reached the signer without a tool name (e.g.
# missing mcp_tool_name in hook data). Fall back to a generic
# tools/call scope so the upstream server still accepts the
# invocation rather than rejecting it as a tools/list-only token.
scopes = ["mcp:tools/call"]
else:
scopes = ["mcp:tools/call", "mcp:tools/list"]
scopes = ["mcp:tools/list"]
return " ".join(scopes)
# ------------------------------------------------------------------
@ -673,6 +687,7 @@ class MCPJWTSigner(CustomGuardrail):
user_api_key_dict: UserAPIKeyAuth,
data: dict,
jwt_claims: Optional[Dict[str, Any]] = None,
call_type: Optional[CallTypesLiteral] = None,
) -> Dict[str, Any]:
"""
Build JWT claims for the outbound MCP access token.
@ -713,7 +728,7 @@ class MCPJWTSigner(CustomGuardrail):
# scope (FR-10)
raw_tool_name: str = data.get("mcp_tool_name", "")
claims["scope"] = self._build_scope(raw_tool_name)
claims["scope"] = self._build_scope(raw_tool_name, call_type=call_type)
# optional_claims passthrough (FR-15)
claims = self._passthrough_optional_claims(claims, jwt_claims)
@ -779,16 +794,20 @@ class MCPJWTSigner(CustomGuardrail):
Verifies the incoming token (when configured), validates required claims,
then signs an outbound JWT and injects it as the Authorization header.
All non-MCP call types pass through unchanged.
Signs outbound MCP tool calls and tools/list requests.
"""
if call_type != "call_mcp_tool":
if call_type not in _MCP_JWT_CALL_TYPES:
return data
hook_data = dict(data)
if call_type == "list_mcp_tools":
hook_data["mcp_tool_name"] = ""
# ------------------------------------------------------------------
# FR-5: Verify incoming token before re-signing
# ------------------------------------------------------------------
jwt_claims: Optional[Dict[str, Any]] = None
raw_token: Optional[str] = data.get("incoming_bearer_token")
raw_token: Optional[str] = hook_data.get("incoming_bearer_token")
if self.access_token_discovery_uri and raw_token:
# Three-dot pattern → JWT; otherwise opaque.
@ -837,7 +856,9 @@ class MCPJWTSigner(CustomGuardrail):
# ------------------------------------------------------------------
# Build outbound access token
# ------------------------------------------------------------------
claims = self._build_claims(user_api_key_dict, data, jwt_claims)
claims = self._build_claims(
user_api_key_dict, hook_data, jwt_claims, call_type=call_type
)
signed_token = jwt.encode(
claims,
@ -848,7 +869,7 @@ class MCPJWTSigner(CustomGuardrail):
# Merge into existing extra_headers — a prior guardrail in the chain may
# have already injected tracing headers or correlation IDs.
existing_headers: Dict[str, str] = data.get("extra_headers") or {}
existing_headers: Dict[str, str] = hook_data.get("extra_headers") or {}
new_headers: Dict[str, str] = {
**existing_headers,
"Authorization": f"Bearer {signed_token}",
@ -875,17 +896,74 @@ class MCPJWTSigner(CustomGuardrail):
claims, self._kid
)
data["extra_headers"] = new_headers
hook_data["extra_headers"] = new_headers
verbose_proxy_logger.debug(
"MCPJWTSigner: signed JWT sub=%s act=%s tool=%s exp=%d "
"verified=%s channel=%s",
"verified=%s channel=%s call_type=%s",
claims.get("sub"),
claims.get("act", {}).get("sub"),
data.get("mcp_tool_name"),
hook_data.get("mcp_tool_name"),
claims["exp"],
jwt_claims is not None,
bool(self.channel_token_audience),
call_type,
)
return data
return hook_data
async def inject_mcp_jwt_headers_for_upstream(
user_api_key_dict: Optional[UserAPIKeyAuth],
extra_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
*,
for_list_tools: bool = False,
mcp_tool_name: str = "",
) -> Dict[str, str]:
"""
Sign outbound MCP headers when MCPJWTSigner is configured.
Used by tools/list paths that do not go through proxy pre_call_hook.
"""
merged = dict(extra_headers or {})
signer = get_mcp_jwt_signer()
if signer is None or user_api_key_dict is None:
return merged
normalized_raw = {k.lower(): v for k, v in (raw_headers or {}).items()}
incoming_bearer_token: Optional[str] = None
auth_hdr = normalized_raw.get("authorization", "")
if auth_hdr.lower().startswith("bearer "):
incoming_bearer_token = auth_hdr[len("bearer ") :]
hook_data: Dict[str, Any] = {
"mcp_tool_name": "" if for_list_tools else mcp_tool_name,
"incoming_bearer_token": incoming_bearer_token,
"extra_headers": merged,
}
call_type: CallTypesLiteral = (
"list_mcp_tools" if for_list_tools else "call_mcp_tool"
)
try:
from litellm.proxy.proxy_server import ( # noqa: PLC0415
proxy_logging_obj as _proxy_logging,
)
shared_cache = (
_proxy_logging.internal_usage_cache.dual_cache
if _proxy_logging is not None
else DualCache()
)
except Exception:
shared_cache = DualCache()
result = await signer.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=shared_cache,
data=hook_data,
call_type=call_type,
)
if isinstance(result, dict) and result.get("extra_headers"):
merged.update(result["extra_headers"])
return merged

View file

@ -1798,7 +1798,10 @@ async def cli_sso_callback(
from fastapi.responses import HTMLResponse
verify_url = str(request.url_for("cli_sso_complete", login_id=key))
verify_url = get_custom_url(
request_base_url=str(request.base_url),
route=f"sso/cli/complete/{key}",
)
html_content = _render_cli_sso_verification_page(
verify_url=verify_url,
browser_complete_token=browser_complete_token,

View file

@ -727,6 +727,76 @@ async def resolve_output_file_ids_to_unified(response, prisma_client) -> None:
pass
async def ensure_batch_response_managed_file_ids(
response,
managed_files_obj,
prisma_client,
verbose_proxy_logger,
user_api_key_dict=None,
db_batch_object=None,
) -> None:
"""Normalize batch file IDs to managed unified IDs before DB persistence."""
await resolve_input_file_id_to_unified(response, prisma_client)
await resolve_output_file_ids_to_unified(response, prisma_client)
if managed_files_obj is None:
return
hidden_params = getattr(response, "_hidden_params", None) or {}
model_id = hidden_params.get("model_id")
if not model_id:
return
model_name = hidden_params.get("model_name")
unified_file_id = hidden_params.get("unified_file_id")
if not model_name and isinstance(unified_file_id, str):
decoded_unified_file_id = (
_is_base64_encoded_unified_file_id(unified_file_id) or unified_file_id
)
target_model_names = get_models_from_unified_file_id(decoded_unified_file_id)
if target_model_names:
model_name = ",".join(target_model_names)
if user_api_key_dict is None and db_batch_object is not None:
from litellm.proxy._types import UserAPIKeyAuth
user_api_key_dict = UserAPIKeyAuth(
user_id=getattr(db_batch_object, "created_by", None) or "default-user-id",
team_id=getattr(db_batch_object, "team_id", None),
)
if user_api_key_dict is None:
return
for file_attr in ("output_file_id", "error_file_id"):
raw_file_id = getattr(response, file_attr, None)
if not raw_file_id or _is_base64_encoded_unified_file_id(raw_file_id):
continue
try:
new_unified_file_id = managed_files_obj.get_unified_output_file_id(
output_file_id=raw_file_id,
model_id=model_id,
model_name=model_name,
)
await managed_files_obj.store_unified_file_id(
file_id=new_unified_file_id,
file_object=None,
litellm_parent_otel_span=getattr(
user_api_key_dict, "parent_otel_span", None
),
model_mappings={model_id: raw_file_id},
user_api_key_dict=user_api_key_dict,
)
setattr(response, file_attr, new_unified_file_id)
verbose_proxy_logger.debug(
f"Converted batch {file_attr} {raw_file_id!r} to managed ID before DB write"
)
except Exception as e:
verbose_proxy_logger.warning(
f"Failed to convert batch {file_attr}={raw_file_id!r} to managed ID "
f"before DB write: {e}"
)
async def get_batch_from_database(
batch_id: str,
unified_batch_id: Union[str, Literal[False]],
@ -800,6 +870,7 @@ async def update_batch_in_database(
verbose_proxy_logger,
db_batch_object=None,
operation: str = "update",
user_api_key_dict=None,
):
"""
Update batch status and object in ManagedObjectTable.
@ -813,6 +884,7 @@ async def update_batch_in_database(
verbose_proxy_logger: Logger instance
db_batch_object: Optional existing database object (for comparison)
operation: Description of operation ("update", "cancel", etc.)
user_api_key_dict: Optional auth context for creating managed file IDs
"""
import litellm.utils
@ -823,6 +895,18 @@ async def update_batch_in_database(
if not prisma_client:
return
# Always normalize the response's file IDs to unified managed IDs
# (mutates in place) so the caller returns unified IDs to the user
# even when we skip the DB update below for an unchanged status.
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=managed_files_obj,
prisma_client=prisma_client,
verbose_proxy_logger=verbose_proxy_logger,
user_api_key_dict=user_api_key_dict,
db_batch_object=db_batch_object,
)
# Only update if status has changed (when db_batch_object is provided)
if db_batch_object and response.status == db_batch_object.status:
return

View file

@ -38,6 +38,35 @@ class LiteLLMDatabaseConnectionPool(Enum):
database_connection_pool_timeout = 60
def _build_db_connection_url_params(
connection_limit: int,
pool_timeout: Optional[Union[int, float]],
connect_timeout: Optional[Union[int, float]] = None,
socket_timeout: Optional[Union[int, float]] = None,
extra_params: Optional[dict] = None,
) -> dict:
"""Build the Prisma DATABASE_URL query params controlling connection pool behavior.
`connect_timeout` / `socket_timeout` map to the Prisma URL params of the same
name (https://www.prisma.io/docs/orm/overview/databases/postgresql) and are
omitted when None so Prisma's defaults apply. `extra_params` is an
untyped passthrough — keys it provides win over the named arguments above,
so it can be used to override any default we set here.
"""
params: dict = {
"connection_limit": connection_limit,
}
if pool_timeout is not None:
params["pool_timeout"] = pool_timeout
if connect_timeout is not None:
params["connect_timeout"] = connect_timeout
if socket_timeout is not None:
params["socket_timeout"] = socket_timeout
if extra_params:
params.update(extra_params)
return params
def append_query_params(url: Optional[str], params: dict) -> str:
from litellm._logging import verbose_proxy_logger
@ -807,6 +836,9 @@ def run_server( # noqa: PLR0915
db_connection_pool_limit = 100
# Starts optional due to config fallback checks; guaranteed non-None before use.
db_connection_timeout: Optional[Union[int, float]] = 60
db_connect_timeout: Optional[Union[int, float]] = None
db_socket_timeout: Optional[Union[int, float]] = None
db_extra_connection_params: Optional[dict] = None
general_settings = {}
### GET DB TOKEN FOR IAM AUTH ###
@ -924,6 +956,11 @@ def run_server( # noqa: PLR0915
db_connection_timeout = (
LiteLLMDatabaseConnectionPool.database_connection_pool_timeout.value
)
db_connect_timeout = general_settings.get("database_connect_timeout")
db_socket_timeout = general_settings.get("database_socket_timeout")
db_extra_connection_params = general_settings.get(
"database_extra_connection_params"
)
if database_url and database_url.startswith("os.environ/"):
original_dir = os.getcwd()
# set the working directory to where this script is
@ -963,27 +1000,26 @@ def run_server( # noqa: PLR0915
try:
from litellm.secret_managers.main import get_secret
connection_url_params = _build_db_connection_url_params(
connection_limit=db_connection_pool_limit,
pool_timeout=db_connection_timeout,
connect_timeout=db_connect_timeout,
socket_timeout=db_socket_timeout,
extra_params=db_extra_connection_params,
)
if os.getenv("DATABASE_URL", None) is not None:
### add connection pool + pool timeout args
params = {
"connection_limit": db_connection_pool_limit,
"pool_timeout": db_connection_timeout,
}
database_url = get_secret("DATABASE_URL", default_value=None)
modified_url = append_query_params(
str(database_url) if database_url else None, params
str(database_url) if database_url else None,
connection_url_params,
)
os.environ["DATABASE_URL"] = modified_url
if os.getenv("DIRECT_URL", None) is not None:
### add connection pool + pool timeout args
params = {
"connection_limit": db_connection_pool_limit,
"pool_timeout": db_connection_timeout,
}
database_url = os.getenv("DIRECT_URL")
modified_url = append_query_params(database_url, params)
modified_url = append_query_params(
database_url, connection_url_params
)
os.environ["DIRECT_URL"] = modified_url
###
subprocess.run(["prisma"], capture_output=True)
is_prisma_runnable = True
except FileNotFoundError:

View file

@ -4328,6 +4328,19 @@ class ProxyConfig:
"health_check_concurrency", None
)
health_check_details = general_settings.get("health_check_details", True)
### INTERACTIONS API SCHEMA ###
_use_legacy_interactions_schema = general_settings.get(
"use_legacy_interactions_schema"
)
if _use_legacy_interactions_schema is not None:
if isinstance(_use_legacy_interactions_schema, str):
litellm.use_legacy_interactions_schema = (
_use_legacy_interactions_schema.lower() == "true"
)
else:
litellm.use_legacy_interactions_schema = bool(
_use_legacy_interactions_schema
)
# Health-check-driven routing (opt-in, passes through to Router later)
_enable_hc_routing = general_settings.get(
"enable_health_check_routing", False

View file

@ -65,8 +65,7 @@ class LiteLLMCompletionTransformationHandler:
litellm_completion_response: Union[
ModelResponse, litellm.CustomStreamWrapper
] = litellm.completion(
**litellm_completion_request,
**kwargs,
**completion_args,
)
if isinstance(litellm_completion_response, ModelResponse):

View file

@ -1115,6 +1115,7 @@ def responses(
stream=stream,
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout if timeout is not None else request_timeout,
**kwargs,
)

View file

@ -208,6 +208,15 @@ if TYPE_CHECKING:
from litellm.router_strategy.quality_router.quality_router import (
QualityRouter,
)
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
)
from litellm.types.llms.base import BaseLiteLLMOpenAIResponseObject
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseInputParam,
ResponsesAPIResponse,
)
Span = Union[_Span, Any]
else:
@ -2246,6 +2255,388 @@ class Router:
return FallbackStreamWrapper(stream_with_fallbacks())
@staticmethod
def _extract_partial_responses_usage(
source_iterator: "BaseResponsesAPIStreamingIterator",
) -> Optional["ResponseAPIUsage"]:
"""
Best-effort: pull partial token usage from a Responses-API streaming
iterator that errored mid-stream, normalized to ResponseAPIUsage so
the caller can combine without crossing token-naming conventions.
Two sources, in priority order:
1. The bridge path (LiteLLMCompletionStreamingIterator) accumulates
chat-completion chunks while streaming — feed them through
stream_chunk_builder to recover chat Usage, then translate
(prompt_tokens → input_tokens, completion_tokens → output_tokens).
2. The native path (ResponsesAPIStreamingIterator) only has a
completed_response object if the stream reached
RESPONSE_COMPLETED before erroring — uncommon mid-stream but
worth checking. Already ResponseAPIUsage-shaped.
Returns None when no partial usage is recoverable.
"""
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
LiteLLMCompletionStreamingIterator,
)
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponseFailedEvent,
ResponseIncompleteEvent,
)
# Bridge subclass is the only iterator that accumulates chat-completion
# chunks. isinstance narrows the type so we can read the attribute
# directly instead of getattr-ing on the base class.
if isinstance(source_iterator, LiteLLMCompletionStreamingIterator):
chunks = source_iterator.collected_chat_completion_chunks
if chunks:
try:
from litellm.main import stream_chunk_builder
built = stream_chunk_builder(chunks=chunks)
# stream_chunk_builder returns ModelResponse |
# TextCompletionResponse | None. ModelResponse sets .usage
# in __init__ rather than declaring it as a class field, so
# static narrowing doesn't expose it. Mirror the sync path
# (_completion_streaming_iterator) and pull via getattr.
chat = getattr(built, "usage", None) if built is not None else None
if chat is not None:
# getattr-with-default because the test path may
# substitute a SimpleNamespace lacking some fields;
# real Usage instances always have them.
prompt = int(getattr(chat, "prompt_tokens", 0) or 0)
completion = int(getattr(chat, "completion_tokens", 0) or 0)
total = int(
getattr(chat, "total_tokens", prompt + completion)
or (prompt + completion)
)
return ResponseAPIUsage(
input_tokens=prompt,
output_tokens=completion,
total_tokens=total,
)
except Exception:
# Builder is best-effort — fall through to native path.
pass
# Native path: completed_response is set only if RESPONSE_COMPLETED
# arrived before the error (uncommon mid-stream but worth checking).
# Already ResponseAPIUsage-shaped — return as-is.
completed = source_iterator.completed_response
if isinstance(
completed,
(ResponseCompletedEvent, ResponseFailedEvent, ResponseIncompleteEvent),
):
return completed.response.usage
return None
@staticmethod
def _combine_responses_fallback_usage(
fallback_item: "BaseLiteLLMOpenAIResponseObject",
partial_usage: "ResponseAPIUsage",
) -> None:
"""
Merge partial-stream usage with fallback-stream usage on a
Responses-API streaming event.
Only mutates events that carry a `response` with a `usage` field
(response.completed / response.failed / response.incomplete). Other
events pass through unchanged.
Both inputs are ResponseAPIUsage-shaped (see
_extract_partial_responses_usage which normalizes the bridge path),
so we can sum input_tokens / output_tokens / total_tokens directly
and produce a clean ResponseAPIUsage — no token-naming split, no
setattr bypass.
"""
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponseFailedEvent,
ResponseIncompleteEvent,
)
if not isinstance(
fallback_item,
(ResponseCompletedEvent, ResponseFailedEvent, ResponseIncompleteEvent),
):
return
response = fallback_item.response
if response.usage is None:
return
fb = response.usage
response.usage = ResponseAPIUsage(
input_tokens=(partial_usage.input_tokens or 0) + (fb.input_tokens or 0),
output_tokens=(partial_usage.output_tokens or 0) + (fb.output_tokens or 0),
total_tokens=(partial_usage.total_tokens or 0) + (fb.total_tokens or 0),
)
@staticmethod
def _build_responses_continuation_input(
input_val: Optional[Union[str, "ResponseInputParam"]],
generated_content: str,
) -> "ResponseInputParam":
"""
Convert Responses-API input + partial assistant output into a
continuation input that asks the fallback model to pick up where the
prior assistant message stopped.
Best effort across providers. The chat-completions path uses
Anthropic's `prefix: True` prefill trick on the assistant message;
the Responses-API input schema has no direct equivalent, so we
append an instruction (developer role) plus a prior assistant
message containing the partial output. Providers without prefill
semantics (OpenAI, Vertex) treat this as conversational context
and may regenerate — same trade-off as the chat-completions path
for non-Anthropic fallbacks.
"""
# base/continuation are List[Any] because ResponseInputParam items
# are a wide Union of TypedDicts (EasyInputMessageParam, Message,
# ResponseOutputMessageParam, ...) — annotating as List[Dict[str, Any]]
# rejects the list() spread of input_val. We cast the combined list to
# ResponseInputParam at the return.
base: List[Any]
if isinstance(input_val, str):
base = [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": input_val}],
}
]
elif isinstance(input_val, list):
base = list(input_val)
else:
base = []
continuation: List[Any] = [
{
"type": "message",
"role": "developer",
"content": [
{
"type": "input_text",
"text": (
"The previous assistant response was interrupted "
"mid-stream. Continue exactly where it stopped — "
"do not repeat any of its content. Your response "
"must read as a seamless continuation."
),
}
],
},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": generated_content}],
},
]
return cast("ResponseInputParam", base + continuation)
async def _aresponses_streaming_iterator(
self,
response: "BaseResponsesAPIStreamingIterator",
initial_kwargs: Dict[str, Any],
) -> "BaseResponsesAPIStreamingIterator":
"""
Wrap a Responses-API streaming iterator so MidStreamFallbackError
triggers the Router's fallback chain (parity with
_acompletion_streaming_iterator for the chat-completions path).
The Responses-API streaming path goes through
_ageneric_api_call_with_fallbacks rather than _acompletion, so the
returned iterator is never wrapped by the chat completions
fallback handler. Without this wrapper, MidStreamFallbackError
raised mid-stream from the underlying CustomStreamWrapper (used by
LiteLLMCompletionStreamingIterator when the Responses API is
served via the completion bridge) propagates unhandled and the
configured cross-provider fallback never fires.
Full parity with the chat-completions path:
- Pre-first-chunk: retry with the original input unchanged.
- Partial content: inject a developer instruction + prior
assistant message carrying the generated text so the fallback
model continues rather than restarts.
- Usage combining: merge partial-stream usage onto the fallback's
response.completed event so accounting reflects both attempts.
- Stream cleanup: shielded aclose() on both source and fallback
iterators on terminate.
"""
from litellm.exceptions import MidStreamFallbackError
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
)
source_iterator = response
class FallbackResponsesStreamWrapper(BaseResponsesAPIStreamingIterator):
"""
Subclasses BaseResponsesAPIStreamingIterator only for isinstance
compatibility (proxy + interactions code paths check the type).
Bypasses the parent constructor and delegates iteration to an
async generator.
"""
def __init__(self, async_generator: AsyncGenerator):
import time
from datetime import datetime
self._async_generator = async_generator
# Mirror every attribute BaseResponsesAPIStreamingIterator.__init__
# would have set. The wrapper bypasses super().__init__ (it has no
# httpx.Response of its own and no provider config to drive), so
# we copy from source_iterator where applicable and use safe
# defaults elsewhere. This keeps inherited methods (e.g.
# _check_max_streaming_duration, _handle_failure) safe to call.
#
# The bridge path (LiteLLMCompletionStreamingIterator used by
# Anthropic/Bedrock/Vertex) does not call super().__init__ and
# is missing many of these attributes — use getattr fallbacks
# so wrapper construction never raises AttributeError. The
# bridge stores the logging object as `litellm_logging_obj`.
self.response = getattr(source_iterator, "response", None)
self.model = getattr(source_iterator, "model", None)
self.logging_obj = getattr(
source_iterator,
"logging_obj",
getattr(source_iterator, "litellm_logging_obj", None),
)
self.finished = False
self.responses_api_provider_config = getattr(
source_iterator, "responses_api_provider_config", None
)
self.completed_response = None
self.start_time = getattr(source_iterator, "start_time", datetime.now())
self._failure_handled = False
self._completed_response_cached = False
self._completed_response_logged = False
self._completed_response_cache_hit = None
self._persist_completed_response_before_logging = True
self._stream_created_time = time.time()
self.litellm_metadata = getattr(
source_iterator, "litellm_metadata", None
)
self.custom_llm_provider = getattr(
source_iterator, "custom_llm_provider", None
)
self.request_data = getattr(source_iterator, "request_data", {}) or {}
self.call_type = getattr(source_iterator, "call_type", None)
# Preserve hidden params so response headers (model_id,
# api_base, additional_headers) keep flowing.
self._hidden_params = dict(
getattr(source_iterator, "_hidden_params", None) or {}
)
def __aiter__(self):
return self
async def __anext__(self):
return await self._async_generator.__anext__()
async def aclose(self):
# async generators always expose aclose — no defensive check needed.
await self._async_generator.aclose()
async def stream_with_fallbacks():
fallback_response = None
try:
async for item in source_iterator:
yield item
except MidStreamFallbackError as e:
partial_usage = Router._extract_partial_responses_usage(source_iterator)
try:
model_group = cast(str, initial_kwargs.get("model"))
fallbacks: Optional[List] = initial_kwargs.get(
"fallbacks", self.fallbacks
)
context_window_fallbacks: Optional[List] = initial_kwargs.get(
"context_window_fallbacks", self.context_window_fallbacks
)
content_policy_fallbacks: Optional[List] = initial_kwargs.get(
"content_policy_fallbacks", self.content_policy_fallbacks
)
# Re-enter via the per-attempt helper so the fallback chain
# picks deployments through
# _ageneric_api_call_with_fallbacks_helper.
# original_generic_function is preserved by the caller so
# the helper knows what underlying API to invoke per attempt.
initial_kwargs["original_function"] = (
self._ageneric_api_call_with_fallbacks_helper
)
if e.is_pre_first_chunk or not e.generated_content:
# No content generated before the error — retry with the
# original input. Adding a continuation prompt would
# waste tokens and confuse the model.
pass
else:
initial_kwargs["input"] = (
Router._build_responses_continuation_input(
initial_kwargs.get("input"),
e.generated_content,
)
)
# The Responses-API path stores observability metadata
# under "litellm_metadata" (not the default "metadata") —
# see _ageneric_api_call_with_fallbacks. Mirroring that
# here ensures model_group, model_group_alias, and trace
# ids land in the same key litellm.aresponses reads from.
self._update_kwargs_before_fallbacks(
model=model_group,
kwargs=initial_kwargs,
metadata_variable_name="litellm_metadata",
)
fallback_response = (
await self.async_function_with_fallbacks_common_utils(
e=e,
disable_fallbacks=False,
fallbacks=fallbacks,
context_window_fallbacks=context_window_fallbacks,
content_policy_fallbacks=content_policy_fallbacks,
model_group=model_group,
args=(),
kwargs=initial_kwargs,
)
)
if hasattr(fallback_response, "__aiter__"):
async for fallback_item in fallback_response: # type: ignore
if partial_usage is not None:
Router._combine_responses_fallback_usage(
fallback_item, partial_usage
)
yield fallback_item
else:
yield fallback_response
except Exception as fallback_error:
verbose_router_logger.error(
f"Responses streaming fallback also failed: {fallback_error}"
)
raise fallback_error
finally:
with anyio.CancelScope(shield=True):
if hasattr(source_iterator, "aclose"):
try:
await source_iterator.aclose() # type: ignore[func-returns-value]
except BaseException as exc:
verbose_router_logger.debug(
"stream_with_fallbacks(aresponses): error closing source: %s",
exc,
)
if fallback_response is not None and hasattr(
fallback_response, "aclose"
):
try:
await fallback_response.aclose()
except BaseException as exc:
verbose_router_logger.debug(
"stream_with_fallbacks(aresponses): error closing fallback: %s",
exc,
)
return FallbackResponsesStreamWrapper(stream_with_fallbacks())
def _completion_streaming_iterator( # noqa: PLR0915
self,
model_response: CustomStreamWrapper,
@ -4292,6 +4683,61 @@ class Router:
self.fail_calls[model] += 1
raise e
async def _aresponses_with_streaming_fallbacks(
self, original_function: Callable, **kwargs: Any
) -> Union["ResponsesAPIResponse", "BaseResponsesAPIStreamingIterator"]:
"""
_ageneric_api_call_with_fallbacks for the Responses API, with the
addition of mid-stream fallback handling.
When stream=True and the underlying call returns a
BaseResponsesAPIStreamingIterator, wrap it with
_aresponses_streaming_iterator so MidStreamFallbackError raised
during iteration triggers the Router's cross-provider fallback chain.
"""
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
)
from litellm.litellm_core_utils.core_helpers import safe_deep_copy
# Snapshot the request kwargs before _ageneric_api_call_with_fallbacks
# mutates them. A shallow copy alone is not enough: the primary
# attempt mutates nested dicts in place — notably `litellm_metadata`,
# which `_update_kwargs_with_deployment` populates with
# deployment-specific fields (`deployment`, `model_info`, `api_base`,
# tags, etc.). Without an explicit copy of that dict, the shallow
# copy would still share its reference, leaking primary-deployment
# metadata into the mid-stream fallback request.
#
# We avoid deep-copying the full kwargs because it can contain
# non-deepcopyable objects (logging handles, async clients, etc.);
# `safe_deep_copy` deep-copies the metadata dicts key-by-key with a
# fallback to the original reference for any non-picklable value.
# The original_generic_function is preserved so the per-attempt
# helper knows which underlying API to call on fallback.
fallback_kwargs: Dict[str, Any] = kwargs.copy()
if isinstance(fallback_kwargs.get("litellm_metadata"), dict):
fallback_kwargs["litellm_metadata"] = safe_deep_copy(
fallback_kwargs["litellm_metadata"]
)
if isinstance(fallback_kwargs.get("metadata"), dict):
fallback_kwargs["metadata"] = safe_deep_copy(fallback_kwargs["metadata"])
fallback_kwargs["original_generic_function"] = original_function
response = await self._ageneric_api_call_with_fallbacks(
original_function=original_function, **kwargs
)
if kwargs.get("stream") and isinstance(
response, BaseResponsesAPIStreamingIterator
):
return await self._aresponses_streaming_iterator(
response=response,
initial_kwargs=fallback_kwargs,
)
return response
def _generic_api_call_with_fallbacks(
self, model: str, original_function: Callable, **kwargs
):
@ -5511,9 +5957,13 @@ class Router:
custom_llm_provider=custom_llm_provider,
**kwargs,
)
elif call_type == "aresponses":
return await self._aresponses_with_streaming_fallbacks(
original_function=original_function,
**kwargs,
)
elif call_type in (
"anthropic_messages",
"aresponses",
"_arealtime",
"_aresponses_websocket",
"acreate_fine_tuning_job",
@ -5670,6 +6120,7 @@ class Router:
from litellm.responses.utils import ResponsesAPIRequestUtils
container_id = kwargs.get("container_id")
_forwarded_model_id = kwargs.get("model_id")
if isinstance(container_id, str):
decoded = ResponsesAPIRequestUtils._decode_container_id(container_id)
original_id = decoded.get("response_id", container_id)
@ -5678,7 +6129,14 @@ class Router:
decoded_provider = decoded.get("custom_llm_provider")
if decoded_provider and kwargs.get("custom_llm_provider") == "openai":
kwargs["custom_llm_provider"] = decoded_provider
model_id = decoded.get("model_id")
# Fall back to the model_id forwarded by the proxy when the container_id
# is a native upstream ID (e.g. Azure hex cntr_) that carries no LiteLLM
# routing payload, so deployment credentials (api_base, api_key) are applied.
model_id = decoded.get("model_id") or (
_forwarded_model_id.strip()
if isinstance(_forwarded_model_id, str) and _forwarded_model_id.strip()
else None
)
if model_id:
kwargs["model"] = model_id
return await self._ageneric_api_call_with_fallbacks(

View file

@ -36,9 +36,13 @@ from litellm.types.interactions.generated import (
GoogleSearchResultContent,
ImageContent,
Interaction,
InteractionCompleted,
InteractionCreated,
InteractionEvent,
InteractionEnvironment,
InteractionInProgress,
InteractionInput,
InteractionRequiresAction,
InteractionsAPIOptionalRequestParams,
InteractionsAPIResponse,
InteractionsAPIStreamingResponse,
@ -50,6 +54,9 @@ from litellm.types.interactions.generated import (
McpServerToolResultContent,
ModelOption,
ResponseModality,
StepDelta,
StepStart,
StepStop,
)
from litellm.types.interactions.generated import (
Status3 as InteractionStatus, # Main request/response types; Content types; Turn for multi-turn conversations; Tool types; Config types; Usage; Status enum; Events for streaming; Agent configs; Model/Agent options; Response modality; Annotation; LiteLLM types; Backwards compat aliases
@ -115,6 +122,14 @@ __all__ = [
"AgentOption",
"ResponseModality",
"Annotation",
# New schema SSE event types (Api-Revision: 2026-05-20)
"StepStart",
"StepDelta",
"StepStop",
"InteractionCreated",
"InteractionInProgress",
"InteractionCompleted",
"InteractionRequiresAction",
# LiteLLM types
"InteractionEnvironment",
"InteractionInput",

View file

@ -1151,9 +1151,114 @@ class InteractionEvent(BaseModel):
)
# ---------------------------------------------------------------
# New schema SSE event types (Api-Revision: 2026-05-20)
# These replace the legacy content.* / interaction.start|complete
# events and will become the only events after June 8, 2026.
# ---------------------------------------------------------------
class StepStart(BaseModel):
"""Emitted when a new step begins (replaces content.start)."""
event_type: Literal["step.start"] = "step.start"
index: Optional[int] = None
step: Optional[Dict[str, Any]] = Field(
None,
description="The initial step data (type, content, signature, etc.).",
)
event_id: Optional[str] = Field(
None,
description="The event_id token to be used to resume the interaction stream.",
)
class StepDelta(BaseModel):
"""Emitted for incremental step content (replaces content.delta)."""
event_type: Literal["step.delta"] = "step.delta"
index: Optional[int] = None
delta: Optional[Dict[str, Any]] = Field(
None,
description="Incremental content delta (e.g. text, arguments_delta for function calls).",
)
event_id: Optional[str] = Field(
None,
description="The event_id token to be used to resume the interaction stream.",
)
class StepStop(BaseModel):
"""Emitted when a step finishes (replaces content.stop)."""
event_type: Literal["step.stop"] = "step.stop"
index: Optional[int] = None
status: Optional[str] = Field(
None,
description="Step completion status (e.g. 'done').",
)
event_id: Optional[str] = Field(
None,
description="The event_id token to be used to resume the interaction stream.",
)
class InteractionCreated(BaseModel):
"""Emitted when the interaction is first created (replaces interaction.start)."""
event_type: Literal["interaction.created"] = "interaction.created"
interaction: Optional[Dict[str, Any]] = None
event_id: Optional[str] = Field(
None,
description="The event_id token to be used to resume the interaction stream.",
)
class InteractionInProgress(BaseModel):
"""Emitted while the interaction is running."""
event_type: Literal["interaction.in_progress"] = "interaction.in_progress"
interaction_id: Optional[str] = None
event_id: Optional[str] = Field(
None,
description="The event_id token to be used to resume the interaction stream.",
)
class InteractionCompleted(BaseModel):
"""Emitted when the interaction finishes (replaces interaction.complete)."""
event_type: Literal["interaction.completed"] = "interaction.completed"
interaction: Optional[Dict[str, Any]] = None
event_id: Optional[str] = Field(
None,
description="The event_id token to be used to resume the interaction stream.",
)
class InteractionRequiresAction(BaseModel):
"""Emitted when the interaction is paused waiting for a tool result."""
event_type: Literal["interaction.requires_action"] = "interaction.requires_action"
interaction_id: Optional[str] = None
event_id: Optional[str] = Field(
None,
description="The event_id token to be used to resume the interaction stream.",
)
class InteractionSseEvent(
RootModel[
Union[
# New schema events (Api-Revision: 2026-05-20)
StepStart,
StepDelta,
StepStop,
InteractionCreated,
InteractionInProgress,
InteractionCompleted,
InteractionRequiresAction,
# Legacy schema events (Api-Revision: 2026-05-07, removed June 8 2026)
InteractionEvent,
InteractionStatusUpdate,
ContentStart,
@ -1164,6 +1269,15 @@ class InteractionSseEvent(
]
):
root: Union[
# New schema events (Api-Revision: 2026-05-20)
StepStart,
StepDelta,
StepStop,
InteractionCreated,
InteractionInProgress,
InteractionCompleted,
InteractionRequiresAction,
# Legacy schema events (Api-Revision: 2026-05-07, removed June 8 2026)
InteractionEvent,
InteractionStatusUpdate,
ContentStart,
@ -1193,6 +1307,11 @@ class InteractionsAPIResponse(BaseLiteLLMOpenAIResponseObject):
Response from the Interactions API.
Wraps the API response with LiteLLM-specific hidden params.
Schema notes:
- New schema (Api-Revision: 2026-05-20, default): response contains ``steps``.
- Legacy schema (Api-Revision: 2026-05-07, removed June 8 2026): response contains ``outputs``.
Both fields are kept here so callers work with either schema.
"""
id: Optional[str] = None
@ -1203,7 +1322,10 @@ class InteractionsAPIResponse(BaseLiteLLMOpenAIResponseObject):
created: Optional[str] = None
updated: Optional[str] = None
role: Optional[str] = None
# Legacy schema field (Api-Revision: 2026-05-07). Remove after June 8, 2026.
outputs: Optional[List[Dict[str, Any]]] = None
# New schema field (Api-Revision: 2026-05-20).
steps: Optional[List[Dict[str, Any]]] = None
usage: Optional[Dict[str, Any]] = None
_hidden_params: dict = PrivateAttr(default_factory=dict)
@ -1213,7 +1335,12 @@ class InteractionsAPIStreamingResponse(BaseLiteLLMOpenAIResponseObject):
"""
Streaming response chunk from the Interactions API.
Event types per OpenAPI spec:
New schema event types (Api-Revision: 2026-05-20):
- interaction.created, interaction.in_progress, interaction.completed,
interaction.requires_action
- step.start, step.delta, step.stop
Legacy event types (Api-Revision: 2026-05-07, removed June 8 2026):
- interaction.start, interaction.status_update, interaction.complete
- content.start, content.delta, content.stop
- error
@ -1228,9 +1355,17 @@ class InteractionsAPIStreamingResponse(BaseLiteLLMOpenAIResponseObject):
created: Optional[str] = None
updated: Optional[str] = None
role: Optional[str] = None
# Legacy schema field (Api-Revision: 2026-05-07). Remove after June 8, 2026.
outputs: Optional[List[Dict[str, Any]]] = None
# New schema field (Api-Revision: 2026-05-20).
steps: Optional[List[Dict[str, Any]]] = None
usage: Optional[Dict[str, Any]] = None
delta: Optional[Dict[str, Any]] = None
# New schema streaming fields
index: Optional[int] = None
step: Optional[Dict[str, Any]] = None
interaction_id: Optional[str] = None
interaction: Optional[Dict[str, Any]] = None
_hidden_params: dict = PrivateAttr(default_factory=dict)

View file

@ -16,15 +16,15 @@ GeminiEmbeddingInput = Union[EmbeddingInput, List[List[str]]]
class FunctionResponse(TypedDict, total=False):
# `id` correlates this response with the originating `functionCall` part.
# Required by Gemini 3.5+ for strict function-calling response matching.
# Supported on Google AI Studio Gemini 3.5+; Vertex AI rejects this field.
id: str
name: Required[str]
response: Optional[dict]
class FunctionCall(TypedDict, total=False):
# `id` is returned by Gemini 3.5+ to correlate the corresponding
# `functionResponse`. Older Gemini models omit this field.
# `id` correlates the corresponding `functionResponse` on Google AI Studio
# Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field.
id: str
name: Required[str]
args: Optional[dict]
@ -52,8 +52,8 @@ class PartType(TypedDict, total=False):
class HttpxFunctionCall(TypedDict, total=False):
# `id` is returned by Gemini 3.5+ to correlate the corresponding
# `functionResponse`. Older Gemini models omit this field.
# `id` correlates the corresponding `functionResponse` on Google AI Studio
# Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field.
id: str
name: Required[str]
args: dict

View file

@ -27344,6 +27344,58 @@
"supports_web_search": true,
"tpm": 800000
},
"openrouter/google/gemini-3.1-flash-lite": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_per_audio_token": 5e-08,
"input_cost_per_audio_token": 5e-07,
"input_cost_per_token": 2.5e-07,
"litellm_provider": "openrouter",
"max_audio_length_hours": 8.4,
"max_audio_per_prompt": 1,
"max_images_per_prompt": 3000,
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
"max_pdf_size_mb": 30,
"max_tokens": 65536,
"max_video_length": 1,
"max_videos_per_prompt": 10,
"mode": "chat",
"output_cost_per_reasoning_token": 1.5e-06,
"output_cost_per_token": 1.5e-06,
"rpm": 2000,
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-lite",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/completions",
"/v1/batch"
],
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_audio_input": true,
"supports_audio_output": false,
"supports_code_execution": true,
"supports_file_search": true,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_url_context": true,
"supports_video_input": true,
"supports_vision": true,
"supports_web_search": true,
"tpm": 800000
},
"openrouter/google/gemini-3.1-pro-preview": {
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
@ -28153,10 +28205,10 @@
"supports_tool_choice": true
},
"openrouter/xiaomi/mimo-v2-flash": {
"input_cost_per_token": 9e-08,
"output_cost_per_token": 2.9e-07,
"input_cost_per_token": 1e-07,
"output_cost_per_token": 3e-07,
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 0.0,
"cache_read_input_token_cost": 1e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 16384,
@ -28166,7 +28218,43 @@
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": false,
"supports_prompt_caching": false
"supports_prompt_caching": true
},
"openrouter/xiaomi/mimo-v2.5-pro": {
"input_cost_per_token": 1e-06,
"output_cost_per_token": 3e-06,
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 2e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": false,
"supports_response_schema": true,
"supports_prompt_caching": true
},
"openrouter/xiaomi/mimo-v2.5": {
"input_cost_per_token": 4e-07,
"output_cost_per_token": 2e-06,
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 8e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true,
"supports_response_schema": true,
"supports_prompt_caching": true
},
"openrouter/z-ai/glm-4.7": {
"input_cost_per_token": 4e-07,

View file

@ -3,7 +3,7 @@ import sys
import pytest
import asyncio
from typing import Optional
from unittest.mock import patch, AsyncMock
from unittest.mock import patch, AsyncMock, MagicMock
from litellm.responses.litellm_completion_transformation.handler import (
LiteLLMCompletionTransformationHandler,
)
@ -130,6 +130,26 @@ def test_multiturn_tool_calls():
print("follow_up_response=", follow_up_response)
def test_response_api_handler_merges_metadata_and_service_tier_without_error():
"""Sync path must merge kwargs like async; double-splat raises TypeError."""
handler = LiteLLMCompletionTransformationHandler()
with patch("litellm.completion", new_callable=MagicMock) as mock_completion:
mock_completion.return_value = ModelResponse(
id="id", created=0, model="test", object="chat.completion", choices=[]
)
handler.response_api_handler(
model="test",
input="hi",
responses_api_request={},
metadata={"trace": "abc"},
service_tier="auto",
)
assert mock_completion.call_count == 1
assert mock_completion.call_args.kwargs["metadata"] == {"trace": "abc"}
assert mock_completion.call_args.kwargs["service_tier"] == "auto"
@pytest.mark.asyncio
async def test_async_response_api_handler_merges_trace_id_without_error():
handler = LiteLLMCompletionTransformationHandler()
@ -158,3 +178,39 @@ async def test_async_response_api_handler_merges_trace_id_without_error():
assert (
mock_acompletion.call_args.kwargs["litellm_trace_id"] == "session-trace"
)
@pytest.mark.asyncio
async def test_aresponses_forwards_timeout_to_acompletion():
"""Regression test: timeout passed to aresponses() must reach acompletion()
on the completion transformation path (Anthropic, Bedrock, Vertex etc.).
Previously, `timeout` was a named param of `responses()` but was NOT
forwarded to `litellm_completion_transformation_handler.response_api_handler`,
so it was silently dropped — `Router(timeout=N)` was a no-op for Anthropic
and similar providers, with calls falling back to the provider SDK default
(~600s for Anthropic).
"""
with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion:
mock_acompletion.return_value = ModelResponse(
id="id",
created=0,
model="anthropic/claude-sonnet-4-5",
object="chat.completion",
choices=[],
)
await litellm.aresponses(
model="anthropic/claude-sonnet-4-5",
input="hello",
timeout=42,
api_key="sk-ant-fake",
)
assert mock_acompletion.call_count == 1
forwarded_timeout = mock_acompletion.call_args.kwargs.get("timeout")
assert forwarded_timeout == 42, (
f"timeout was not forwarded to acompletion (got {forwarded_timeout!r}); "
"this means Router(timeout=N) silently fails for providers on the "
"completion transformation path."
)

View file

@ -31,10 +31,8 @@ import litellm
# the cassette state the branch is being tested with.
from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
_local_cost_map = GetModelCostMap.load_local_model_cost_map()
for _k, _v in _local_cost_map.items():
for _k, _v in GetModelCostMap.load_local_model_cost_map().items():
litellm.model_cost.setdefault(_k, _v)
del _local_cost_map
from tests._vcr_conftest_common import ( # noqa: E402,F401
VerboseReporterState,

View file

@ -382,6 +382,11 @@ async def test_mcp_http_transport_tool_not_found():
}
)
# Mapping populated for this server but not for the requested tool
test_manager.tool_name_to_mcp_server_name_mapping["gmail_send_email"] = (
"test_http_server"
)
# Try to call a tool that doesn't exist in mapping
with pytest.raises(ValueError, match="Tool nonexistent_tool not found"):
await test_manager.call_tool(
@ -881,6 +886,7 @@ async def test_get_tools_from_mcp_servers():
extra_headers=None,
add_prefix=False,
raw_headers=None,
user_api_key_auth=None,
):
if server.server_id == "server1_id":
return [mock_tool_1]
@ -1856,6 +1862,7 @@ async def test_get_tools_for_single_server():
extra_headers=None,
add_prefix=False,
raw_headers=None,
user_api_key_auth=None,
)
# Verify the result

View file

@ -0,0 +1,268 @@
"""
Unit tests for the Responses-API streaming-fallback helpers added to Router
in PR #28215 (fix(router): wrap aresponses streaming iterator for mid-stream
fallbacks).
Targets the four helpers introduced on Router:
- _extract_partial_responses_usage
- _combine_responses_fallback_usage
- _build_responses_continuation_input
- _aresponses_streaming_iterator
"""
import os
import sys
from typing import Any, AsyncIterator, List
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../.."))
from litellm import Router
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
def _make_router() -> Router:
return Router(
model_list=[
{
"model_name": "primary",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-test",
},
},
{
"model_name": "fallback",
"litellm_params": {
"model": "openai/gpt-4o",
"api_key": "sk-test",
},
},
]
)
def _make_completed_event(
input_tokens: int, output_tokens: int, total_tokens: int
) -> ResponseCompletedEvent:
response = ResponsesAPIResponse.model_construct(
usage=ResponseAPIUsage(
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
)
)
return ResponseCompletedEvent.model_construct(
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
response=response,
)
# -------- _extract_partial_responses_usage --------
def test_extract_partial_responses_usage_native_completed():
"""Native path: completed_response carries usage → returned as-is."""
completed = _make_completed_event(11, 7, 18)
source = MagicMock()
source.completed_response = completed
usage = Router._extract_partial_responses_usage(source)
assert usage is not None
assert usage.input_tokens == 11
assert usage.output_tokens == 7
assert usage.total_tokens == 18
def test_extract_partial_responses_usage_no_completed_response():
"""Native path: no completed_response → returns None."""
source = MagicMock()
source.completed_response = None
usage = Router._extract_partial_responses_usage(source)
assert usage is None
# -------- _combine_responses_fallback_usage --------
def test_combine_responses_fallback_usage_sums_completed_event():
"""Partial-stream usage is summed into the fallback event's usage."""
fallback_event = _make_completed_event(5, 3, 8)
partial = ResponseAPIUsage(input_tokens=11, output_tokens=7, total_tokens=18)
Router._combine_responses_fallback_usage(fallback_event, partial)
combined = fallback_event.response.usage
assert combined is not None
assert combined.input_tokens == 16
assert combined.output_tokens == 10
assert combined.total_tokens == 26
def test_combine_responses_fallback_usage_passthrough_for_unknown_event():
"""Events that are not completed/failed/incomplete are not mutated."""
other = MagicMock() # not a ResponseCompletedEvent etc. → isinstance false
partial = ResponseAPIUsage(input_tokens=1, output_tokens=1, total_tokens=2)
Router._combine_responses_fallback_usage(other, partial)
# No mutation expected on the unknown event — call is a no-op.
# -------- _build_responses_continuation_input --------
def test_build_responses_continuation_input_from_string():
out = Router._build_responses_continuation_input(
"Hello world", "partial assistant text"
)
assert len(out) == 3
assert out[0]["role"] == "user"
assert out[0]["content"][0]["text"] == "Hello world"
assert out[1]["role"] == "developer"
assert out[2]["role"] == "assistant"
assert out[2]["content"][0]["text"] == "partial assistant text"
def test_build_responses_continuation_input_from_list_preserves_items():
existing: List[Any] = [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "msg1"}],
}
]
out = Router._build_responses_continuation_input(existing, "partial")
assert len(out) == 3
assert out[0]["content"][0]["text"] == "msg1"
assert out[1]["role"] == "developer"
assert out[2]["role"] == "assistant"
def test_build_responses_continuation_input_from_none():
out = Router._build_responses_continuation_input(None, "partial")
assert len(out) == 2
assert out[0]["role"] == "developer"
assert out[1]["role"] == "assistant"
# -------- _aresponses_streaming_iterator (passthrough smoke test) --------
@pytest.mark.asyncio
async def test_aresponses_streaming_iterator_passthrough():
"""
Without MidStreamFallbackError, the wrapper yields source events
unchanged and returns a BaseResponsesAPIStreamingIterator subclass.
"""
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
)
events = [_make_completed_event(1, 1, 2)]
class _FakeSource:
"""Minimal source iterator. Provides every attribute the wrapper
constructor reads from source_iterator."""
def __init__(self) -> None:
self._i = 0
self.completed_response = None
self.response = MagicMock()
self.model = "openai/gpt-4o-mini"
self.logging_obj = MagicMock()
self.responses_api_provider_config = MagicMock()
self.start_time = 0.0
self.litellm_metadata = {}
self.custom_llm_provider = "openai"
self.request_data = {}
self.call_type = "aresponses"
self._hidden_params: dict = {}
def __aiter__(self) -> AsyncIterator[Any]:
return self
async def __anext__(self):
if self._i >= len(events):
raise StopAsyncIteration
ev = events[self._i]
self._i += 1
return ev
async def aclose(self):
return None
router = _make_router()
source = _FakeSource()
wrapper = await router._aresponses_streaming_iterator(
source, initial_kwargs={"model": "primary"}
)
assert isinstance(wrapper, BaseResponsesAPIStreamingIterator)
collected = [ev async for ev in wrapper]
assert len(collected) == 1
assert collected[0].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
# -------- _aresponses_with_streaming_fallbacks --------
@pytest.mark.asyncio
async def test_aresponses_with_streaming_fallbacks_non_streaming_passthrough():
"""Non-streaming response is returned unchanged, no wrap."""
router = _make_router()
plain_response = MagicMock()
async def fake_original(**_kwargs):
return plain_response
with patch.object(
router,
"_ageneric_api_call_with_fallbacks",
new=AsyncMock(return_value=plain_response),
):
out = await router._aresponses_with_streaming_fallbacks(
original_function=fake_original,
model="primary",
stream=False,
)
assert out is plain_response
@pytest.mark.asyncio
async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator():
"""Streaming response is wrapped via _aresponses_streaming_iterator."""
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
)
router = _make_router()
streaming_iter = MagicMock(spec=BaseResponsesAPIStreamingIterator)
wrapped = MagicMock(spec=BaseResponsesAPIStreamingIterator)
async def fake_original(**_kwargs):
return streaming_iter
with patch.object(
router,
"_ageneric_api_call_with_fallbacks",
new=AsyncMock(return_value=streaming_iter),
), patch.object(
router,
"_aresponses_streaming_iterator",
new=AsyncMock(return_value=wrapped),
) as mock_wrap:
out = await router._aresponses_with_streaming_fallbacks(
original_function=fake_original,
model="primary",
stream=True,
)
assert out is wrapped
mock_wrap.assert_awaited_once()

View file

@ -109,6 +109,31 @@ class TestAzureContainerConfig:
assert "/openai/v1/containers" in url
def test_get_complete_url_strips_responses_path_and_preserves_api_version(self):
"""When api_base is the responses endpoint URL, get_complete_url must:
- strip /openai/responses (no double-path)
- use the api-version from api_base query string, NOT the deployment's
older api_version (e.g. 2024-08-01-preview → containers need 2025-04-01-preview)
"""
api_base = "https://my-resource.cognitiveservices.azure.com/openai/responses?api-version=2025-04-01-preview"
url = self.config.get_complete_url(
api_base=api_base,
litellm_params={"api_version": "2024-08-01-preview"},
)
assert (
"/openai/responses/openai/containers" not in url
), "path must not double /openai/responses"
assert "my-resource.cognitiveservices.azure.com" in url
assert "/openai/containers" in url or "/openai/v1/containers" in url
assert (
"2025-04-01-preview" in url
), "must use version from api_base, not litellm_params"
assert (
"2024-08-01-preview" not in url
), "must not fall back to older chat api_version"
def test_get_complete_url_raises_without_api_base(self, monkeypatch):
monkeypatch.delenv("AZURE_API_BASE", raising=False)
monkeypatch.setattr(litellm, "api_base", None)
@ -531,6 +556,92 @@ class TestAzureContainerKnownFailureRegressions:
assert qs.get("api-version") == ["v1"]
assert qs.get("foo") == ["bar"]
@pytest.mark.asyncio
async def test_regression_no_container_id_does_not_use_user_supplied_model_id(
self, monkeypatch
):
"""Operations without container_id (create, list) must NOT route via
_ageneric_api_call_with_fallbacks using a caller-supplied model_id.
Security boundary: only the path that holds a validated container_id
is trusted to fall back to the forwarded model_id. A caller setting
model_id without container_id on POST /v1/containers must not gain
access to an arbitrary deployment UUID.
"""
from litellm.router import Router
router = Router(
model_list=[
{
"model_name": "azure-model",
"litellm_params": {
"model": "azure/gpt-4",
"api_base": "https://my-resource.cognitiveservices.azure.com",
"api_key": "test-key",
"api_version": "2025-04-01-preview",
},
"model_info": {"id": "deployment-uuid-123"},
}
]
)
fallback_called = {"called": False}
async def _mock_fallback(original_function, **kwargs):
fallback_called["called"] = True
return {}
monkeypatch.setattr(router, "_ageneric_api_call_with_fallbacks", _mock_fallback)
original_called = {"called": False}
async def _noop(**kwargs):
original_called["called"] = True
return {}
# No container_id — simulates create/list; caller injects a model_id
await router._init_containers_api_endpoints(
original_function=_noop,
model_id="deployment-uuid-123",
custom_llm_provider="azure",
)
assert not fallback_called["called"], (
"_ageneric_api_call_with_fallbacks must NOT be called when "
"container_id is absent, even if model_id is supplied"
)
assert original_called["called"], "original_function must be called directly"
def test_regression_httpx_empty_params_strips_query_string(self):
"""httpx erases the URL query-string when params={} (empty dict) is passed.
Root cause of the Azure container 404s on POST/DELETE:
_build_query_params returns {} when the endpoint has no extra params;
passing that {} as params= to httpx wiped ?api-version=2025-04-01-preview.
Fix: every container httpx call now uses `params or None` so an empty
dict falls back to None, which tells httpx to leave the URL untouched.
"""
url = (
"https://resource.cognitiveservices.azure.com"
"/openai/containers/cntr_123?api-version=2025-04-01-preview"
)
client = httpx.AsyncClient()
req_none = client.build_request("DELETE", url, params=None)
assert "api-version=2025-04-01-preview" in str(req_none.url)
req_empty = client.build_request("DELETE", url, params={})
assert "api-version" not in str(
req_empty.url
), "Documents root cause: params={} strips the query string"
effective: dict = {}
req_guarded = client.build_request("DELETE", url, params=effective or None)
assert "api-version=2025-04-01-preview" in str(
req_guarded.url
), "`params or None` must preserve ?api-version"
def test_regression_proxy_resolves_azure_text_same_as_azure(self):
"""Router/proxy treat azure_text like azure for container config."""
from litellm.proxy.container_endpoints.handler_factory import (
@ -770,3 +881,143 @@ class TestAzureContainerKnownFailureRegressions:
assert captured["data"]["container_id"] == "cntr_123"
assert captured["data"]["custom_llm_provider"] == "azure"
assert captured["data"]["model_id"] == "model_abc123"
@pytest.mark.asyncio
async def test_regression_get_container_forwarding_params_sets_model_id_for_managed_id(
self,
):
"""get_container_forwarding_params must extract model_id from a
LiteLLM-managed encoded container ID and include it in the forwarding
dict. This is the proxy-side half of the native-Azure-ID routing fix:
the router's _init_containers_api_endpoints reads kwargs["model_id"]
which is set here.
"""
from litellm.proxy.container_endpoints.ownership import (
get_container_forwarding_params,
)
encoded_id = ResponsesAPIRequestUtils._build_container_id(
custom_llm_provider="azure",
model_id="deployment-uuid-123",
container_id="cntr_6a058b43d24c8190a226cfb1d35405b20115fb7875ff11df",
)
params = await get_container_forwarding_params(
container_id=encoded_id,
original_container_id="cntr_6a058b43d24c8190a226cfb1d35405b20115fb7875ff11df",
custom_llm_provider="azure",
)
assert (
params.get("model_id") == "deployment-uuid-123"
), "model_id must be forwarded to the router for managed container IDs"
assert params.get("container_id") == (
"cntr_6a058b43d24c8190a226cfb1d35405b20115fb7875ff11df"
)
assert params.get("custom_llm_provider") == "azure"
@pytest.mark.asyncio
async def test_regression_get_container_forwarding_params_recovers_model_id_for_native_id(
self, monkeypatch
):
"""Native Azure IDs (``cntr_<hex>``) cannot be decoded, so model_id
must be recovered from the ownership row's ``unified_object_id`` —
the encoded form captured at create time when the router selected a
specific deployment. Without this, the router-side fallback for
native IDs in ``_init_containers_api_endpoints`` is dead code.
"""
from types import SimpleNamespace
from unittest.mock import AsyncMock
from litellm.proxy.container_endpoints import ownership
from litellm.proxy.container_endpoints.ownership import (
get_container_forwarding_params,
)
native_id = "cntr_6a058b43d24c8190a226cfb1d35405b20115fb7875ff11df"
encoded_stored_id = ResponsesAPIRequestUtils._build_container_id(
custom_llm_provider="azure",
model_id="deployment-uuid-123",
container_id=native_id,
)
ownership._CONTAINER_STORED_ID_CACHE.flush_cache()
ownership._CONTAINER_OWNER_CACHE.flush_cache()
table = AsyncMock()
table.find_first.return_value = SimpleNamespace(
created_by="user-1",
file_purpose=ownership.CONTAINER_OBJECT_PURPOSE,
unified_object_id=encoded_stored_id,
)
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
)
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
params = await get_container_forwarding_params(
container_id=native_id,
original_container_id=native_id,
custom_llm_provider="azure",
)
assert params.get("model_id") == "deployment-uuid-123", (
"model_id must be recovered from the stored unified_object_id "
"for native upstream container IDs"
)
assert params.get("container_id") == native_id
assert params.get("custom_llm_provider") == "azure"
@pytest.mark.asyncio
async def test_regression_native_azure_container_id_uses_forwarded_model_id(
self, monkeypatch
):
"""Native Azure container IDs (cntr_ + hex, no LiteLLM payload) must
still route through _ageneric_api_call_with_fallbacks using the
model_id forwarded from the proxy ownership check so that deployment
credentials (api_base) are applied."""
from litellm.router import Router
router = Router(
model_list=[
{
"model_name": "azure-model",
"litellm_params": {
"model": "azure/gpt-4",
"api_base": "https://my-resource.cognitiveservices.azure.com",
"api_key": "test-key",
"api_version": "2025-04-01-preview",
},
"model_info": {"id": "deployment-uuid-123"},
}
]
)
called_with: dict = {}
async def _mock_fallback(original_function, **kwargs):
called_with.update(kwargs)
return {}
monkeypatch.setattr(router, "_ageneric_api_call_with_fallbacks", _mock_fallback)
native_azure_id = "cntr_6a058b43d24c8190a226cfb1d35405b20115fb7875ff11df"
async def _noop(**kwargs):
return {}
await router._init_containers_api_endpoints(
original_function=_noop,
container_id=native_azure_id,
model_id="deployment-uuid-123",
custom_llm_provider="azure",
)
assert called_with.get("model") == "deployment-uuid-123", (
"_ageneric_api_call_with_fallbacks must be called with the forwarded "
"model_id when the container_id carries no LiteLLM routing payload"
)

View file

@ -0,0 +1,260 @@
"""Regression: update_batch_in_database must not persist raw provider output_file_id."""
import json
from types import SimpleNamespace
from typing import Optional
import pytest
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.openai_files_endpoints.common_utils import (
ensure_batch_response_managed_file_ids,
update_batch_in_database,
)
from litellm.types.utils import LiteLLMBatch
def _build_batch_response(
*,
batch_id: str = "batch_managed_ids_test",
status: str = "completed",
output_file_id: Optional[str] = "file-rawoutput789",
error_file_id: Optional[str] = None,
hidden_params: Optional[dict] = None,
) -> LiteLLMBatch:
batch = LiteLLMBatch(
id=batch_id,
object="batch",
status=status,
endpoint="/v1/chat/completions",
input_file_id="file-input123",
output_file_id=output_file_id,
error_file_id=error_file_id,
completion_window="24h",
created_at=1234567890,
)
if hidden_params is not None:
batch._hidden_params = hidden_params # type: ignore[attr-defined]
return batch
def _build_managed_files_mock(unified_id: str = "file-bWFuYWdlZF9vdXRwdXRfaWQ="):
mock = MagicMock()
mock.get_unified_output_file_id = MagicMock(return_value=unified_id)
mock.store_unified_file_id = AsyncMock()
return mock
def _build_prisma_mock():
mock = MagicMock()
mock.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
mock.db.litellm_managedobjecttable.update = AsyncMock()
return mock
@pytest.mark.asyncio
async def test_update_batch_in_database_stores_unified_output_file_id():
raw_output_file_id = "file-rawoutput789"
unified_output_file_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ="
batch_id = "batch_managed_ids_test"
unified_batch_id = (
"litellm_proxy;model_id:my-model;llm_batch_id:batch_managed_ids_test"
)
response = _build_batch_response(
batch_id=batch_id,
output_file_id=raw_output_file_id,
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
)
mock_managed_files = _build_managed_files_mock(unified_id=unified_output_file_id)
mock_prisma = _build_prisma_mock()
await update_batch_in_database(
batch_id=batch_id,
unified_batch_id=unified_batch_id,
response=response,
managed_files_obj=mock_managed_files,
prisma_client=mock_prisma,
verbose_proxy_logger=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"),
)
stored = json.loads(
mock_prisma.db.litellm_managedobjecttable.update.call_args.kwargs["data"][
"file_object"
]
)
assert stored["output_file_id"] == unified_output_file_id
assert stored["output_file_id"] != raw_output_file_id
@pytest.mark.asyncio
async def test_ensure_batch_response_normalizes_error_file_id():
"""Both output_file_id and error_file_id must be normalized to managed IDs."""
unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ="
response = _build_batch_response(
output_file_id="file-raw-output",
error_file_id="file-raw-error",
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
)
mock_managed_files = _build_managed_files_mock(unified_id=unified_id)
mock_prisma = _build_prisma_mock()
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=mock_managed_files,
prisma_client=mock_prisma,
verbose_proxy_logger=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"),
)
assert response.output_file_id == unified_id
assert response.error_file_id == unified_id
assert mock_managed_files.get_unified_output_file_id.call_count == 2
@pytest.mark.asyncio
async def test_ensure_batch_response_swallows_conversion_errors():
"""When the managed-files conversion raises, the failure is logged, not propagated."""
raw_output_file_id = "file-raw-output"
response = _build_batch_response(
output_file_id=raw_output_file_id,
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
)
mock_managed_files = MagicMock()
mock_managed_files.get_unified_output_file_id = MagicMock(
side_effect=RuntimeError("boom")
)
mock_managed_files.store_unified_file_id = AsyncMock()
mock_logger = MagicMock()
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=mock_managed_files,
prisma_client=_build_prisma_mock(),
verbose_proxy_logger=mock_logger,
user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"),
)
assert response.output_file_id == raw_output_file_id
mock_logger.warning.assert_called()
@pytest.mark.asyncio
async def test_ensure_batch_response_builds_auth_from_db_batch_object():
"""If user_api_key_dict is omitted, fall back to created_by/team_id on db_batch_object."""
unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ="
response = _build_batch_response(
output_file_id="file-raw-output",
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
)
mock_managed_files = _build_managed_files_mock(unified_id=unified_id)
db_batch_object = SimpleNamespace(
created_by="user-from-db", team_id="team-from-db", status="completed"
)
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=mock_managed_files,
prisma_client=_build_prisma_mock(),
verbose_proxy_logger=MagicMock(),
db_batch_object=db_batch_object,
)
forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[
"user_api_key_dict"
]
assert forwarded_auth.user_id == "user-from-db"
assert forwarded_auth.team_id == "team-from-db"
@pytest.mark.asyncio
async def test_ensure_batch_response_resolves_model_name_from_unified_file_id():
"""When hidden_params lacks model_name, derive it from unified_file_id."""
unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ="
response = _build_batch_response(
output_file_id="file-raw-output",
hidden_params={
"model_id": "my-model",
"unified_file_id": "litellm_proxy:application/octet-stream;unified_id,abc;target_model_names,gpt-4o-mini,gemini-2.0-flash",
},
)
mock_managed_files = _build_managed_files_mock(unified_id=unified_id)
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=mock_managed_files,
prisma_client=_build_prisma_mock(),
verbose_proxy_logger=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"),
)
assert (
mock_managed_files.get_unified_output_file_id.call_args.kwargs["model_name"]
== "gpt-4o-mini,gemini-2.0-flash"
)
@pytest.mark.asyncio
async def test_ensure_batch_response_returns_early_without_managed_files_obj():
"""Without managed_files_obj, the helper is a no-op (no conversion attempted)."""
response = _build_batch_response(
output_file_id="file-raw-output",
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
)
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=None,
prisma_client=_build_prisma_mock(),
verbose_proxy_logger=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"),
)
assert response.output_file_id == "file-raw-output"
@pytest.mark.asyncio
async def test_ensure_batch_response_returns_early_without_model_id():
"""Without model_id in hidden_params, the helper cannot create managed IDs."""
response = _build_batch_response(
output_file_id="file-raw-output",
hidden_params={"model_name": "openai/gpt-4o"},
)
mock_managed_files = _build_managed_files_mock()
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=mock_managed_files,
prisma_client=_build_prisma_mock(),
verbose_proxy_logger=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"),
)
assert response.output_file_id == "file-raw-output"
mock_managed_files.get_unified_output_file_id.assert_not_called()
@pytest.mark.asyncio
async def test_ensure_batch_response_returns_early_without_auth():
"""Without user_api_key_dict or db_batch_object, no conversion is attempted."""
response = _build_batch_response(
output_file_id="file-raw-output",
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
)
mock_managed_files = _build_managed_files_mock()
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=mock_managed_files,
prisma_client=_build_prisma_mock(),
verbose_proxy_logger=MagicMock(),
)
assert response.output_file_id == "file-raw-output"
mock_managed_files.get_unified_output_file_id.assert_not_called()

View file

@ -1,10 +1,11 @@
"""
Tests for Gemini Interactions API transformation.
Covers credential leak prevention changes:
- validate_environment sets x-goog-api-key header
- get_complete_url excludes API key from URL
- get/delete/cancel interaction request URLs exclude API key
Covers:
- validate_environment: x-goog-api-key header, Api-Revision schema selection
- get_complete_url: API key excluded from URL
- get/delete/cancel interaction request URLs
- transform_request: response_mime_type coalescing, image_config migration
"""
import os
@ -15,6 +16,7 @@ import pytest
sys.path.insert(0, os.path.abspath("../../.."))
import litellm
from litellm.interactions.litellm_responses_transformation.streaming_iterator import (
LiteLLMResponsesInteractionsStreamingIterator,
)
@ -22,7 +24,6 @@ from litellm.llms.gemini.interactions.transformation import (
GoogleAIStudioInteractionsConfig,
)
from litellm.types.llms.openai import (
ContentPartAddedEvent,
OutputTextDeltaEvent,
ResponseCompletedEvent,
ResponseCreatedEvent,
@ -85,6 +86,30 @@ class TestValidateEnvironment:
assert headers["X-Custom"] == "value"
assert headers["x-goog-api-key"] == "test-key"
def test_api_revision_new_schema_by_default(self, config):
# Default: use_legacy_interactions_schema=False → new steps schema
original = litellm.use_legacy_interactions_schema
try:
litellm.use_legacy_interactions_schema = False
headers = config.validate_environment(
headers={}, model="gemini-2.5-flash", litellm_params=None
)
assert headers["Api-Revision"] == "2026-05-20"
finally:
litellm.use_legacy_interactions_schema = original
def test_api_revision_legacy_schema_when_flag_set(self, config):
# Flag on → legacy outputs schema until June 8, 2026
original = litellm.use_legacy_interactions_schema
try:
litellm.use_legacy_interactions_schema = True
headers = config.validate_environment(
headers={}, model="gemini-2.5-flash", litellm_params=None
)
assert headers["Api-Revision"] == "2026-05-07"
finally:
litellm.use_legacy_interactions_schema = original
class TestGetCompleteUrl:
def test_url_excludes_api_key(self, config):
@ -127,7 +152,12 @@ class TestTransformRequest:
request_body = config.transform_request(
model=None,
agent="my-custom-slides-agent",
input=[{"type": "text", "text": "Create a 5-slide presentation about AI trends."}],
input=[
{
"type": "text",
"text": "Create a 5-slide presentation about AI trends.",
}
],
optional_params={
"environment": "remote",
"stream": False,
@ -172,158 +202,7 @@ class TestTransformRequest:
)
assert request_body["environment"] == env_id
class TestStreamingIterator:
def _make_iterator(self) -> LiteLLMResponsesInteractionsStreamingIterator:
return LiteLLMResponsesInteractionsStreamingIterator(
model="gpt-5.4",
litellm_custom_stream_wrapper=MagicMock(),
request_input="hi",
optional_params={},
)
def _make_text_delta(
self, text: str, item_id: str = "item_1"
) -> OutputTextDeltaEvent:
event = MagicMock(spec=OutputTextDeltaEvent)
event.delta = text
event.item_id = item_id
return event
def _make_part_added(self, item_id: str = "item_1") -> ContentPartAddedEvent:
event = MagicMock(spec=ContentPartAddedEvent)
event.item_id = item_id
return event
def _make_response_created(self) -> ResponseCreatedEvent:
event = MagicMock(spec=ResponseCreatedEvent)
event.response = MagicMock(id="resp_123")
return event
def test_content_delta_includes_type_field(self):
"""content.delta events must carry delta.type='text' so the UI can display them."""
it = self._make_iterator()
it.sent_interaction_start = True
it.sent_content_start = True
chunk = it._transform_responses_chunk_to_interactions_chunk(
self._make_text_delta("Hello")
)
assert chunk is not None
assert chunk.event_type == "content.delta"
assert chunk.delta == {"type": "text", "text": "Hello"}
def test_response_part_added_emits_content_start(self):
"""ContentPartAddedEvent (arrives before text deltas) should emit content.start
so the first OutputTextDeltaEvent immediately emits content.delta without dropping text.
"""
it = self._make_iterator()
it.sent_interaction_start = True
chunk = it._transform_responses_chunk_to_interactions_chunk(
self._make_part_added()
)
assert chunk is not None
assert chunk.event_type == "content.start"
assert it.sent_content_start is True
def test_first_text_delta_not_dropped_when_part_added_seen(self):
"""After ContentPartAddedEvent, the first text delta must yield content.delta
(not content.start), preserving the token text."""
it = self._make_iterator()
it.sent_interaction_start = True
it._transform_responses_chunk_to_interactions_chunk(self._make_part_added())
chunk = it._transform_responses_chunk_to_interactions_chunk(
self._make_text_delta("Hello")
)
assert chunk is not None
assert chunk.event_type == "content.delta"
assert chunk.delta is not None
assert chunk.delta.get("text") == "Hello"
def test_part_added_emits_interaction_start_fallback_when_not_sent(self):
"""If ContentPartAddedEvent arrives before any ResponseCreatedEvent,
the iterator must emit interaction.start before content.start to honor
the documented event ordering contract."""
it = self._make_iterator()
chunk = it._transform_responses_chunk_to_interactions_chunk(
self._make_part_added(item_id="item_42")
)
assert chunk is not None
assert chunk.event_type == "interaction.start"
assert chunk.id == "item_42"
assert chunk.status == "in_progress"
assert chunk.model == "gpt-5.4"
assert it.sent_interaction_start is True
assert it.sent_content_start is False
def test_part_added_returns_none_when_already_started(self):
"""A second ContentPartAddedEvent (after content.start was already emitted)
should be a no-op so we don't re-emit content.start."""
it = self._make_iterator()
it.sent_interaction_start = True
it.sent_content_start = True
chunk = it._transform_responses_chunk_to_interactions_chunk(
self._make_part_added()
)
assert chunk is None
def test_part_added_without_item_id_falls_back_to_self_id(self):
"""When ContentPartAddedEvent has no item_id and we emit the interaction.start
fallback, the id must default to an interaction_<id(self)> string."""
it = self._make_iterator()
event = MagicMock(spec=ContentPartAddedEvent)
event.item_id = None
chunk = it._transform_responses_chunk_to_interactions_chunk(event)
assert chunk is not None
assert chunk.event_type == "interaction.start"
assert chunk.id == f"interaction_{id(it)}"
def test_first_text_delta_not_dropped_when_no_prior_start_events(self):
"""When OutputTextDeltaEvent arrives before any ResponseCreatedEvent or
ContentPartAddedEvent, the iterator must emit interaction.start *and*
immediately follow with a content.start that carries this delta's text,
so the first token is never silently dropped from the stream."""
events = [
self._make_text_delta("Hello"),
self._make_text_delta(" World"),
]
wrapper = MagicMock()
wrapper.__iter__ = lambda self: iter(events)
wrapper.__next__ = lambda self, _it=iter(events): next(_it)
it = LiteLLMResponsesInteractionsStreamingIterator(
model="gpt-5.4",
litellm_custom_stream_wrapper=wrapper,
request_input="hi",
optional_params={},
)
first = it._transform_responses_chunk_to_interactions_chunk(events[0])
assert first is not None
assert first.event_type == "interaction.start"
assert it.sent_interaction_start is True
assert it.sent_content_start is True
assert len(it._pending_events) == 1
pending = it._pending_events[0]
assert pending.event_type == "content.start"
assert pending.delta == {"type": "text", "text": "Hello"}
second = it._transform_responses_chunk_to_interactions_chunk(events[1])
assert second is not None
assert second.event_type == "content.delta"
assert second.delta == {"type": "text", "text": " World"}
class TestTransformRequest:
def test_stream_param_included_in_request_body(self, config):
"""When stream=True is in optional_params, the request body must include it
so the proxy forwards the SSE streaming flag to Google's backend."""
@ -352,6 +231,273 @@ class TestTransformRequest:
assert "stream" not in body
class TestStreamingIterator:
def _make_iterator(
self, use_legacy: bool = False
) -> LiteLLMResponsesInteractionsStreamingIterator:
original = litellm.use_legacy_interactions_schema
litellm.use_legacy_interactions_schema = use_legacy
try:
return LiteLLMResponsesInteractionsStreamingIterator(
model="gpt-5.4",
litellm_custom_stream_wrapper=MagicMock(),
request_input="hi",
optional_params={},
)
finally:
litellm.use_legacy_interactions_schema = original
def _make_text_delta(
self, text: str, item_id: str = "item_1"
) -> OutputTextDeltaEvent:
event = MagicMock(spec=OutputTextDeltaEvent)
event.delta = text
event.item_id = item_id
return event
def _make_response_created(self) -> ResponseCreatedEvent:
event = MagicMock(spec=ResponseCreatedEvent)
event.response = MagicMock(id="resp_123")
return event
def test_step_delta_includes_type_field(self):
"""step.delta events must carry delta.type='text' so the UI can display them."""
it = self._make_iterator(use_legacy=False)
it.sent_interaction_start = True
it.sent_content_start = True
chunk = it._transform_responses_chunk_to_interactions_chunk(
self._make_text_delta("Hello")
)
assert chunk is not None
assert chunk.event_type == "step.delta"
assert chunk.delta == {"type": "text", "text": "Hello"}
def test_content_delta_legacy_schema(self):
"""Legacy schema emits content.delta with type and text fields."""
it = self._make_iterator(use_legacy=True)
it.sent_interaction_start = True
it.sent_content_start = True
chunk = it._transform_responses_chunk_to_interactions_chunk(
self._make_text_delta("Hello")
)
assert chunk is not None
assert chunk.event_type == "content.delta"
assert chunk.delta == {"type": "text", "text": "Hello"}
def test_response_created_emits_interaction_created(self):
it = self._make_iterator(use_legacy=False)
chunk = it._transform_responses_chunk_to_interactions_chunk(
self._make_response_created()
)
assert chunk is not None
assert chunk.event_type == "interaction.created"
assert chunk.id == "resp_123"
assert it.sent_interaction_start is True
def test_response_created_emits_interaction_start_legacy(self):
it = self._make_iterator(use_legacy=True)
chunk = it._transform_responses_chunk_to_interactions_chunk(
self._make_response_created()
)
assert chunk is not None
assert chunk.event_type == "interaction.start"
assert chunk.id == "resp_123"
def test_text_delta_sequence_new_schema(self):
"""First chunk yields created + step.start + step.delta; later chunks yield step.delta."""
it = self._make_iterator(use_legacy=False)
first_events = it._events_for_chunk(self._make_text_delta("Hello"))
assert [e.event_type for e in first_events] == [
"interaction.created",
"step.start",
"step.delta",
]
assert first_events[-1].delta == {"type": "text", "text": "Hello"}
assert it.sent_interaction_start is True
assert it.sent_content_start is True
second_events = it._events_for_chunk(self._make_text_delta(" World"))
assert [e.event_type for e in second_events] == ["step.delta"]
assert second_events[0].delta == {"type": "text", "text": " World"}
third_events = it._events_for_chunk(self._make_text_delta("!"))
assert [e.event_type for e in third_events] == ["step.delta"]
assert third_events[0].delta == {"type": "text", "text": "!"}
def test_text_delta_sequence_legacy_schema(self):
"""Legacy: first chunk yields interaction.start + content.start + content.delta."""
it = self._make_iterator(use_legacy=True)
first_events = it._events_for_chunk(self._make_text_delta("Hello"))
assert [e.event_type for e in first_events] == [
"interaction.start",
"content.start",
"content.delta",
]
assert first_events[-1].delta == {"type": "text", "text": "Hello"}
second_events = it._events_for_chunk(self._make_text_delta(" World"))
assert [e.event_type for e in second_events] == ["content.delta"]
assert second_events[0].delta == {"type": "text", "text": " World"}
def test_first_text_delta_without_item_id_uses_fallback_id(self):
it = self._make_iterator(use_legacy=False)
event = self._make_text_delta("Hi")
event.item_id = None
events = it._events_for_chunk(event)
assert events[0].event_type == "interaction.created"
assert events[0].id == f"interaction_{id(it)}"
def test_first_text_delta_emits_text_via_compat_shim(self):
"""The legacy single-chunk shim must surface the synthetic events AND the delta."""
it = self._make_iterator(use_legacy=False)
first = it._transform_responses_chunk_to_interactions_chunk(
self._make_text_delta("Hello")
)
assert first is not None
assert first.event_type == "interaction.created"
second = it.__next__() if it._pending_events else None
assert second is not None
assert second.event_type == "step.start"
third = it.__next__() if it._pending_events else None
assert third is not None
assert third.event_type == "step.delta"
assert third.delta == {"type": "text", "text": "Hello"}
def test_response_created_then_text_delta_emits_step_start_and_delta(self):
"""Realistic flow: response.created arrives first, then text delta."""
it = self._make_iterator(use_legacy=False)
first = it._events_for_chunk(self._make_response_created())
assert [e.event_type for e in first] == ["interaction.created"]
second = it._events_for_chunk(self._make_text_delta("Hello"))
assert [e.event_type for e in second] == ["step.start", "step.delta"]
assert second[-1].delta == {"type": "text", "text": "Hello"}
def test_no_text_token_is_dropped_during_streaming(self):
"""Concatenated step.delta payloads must equal the upstream text."""
it = self._make_iterator(use_legacy=False)
chunks = ["Hello", " ", "world", "!"]
emitted_text = ""
for c in chunks:
for ev in it._events_for_chunk(self._make_text_delta(c)):
if ev.event_type == "step.delta":
assert ev.delta is not None
emitted_text += ev.delta["text"]
assert emitted_text == "Hello world!"
def test_stop_iteration_fallback_emits_completion_event(self):
"""If upstream ends without ResponseCompletedEvent, terminal events still flow."""
from unittest.mock import MagicMock
text_event = self._make_text_delta("hi")
sync_iter = MagicMock()
sync_iter.__iter__ = lambda self: self
sync_iter.__next__ = MagicMock(side_effect=[text_event, StopIteration])
original = litellm.use_legacy_interactions_schema
litellm.use_legacy_interactions_schema = False
try:
it = LiteLLMResponsesInteractionsStreamingIterator(
model="gpt-5.4",
litellm_custom_stream_wrapper=sync_iter,
request_input="hi",
optional_params={},
)
finally:
litellm.use_legacy_interactions_schema = original
emitted: list = []
try:
while True:
emitted.append(next(it))
except StopIteration:
pass
event_types = [e.event_type for e in emitted]
assert event_types == [
"interaction.created",
"step.start",
"step.delta",
"step.stop",
"interaction.completed",
]
terminal = emitted[-1]
assert terminal.steps == [
{
"type": "model_output",
"content": [{"type": "text", "text": "hi"}],
}
]
# EOF-flushed terminal event must carry the same id as interaction.created.
assert terminal.id == emitted[0].id == "item_1"
def test_response_completed_emits_stop_then_completion(self):
"""ResponseCompletedEvent expands into step.stop + interaction.completed."""
from unittest.mock import MagicMock
text_event = self._make_text_delta("hi")
completed = MagicMock(spec=ResponseCompletedEvent)
completed.response = MagicMock(id="resp_999")
sync_iter = MagicMock()
sync_iter.__iter__ = lambda self: self
sync_iter.__next__ = MagicMock(side_effect=[text_event, completed])
original = litellm.use_legacy_interactions_schema
litellm.use_legacy_interactions_schema = False
try:
it = LiteLLMResponsesInteractionsStreamingIterator(
model="gpt-5.4",
litellm_custom_stream_wrapper=sync_iter,
request_input="hi",
optional_params={},
)
finally:
litellm.use_legacy_interactions_schema = original
emitted: list = []
try:
while True:
emitted.append(next(it))
except StopIteration:
pass
event_types = [e.event_type for e in emitted]
assert event_types == [
"interaction.created",
"step.start",
"step.delta",
"step.stop",
"interaction.completed",
]
# StopIteration fallback path must NOT add a duplicate completion event.
assert event_types.count("interaction.completed") == 1
# When the stream starts directly with a text delta (no preceding
# response.created), the terminal events must reuse the id derived from
# the first chunk's item_id rather than switching to response.id, so
# consumers can correlate the start and completion events by id.
assert emitted[0].id == "item_1"
assert emitted[-1].id == "item_1"
class TestInteractionOperationUrls:
"""Test that get/delete/cancel interaction URLs exclude API key."""
@ -410,3 +556,152 @@ class TestInteractionOperationUrls:
litellm_params=GenericLiteLLMParams(api_key=None),
headers={},
)
class TestTransformRequestSchemaCoalescing:
"""Test new-schema request coalescing (Api-Revision: 2026-05-20)."""
def test_response_mime_type_folded_into_response_format(self, config):
original = litellm.use_legacy_interactions_schema
try:
litellm.use_legacy_interactions_schema = False
body = config.transform_request(
model="gemini/gemini-2.5-flash",
agent=None,
input="summarise",
optional_params={
"response_mime_type": "application/json",
"response_format": {"type": "object", "properties": {}},
},
litellm_params=GenericLiteLLMParams(),
headers={},
)
finally:
litellm.use_legacy_interactions_schema = original
# response_mime_type must not appear as a top-level body key
assert "response_mime_type" not in body
rf = body["response_format"]
assert rf["type"] == "text"
assert rf["mime_type"] == "application/json"
assert "schema" in rf
def test_image_config_moved_to_response_format(self, config):
original = litellm.use_legacy_interactions_schema
try:
litellm.use_legacy_interactions_schema = False
body = config.transform_request(
model="gemini/gemini-2.5-flash",
agent=None,
input="draw a sunset",
optional_params={
"generation_config": {
"temperature": 0.7,
"image_config": {"aspect_ratio": "1:1", "image_size": "1K"},
}
},
litellm_params=GenericLiteLLMParams(),
headers={},
)
finally:
litellm.use_legacy_interactions_schema = original
# image_config removed from generation_config
assert "image_config" not in body.get("generation_config", {})
# moved into response_format with type=image
rf = body["response_format"]
assert rf["type"] == "image"
assert rf["aspect_ratio"] == "1:1"
def test_response_mime_type_skipped_when_response_format_is_list(self, config):
"""Lists are already polymorphic; do not wrap them into schema."""
original = litellm.use_legacy_interactions_schema
try:
litellm.use_legacy_interactions_schema = False
rf_list = [
{"type": "text", "mime_type": "application/json"},
{"type": "image", "aspect_ratio": "1:1"},
]
body = config.transform_request(
model="gemini/gemini-2.5-flash",
agent=None,
input="multimodal",
optional_params={
"response_format": rf_list,
"response_mime_type": "application/json",
},
litellm_params=GenericLiteLLMParams(),
headers={},
)
finally:
litellm.use_legacy_interactions_schema = original
assert body["response_format"] == rf_list
assert "response_mime_type" not in body
def test_image_config_appended_to_response_format_list_without_mutating_input(
self, config
):
"""When response_format is already a list, image_config must not mutate optional_params."""
original = litellm.use_legacy_interactions_schema
try:
litellm.use_legacy_interactions_schema = False
text_rf = {"type": "text", "mime_type": "application/json"}
optional_params = {
"response_format": [text_rf],
"generation_config": {
"image_config": {"aspect_ratio": "16:9", "image_size": "2K"},
},
}
original_rf = optional_params["response_format"]
body = config.transform_request(
model="gemini/gemini-2.5-flash",
agent=None,
input="draw and summarise",
optional_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert optional_params["response_format"] is original_rf
assert len(optional_params["response_format"]) == 1
assert body["response_format"] == [
text_rf,
{"type": "image", "aspect_ratio": "16:9", "image_size": "2K"},
]
# Retry must not append a second image entry into the caller's list.
body_retry = config.transform_request(
model="gemini/gemini-2.5-flash",
agent=None,
input="draw and summarise",
optional_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert len(optional_params["response_format"]) == 1
assert body_retry["response_format"] == body["response_format"]
finally:
litellm.use_legacy_interactions_schema = original
def test_legacy_schema_passes_fields_unchanged(self, config):
original = litellm.use_legacy_interactions_schema
try:
litellm.use_legacy_interactions_schema = True
body = config.transform_request(
model="gemini/gemini-2.5-flash",
agent=None,
input="hello",
optional_params={
"response_mime_type": "application/json",
"generation_config": {"image_config": {"aspect_ratio": "16:9"}},
},
litellm_params=GenericLiteLLMParams(),
headers={},
)
finally:
litellm.use_legacy_interactions_schema = original
assert body["response_mime_type"] == "application/json"
assert body["generation_config"]["image_config"]["aspect_ratio"] == "16:9"

View file

@ -2476,6 +2476,120 @@ def test_reasoning_effort_does_not_set_output_config_for_older_models():
), f"output_config should not be set for {model}"
@pytest.mark.parametrize(
"reasoning_effort_value",
[
# String shape — what callers send when using `reasoning_effort="low"` directly.
"low",
# Dict shape with `effort` only — what the Responses->Chat parser produces
# when `reasoning={"effort": "low"}` is set without `summary`.
{"effort": "low"},
# Dict shape with `effort` AND `summary` — what the Responses->Chat parser
# produces when callers send `Reasoning(effort="low", summary="concise")`.
# PR #25359 added the dict-keeping branch for this case, but the Anthropic
# transformation must coerce the dict back to a string before mapping.
{"effort": "low", "summary": "concise"},
{"effort": "low", "summary": "detailed"},
],
)
def test_reasoning_effort_accepts_dict_shape_for_adaptive_model(reasoning_effort_value):
"""
Adaptive-thinking (Claude 4.6+) branch: dict-shape reasoning_effort must
map to ``thinking.type='adaptive'`` + ``output_config.effort``.
Regression test for the dict-shape ``reasoning_effort`` produced by the
Responses->Chat parser when ``summary`` is set on the request's
``reasoning`` field. Before this fix, the Anthropic transformation guarded
on ``isinstance(value, str)`` and silently dropped the param — disabling
extended thinking entirely.
"""
config = AnthropicConfig()
result = config.map_openai_params(
non_default_params={"reasoning_effort": reasoning_effort_value},
optional_params={},
model="claude-sonnet-4-6-20260219",
drop_params=False,
)
# thinking must be set (adaptive for 4.6+)
assert "thinking" in result, (
f"thinking missing for reasoning_effort={reasoning_effort_value!r}"
)
assert result["thinking"]["type"] == "adaptive"
# output_config must carry the mapped effort
assert "output_config" in result, (
f"output_config missing for reasoning_effort={reasoning_effort_value!r}"
)
assert result["output_config"]["effort"] == "low"
@pytest.mark.parametrize(
"reasoning_effort_value",
[
"low",
{"effort": "low"},
{"effort": "low", "summary": "concise"},
],
)
def test_reasoning_effort_accepts_dict_shape_for_non_adaptive_model(reasoning_effort_value):
"""
Non-adaptive (pre-4.6) branch: dict-shape reasoning_effort must still map
to ``thinking.type='enabled'`` + ``budget_tokens``. ``output_config`` must
NOT be set on these models.
"""
config = AnthropicConfig()
result = config.map_openai_params(
non_default_params={"reasoning_effort": reasoning_effort_value},
optional_params={},
model="claude-sonnet-4-5-20250929",
drop_params=False,
)
assert "thinking" in result, (
f"thinking missing for reasoning_effort={reasoning_effort_value!r}"
)
assert result["thinking"]["type"] == "enabled"
assert "budget_tokens" in result["thinking"]
assert result["thinking"]["budget_tokens"] > 0
# Older models must not get adaptive-thinking output_config
assert "output_config" not in result, (
f"output_config should not be set for non-adaptive model "
f"(reasoning_effort={reasoning_effort_value!r})"
)
@pytest.mark.parametrize(
"bad_value",
[
{"summary": "concise"}, # missing effort
{"effort": None}, # explicit None effort
{"effort": 123}, # non-string effort
],
)
def test_reasoning_effort_unparseable_dict_is_dropped(bad_value):
"""
A dict shape that doesn't carry a usable ``effort`` key (e.g. only
``summary`` is set, or the value is some other unexpected type) should be
silently dropped — not crash, not partially apply.
"""
config = AnthropicConfig()
result = config.map_openai_params(
non_default_params={"reasoning_effort": bad_value},
optional_params={},
model="claude-sonnet-4-6-20260219",
drop_params=False,
)
assert "thinking" not in result, (
f"thinking should not be set for bad value {bad_value!r}"
)
assert "output_config" not in result, (
f"output_config should not be set for bad value {bad_value!r}"
)
@pytest.mark.parametrize(
"model",
[

View file

@ -2097,6 +2097,125 @@ def test_is_gemini_3_or_newer():
assert VertexGeminiConfig._is_gemini_3_or_newer("") == False
def test_forward_gemini_function_call_id_vertex_vs_google_ai_studio():
"""Vertex AI rejects `id` on function_call/function_response; Google AI Studio accepts it on Gemini 3.5+."""
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
model = "gemini-3.5-flash"
assert (
VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai") is False
)
assert (
VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai_beta")
is False
)
assert VertexGeminiConfig._forward_gemini_function_call_id(model, "gemini") is True
assert VertexGeminiConfig._forward_gemini_function_call_id(model, None) is False
assert (
VertexGeminiConfig._forward_gemini_function_call_id(
"gemini-2.5-flash", "gemini"
)
is False
)
def test_vertex_ai_gemini_35_tool_calls_omit_function_call_id():
"""Regression: Vertex must not send OpenAI tool_call id inside Gemini function_call parts."""
from litellm.llms.vertex_ai.gemini.transformation import (
_gemini_convert_messages_with_history,
)
messages = [
{"role": "user", "content": "Explore this directory"},
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_50e7e0fe0989464a89f188eda443",
"type": "function",
"function": {
"name": "read",
"arguments": '{"filePath": "/tmp"}',
},
}
],
},
{
"role": "tool",
"tool_call_id": "call_50e7e0fe0989464a89f188eda443",
"content": "ok",
},
]
contents = _gemini_convert_messages_with_history(
messages=messages,
model="gemini-3.5-flash",
custom_llm_provider="vertex_ai",
)
for content in contents:
for part in content.get("parts", []):
fc = part.get("function_call")
if fc is not None:
assert "id" not in fc, f"Vertex payload must not include id: {fc}"
fr = part.get("function_response")
if fr is not None:
assert "id" not in fr, f"Vertex payload must not include id: {fr}"
def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id():
from litellm.llms.vertex_ai.gemini.transformation import (
_gemini_convert_messages_with_history,
)
tool_call_id = "call_50e7e0fe0989464a89f188eda443"
messages = [
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": tool_call_id,
"type": "function",
"function": {
"name": "read",
"arguments": '{"filePath": "/tmp"}',
},
}
],
},
{
"role": "tool",
"tool_call_id": tool_call_id,
"content": "ok",
},
]
contents = _gemini_convert_messages_with_history(
messages=messages,
model="gemini-3.5-flash",
custom_llm_provider="gemini",
)
function_call_ids = []
function_response_ids = []
for content in contents:
for part in content.get("parts", []):
fc = part.get("function_call")
if fc is not None:
function_call_ids.append(fc.get("id"))
fr = part.get("function_response")
if fr is not None:
function_response_ids.append(fr.get("id"))
assert function_call_ids == [tool_call_id]
assert function_response_ids == [tool_call_id]
def test_reasoning_effort_maps_to_thinking_level_gemini_3():
"""Test that reasoning_effort maps to thinking_level AND includeThoughts for Gemini 3+ models"""
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
@ -3531,7 +3650,12 @@ def test_video_metadata_supported_for_all_gemini_models():
}
]
for model in ["gemini-1.5-pro", "gemini-2.5-flash", "gemini-2.5-pro", "gemini-3-pro-preview"]:
for model in [
"gemini-1.5-pro",
"gemini-2.5-flash",
"gemini-2.5-pro",
"gemini-3-pro-preview",
]:
contents = _gemini_convert_messages_with_history(messages=messages, model=model)
file_part = None
@ -3541,19 +3665,25 @@ def test_video_metadata_supported_for_all_gemini_models():
break
assert file_part is not None, f"{model}: file part should exist"
assert "video_metadata" in file_part, f"{model}: video_metadata should be present"
assert (
"video_metadata" in file_part
), f"{model}: video_metadata should be present"
assert file_part["video_metadata"]["fps"] == 5, f"{model}: fps should be 5"
# Per-part media_resolution is Gemini 3+ only; 2.x uses generation_config global
for model in ["gemini-3-pro-preview"]:
contents = _gemini_convert_messages_with_history(messages=messages, model=model)
file_part = next(p for p in contents[0]["parts"] if "file_data" in p)
assert "media_resolution" in file_part, f"{model}: media_resolution should be present"
assert (
"media_resolution" in file_part
), f"{model}: media_resolution should be present"
for model in ["gemini-1.5-pro", "gemini-2.5-flash", "gemini-2.5-pro"]:
contents = _gemini_convert_messages_with_history(messages=messages, model=model)
file_part = next(p for p in contents[0]["parts"] if "file_data" in p)
assert "media_resolution" not in file_part, f"{model}: per-part media_resolution should not be set"
assert (
"media_resolution" not in file_part
), f"{model}: per-part media_resolution should not be set"
def test_chunk_parser_handles_prompt_feedback_block():
@ -4186,8 +4316,9 @@ def test_vertex_ai_usage_metadata_with_document_tokens_in_prompt():
# DOCUMENT tokens should be included in text_tokens: 8 (TEXT) + 774 (DOCUMENT) = 782
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.text_tokens == 782, \
"DOCUMENT modality tokens should be added to text_tokens (8 TEXT + 774 DOCUMENT = 782)"
assert (
result.prompt_tokens_details.text_tokens == 782
), "DOCUMENT modality tokens should be added to text_tokens (8 TEXT + 774 DOCUMENT = 782)"
# Verify completion token details
assert result.completion_tokens_details is not None
@ -4222,8 +4353,9 @@ def test_vertex_ai_usage_metadata_with_document_tokens_cached():
# DOCUMENT cached tokens map to cached_text_tokens, so:
# text_tokens = (8 TEXT + 774 DOCUMENT) - 400 cached = 382
assert result.prompt_tokens_details.text_tokens == 382, \
"text_tokens should be (8 + 774) - 400 cached = 382"
assert (
result.prompt_tokens_details.text_tokens == 382
), "text_tokens should be (8 + 774) - 400 cached = 382"
assert result.prompt_tokens_details.cached_tokens == 400
@ -4693,7 +4825,9 @@ def test_mid_stream_429_error_raises_during_iteration():
{
"content": {
"role": "model",
"parts": [{"text": "Let me think about this...", "thought": True}],
"parts": [
{"text": "Let me think about this...", "thought": True}
],
},
"index": 0,
}
@ -4713,7 +4847,9 @@ def test_mid_stream_429_error_raises_during_iteration():
{
"content": {
"role": "model",
"parts": [{"text": "I'll generate the image now.", "thought": True}],
"parts": [
{"text": "I'll generate the image now.", "thought": True}
],
},
"index": 0,
}

View file

@ -1229,6 +1229,143 @@ def test_validate_trusted_redirect_uri_rejects_fragment_and_bad_scheme():
assert exc.value.status_code == 400, uri
def test_validate_trusted_redirect_uri_accepts_cursor_native_callback():
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
req = _make_trusted_request("http://localhost:4000/")
validate_trusted_redirect_uri(req, "cursor://anysphere.cursor-mcp/oauth/callback")
def test_validate_trusted_redirect_uri_rejects_unlisted_native_callback(
monkeypatch,
):
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
monkeypatch.setenv("MCP_TRUSTED_NATIVE_REDIRECT_URIS", "")
# Clear defaults by patching — env-only path for this test
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.oauth_utils._DEFAULT_NATIVE_REDIRECT_URIS",
[],
)
req = _make_trusted_request("http://localhost:4000/")
with pytest.raises(HTTPException) as exc:
validate_trusted_redirect_uri(
req, "cursor://anysphere.cursor-mcp/oauth/callback"
)
assert exc.value.status_code == 400
def test_validate_trusted_redirect_uri_accepts_env_native_redirect_uri(
monkeypatch,
):
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.oauth_utils._DEFAULT_NATIVE_REDIRECT_URIS",
[],
)
monkeypatch.setenv(
"MCP_TRUSTED_NATIVE_REDIRECT_URIS",
"vscode://my-app/oauth/callback",
)
req = _make_trusted_request("http://localhost:4000/")
validate_trusted_redirect_uri(req, "vscode://my-app/oauth/callback")
def test_validate_trusted_redirect_uri_rejects_native_callback_with_fragment():
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
req = _make_trusted_request("http://localhost:4000/")
with pytest.raises(HTTPException) as exc:
validate_trusted_redirect_uri(
req, "cursor://anysphere.cursor-mcp/oauth/callback#frag"
)
assert exc.value.status_code == 400
def test_validate_trusted_redirect_uri_rejects_native_callback_with_query():
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
req = _make_trusted_request("http://localhost:4000/")
with pytest.raises(HTTPException) as exc:
validate_trusted_redirect_uri(
req,
"cursor://anysphere.cursor-mcp/oauth/callback?injected=anything",
)
assert exc.value.status_code == 400
def test_validate_trusted_redirect_uri_native_path_case_insensitive(monkeypatch):
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.oauth_utils._DEFAULT_NATIVE_REDIRECT_URIS",
[],
)
monkeypatch.setenv(
"MCP_TRUSTED_NATIVE_REDIRECT_URIS",
"myapp://host/MyPath",
)
req = _make_trusted_request("http://localhost:4000/")
validate_trusted_redirect_uri(req, "myapp://host/MyPath")
def test_validate_trusted_redirect_uri_native_wildcard_respects_path_boundary(
monkeypatch,
):
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.oauth_utils._DEFAULT_NATIVE_REDIRECT_URIS",
[],
)
monkeypatch.setenv(
"MCP_TRUSTED_NATIVE_REDIRECT_URIS",
"cursor://anysphere.cursor-mcp/oauth/callback*",
)
req = _make_trusted_request("http://localhost:4000/")
validate_trusted_redirect_uri(
req, "cursor://anysphere.cursor-mcp/oauth/callback/extra"
)
with pytest.raises(HTTPException):
validate_trusted_redirect_uri(
req, "cursor://anysphere.cursor-mcp/oauth/callback-2"
)
def test_validate_trusted_redirect_uri_native_wildcard_directory_prefix(
monkeypatch,
):
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.oauth_utils._DEFAULT_NATIVE_REDIRECT_URIS",
[],
)
monkeypatch.setenv(
"MCP_TRUSTED_NATIVE_REDIRECT_URIS",
"cursor://anysphere.cursor-mcp/oauth/*",
)
req = _make_trusted_request("http://localhost:4000/")
validate_trusted_redirect_uri(req, "cursor://anysphere.cursor-mcp/oauth/callback")
def test_validate_trusted_redirect_uri_rejects_scheme_mismatch_on_same_host():
"""Regression: an attacker who can serve http on the proxy's own
host (e.g. by MITMing an unencrypted LAN hop) must not be able to

View file

@ -774,6 +774,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
extra_headers=None,
add_prefix=True,
raw_headers=None,
user_api_key_auth=None,
):
if server.name == "working_server":
# Working server returns tools
@ -879,6 +880,7 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
extra_headers=None,
add_prefix=True,
raw_headers=None,
user_api_key_auth=None,
):
# All servers fail
raise Exception(f"Server {server.name} connection failed")
@ -1339,6 +1341,7 @@ async def test_list_tools_single_server_unprefixed_names():
extra_headers=None,
add_prefix=False,
raw_headers=None,
user_api_key_auth=None,
):
tool = MagicMock()
tool.name = f"{server.alias}-toolA" if add_prefix else "toolA"
@ -1420,6 +1423,7 @@ async def test_list_tools_multiple_servers_prefixed_names():
extra_headers=None,
add_prefix=True,
raw_headers=None,
user_api_key_auth=None,
):
tool = MagicMock()
# When multiple servers, add_prefix should be True -> prefixed names
@ -1686,6 +1690,7 @@ async def test_list_tools_filters_by_key_team_permissions():
extra_headers=None,
add_prefix=False,
raw_headers=None,
user_api_key_auth=None,
):
# Return 4 tools, but only 2 should be allowed
tool1 = MagicMock()
@ -1795,6 +1800,7 @@ async def test_list_tools_with_team_tool_permissions_inheritance():
extra_headers=None,
add_prefix=False,
raw_headers=None,
user_api_key_auth=None,
):
# Return 4 tools
tool1 = MagicMock()
@ -1890,6 +1896,7 @@ async def test_list_tools_with_no_tool_permissions_shows_all():
extra_headers=None,
add_prefix=False,
raw_headers=None,
user_api_key_auth=None,
):
# Return 3 tools
tool1 = MagicMock()
@ -1988,6 +1995,7 @@ async def test_list_tools_strips_prefix_when_matching_permissions():
extra_headers=None,
add_prefix=True,
raw_headers=None,
user_api_key_auth=None,
):
# Return tools WITH prefix (as they come from MCP server)
tool1 = MagicMock()

View file

@ -322,6 +322,7 @@ class TestMCPServerManager:
mcp_auth_header=None,
mcp_protocol_version=None,
raw_headers=None,
user_api_key_auth=None,
):
if server.name == "github":
tool1 = MagicMock()
@ -376,6 +377,7 @@ class TestMCPServerManager:
mcp_auth_header=None,
mcp_protocol_version=None,
raw_headers=None,
user_api_key_auth=None,
):
assert mcp_auth_header == "legacy-token" # Should use legacy header
tool = MagicMock()
@ -414,6 +416,7 @@ class TestMCPServerManager:
mcp_auth_header=None,
mcp_protocol_version=None,
raw_headers=None,
user_api_key_auth=None,
):
assert (
mcp_auth_header == "server-specific-token"
@ -1004,6 +1007,7 @@ class TestMCPServerManager:
mcp_auth_header=None,
mcp_protocol_version=None,
raw_headers=None,
user_api_key_auth=None,
):
assert (
mcp_auth_header == "server-specific-token"
@ -1801,6 +1805,258 @@ class TestMCPServerManager:
assert len(tools_unprefixed) == 1
assert tools_unprefixed[0].name == "send_email"
@pytest.mark.asyncio
async def test_get_tools_from_server_jwt_skipped_when_mcp_auth_header_set(self):
"""When a per-user mcp_auth_header is resolved, JWT injection must be skipped.
MCPClient._get_auth_headers() applies extra_headers AFTER writing
Authorization from auth_value, so an injected JWT would clobber the
user's per-server OAuth token. Regression test for that interaction.
"""
from litellm.proxy._types import UserAPIKeyAuth
manager = MCPServerManager()
server = MCPServer(
server_id="zapier",
name="zapier",
transport=MCPTransport.http,
)
manager._create_mcp_client = AsyncMock(return_value=object())
manager._fetch_tools_with_timeout = AsyncMock(return_value=[])
user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="alice")
with (
patch(
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer",
return_value=MagicMock(),
),
patch(
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.inject_mcp_jwt_headers_for_upstream",
new=AsyncMock(return_value={"Authorization": "Bearer signed-jwt"}),
) as mock_inject,
):
# Case A: mcp_auth_header present -> JWT must NOT be injected
await manager._get_tools_from_server(
server,
mcp_auth_header="oauth-user-token",
user_api_key_auth=user_auth,
)
mock_inject.assert_not_called()
# Case B: no mcp_auth_header -> JWT injection runs as before
await manager._get_tools_from_server(
server,
user_api_key_auth=user_auth,
)
mock_inject.assert_awaited_once()
def test_resolve_mcp_server_for_tool_call_via_prefixed_name(self):
"""Resolution succeeds when the prefixed tool name is in the mapping."""
manager = MCPServerManager()
server = MCPServer(
server_id="jira",
name="jira",
transport=MCPTransport.http,
)
manager.registry = {"jira": server}
manager.tool_name_to_mcp_server_name_mapping["jira-search_issues"] = "jira"
manager.tool_name_to_mcp_server_name_mapping["search_issues"] = "jira"
resolved = manager._resolve_mcp_server_for_tool_call("jira", "search_issues")
assert resolved is server
def test_resolve_mcp_server_for_tool_call_via_alias(self):
"""Resolution falls back to alias/server_name match in the registry."""
manager = MCPServerManager()
server = MCPServer(
server_id="srv-uuid-123",
name="zapier",
alias="zapier-alias",
transport=MCPTransport.http,
)
manager.registry = {"srv-uuid-123": server}
manager.tool_name_to_mcp_server_name_mapping["create_zap"] = "zapier"
resolved = manager._resolve_mcp_server_for_tool_call(
"zapier-alias", "create_zap"
)
assert resolved is server
def test_resolve_mcp_server_for_tool_call_unknown_tool_with_empty_mapping(self):
"""Server-name match alone must not let unknown tools through when the
mapping has no entries for that server (e.g. listing has not completed
or the server is OAuth2 and the user has not yet listed tools).
"""
manager = MCPServerManager()
server = MCPServer(
server_id="srv-uuid-123",
name="zapier",
alias="zapier-alias",
transport=MCPTransport.http,
)
manager.registry = {"srv-uuid-123": server}
with pytest.raises(ValueError, match="Tool create_zap not found"):
manager._resolve_mcp_server_for_tool_call("zapier-alias", "create_zap")
def test_resolve_mcp_server_for_tool_call_fallback_to_unprefixed_lookup(self):
"""Fallback to unprefixed _get_mcp_server_from_tool_name when other paths fail."""
manager = MCPServerManager()
server = MCPServer(
server_id="linear",
name="linear",
transport=MCPTransport.http,
)
manager.registry = {"linear": server}
manager.tool_name_to_mcp_server_name_mapping["create_issue"] = "linear"
# server_name is empty so the fallback unprefixed lookup runs and matches.
resolved = manager._resolve_mcp_server_for_tool_call("", "create_issue")
assert resolved is server
def test_resolve_mcp_server_for_tool_call_raises_when_not_found(self):
"""ValueError is raised when no resolution path finds the tool."""
manager = MCPServerManager()
with pytest.raises(ValueError, match="Tool .* not found"):
manager._resolve_mcp_server_for_tool_call("nonexistent", "ghost_tool")
def test_resolve_mcp_server_for_tool_call_unknown_tool_with_known_server(self):
"""Server-name match alone must not let unknown tools slip through.
If the registry has tools for this server but neither the prefixed nor
unprefixed tool name is in the mapping, raise rather than returning the
server (would otherwise allow tool enumeration via name spoofing).
"""
manager = MCPServerManager()
server = MCPServer(
server_id="github",
name="github",
transport=MCPTransport.http,
)
manager.registry = {"github": server}
# Mapping has *some* tools for github but not "missing_tool".
manager.tool_name_to_mcp_server_name_mapping["github-list_repos"] = "github"
manager.tool_name_to_mcp_server_name_mapping["list_repos"] = "github"
with pytest.raises(ValueError, match="Tool missing_tool not found"):
manager._resolve_mcp_server_for_tool_call("github", "missing_tool")
@pytest.mark.asyncio
async def test_resolve_oauth2_headers_skipped_when_not_user_oauth(self):
"""Returns input headers unchanged when server does not need user OAuth."""
from litellm.proxy._types import UserAPIKeyAuth
manager = MCPServerManager()
server = MCPServer(
server_id="plain",
name="plain",
transport=MCPTransport.http,
)
# needs_user_oauth_token defaults to False.
user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="bob")
result = await manager._resolve_oauth2_headers_for_tool_call(
server, oauth2_headers=None, user_api_key_auth=user_auth
)
assert result is None
@pytest.mark.asyncio
async def test_resolve_oauth2_headers_returns_client_supplied_token(self):
"""Returns the client's oauth2_headers as-is when already set."""
from litellm.proxy._types import UserAPIKeyAuth
manager = MCPServerManager()
server = MCPServer(
server_id="oauth-srv",
name="oauth-srv",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
)
assert server.needs_user_oauth_token is True
user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="alice")
supplied = {"Authorization": "Bearer client-supplied"}
result = await manager._resolve_oauth2_headers_for_tool_call(
server, oauth2_headers=supplied, user_api_key_auth=user_auth
)
assert result is supplied
@pytest.mark.asyncio
async def test_resolve_oauth2_headers_looks_up_stored_token(self):
"""Falls back to stored per-user OAuth headers when no token is supplied."""
from litellm.proxy._types import UserAPIKeyAuth
manager = MCPServerManager()
server = MCPServer(
server_id="oauth-srv",
name="oauth-srv",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
)
user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="alice")
stored = {"Authorization": "Bearer stored-user-token"}
with patch(
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
new=AsyncMock(return_value=stored),
) as mock_lookup:
result = await manager._resolve_oauth2_headers_for_tool_call(
server, oauth2_headers=None, user_api_key_auth=user_auth
)
assert result == stored
mock_lookup.assert_awaited_once()
@pytest.mark.asyncio
async def test_resolve_oauth2_headers_swallows_lookup_exception(self):
"""Returns supplied headers (None) when the stored-token lookup raises."""
from litellm.proxy._types import UserAPIKeyAuth
manager = MCPServerManager()
server = MCPServer(
server_id="oauth-srv",
name="oauth-srv",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
)
user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="alice")
with patch(
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
new=AsyncMock(side_effect=RuntimeError("redis down")),
):
result = await manager._resolve_oauth2_headers_for_tool_call(
server, oauth2_headers=None, user_api_key_auth=user_auth
)
assert result is None
@pytest.mark.asyncio
async def test_resolve_oauth2_headers_no_user_id(self):
"""Skip lookup entirely when user_api_key_auth has no user_id."""
from litellm.proxy._types import UserAPIKeyAuth
manager = MCPServerManager()
server = MCPServer(
server_id="oauth-srv",
name="oauth-srv",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
)
# user_id is None -> lookup must not happen
user_auth = UserAPIKeyAuth(api_key="sk-test")
with patch(
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
new=AsyncMock(return_value={"Authorization": "Bearer x"}),
) as mock_lookup:
result = await manager._resolve_oauth2_headers_for_tool_call(
server, oauth2_headers=None, user_api_key_auth=user_auth
)
assert result is None
mock_lookup.assert_not_called()
def test_create_prefixed_tools_updates_mapping_for_both_forms(self):
"""_create_prefixed_tools should populate mapping for prefixed and original names even when not adding prefix in output."""
manager = MCPServerManager()

View file

@ -1,5 +1,6 @@
import json
from typing import Any, Dict, Optional
from unittest.mock import MagicMock
import pytest
from fastapi import HTTPException
@ -796,6 +797,25 @@ class TestCallToolRestAPI:
raising=False,
)
mock_server = MagicMock()
mock_server.server_id = "server-1"
def fake_get_mcp_server_by_id(server_id):
return mock_server if server_id == "server-1" else None
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_id",
fake_get_mcp_server_by_id,
raising=False,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_name",
lambda *args, **kwargs: None,
raising=False,
)
request_payload = {
"server_id": "server-1",
"name": "demo-tool",

View file

@ -219,7 +219,7 @@ def test_build_claims_scope_with_tool():
def test_build_claims_scope_without_tool():
"""_build_claims() includes mcp:tools/list when no specific tool is called."""
"""_build_claims() emits only mcp:tools/list when no specific tool is called."""
signer = _make_signer()
user_dict = _make_user_api_key_dict()
data: Dict[str, Any] = {}
@ -227,10 +227,11 @@ def test_build_claims_scope_without_tool():
claims = signer._build_claims(user_dict, data)
scopes = set(claims["scope"].split())
assert "mcp:tools/call" in scopes
assert "mcp:tools/list" in scopes
# List-only JWTs must NOT carry mcp:tools/call — least-privilege
assert "mcp:tools/call" not in scopes
# No per-tool call scope when no tool name was given
assert not any(s.endswith(":call") and s != "mcp:tools/call" for s in scopes)
assert not any(s.endswith(":call") for s in scopes)
def test_build_claims_act_fallback_to_litellm_proxy():
@ -338,7 +339,7 @@ async def test_hook_skips_non_mcp_call_types():
user_dict = _make_user_api_key_dict()
data = {"messages": [{"role": "user", "content": "hello"}]}
for call_type in ("completion", "acompletion", "embedding", "list_mcp_tools"):
for call_type in ("completion", "acompletion", "embedding"):
original_data = {**data}
result = await signer.async_pre_call_hook(
user_api_key_dict=user_dict,
@ -351,6 +352,33 @@ async def test_hook_skips_non_mcp_call_types():
), f"extra_headers should not be set for {call_type}"
@pytest.mark.asyncio
async def test_hook_signs_list_mcp_tools():
"""async_pre_call_hook() signs JWT for list_mcp_tools with list scope."""
signer = _make_signer(
issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300
)
user_dict = _make_user_api_key_dict(user_id="alice", team_id="backend")
data = {"mcp_tool_name": "should_be_cleared"}
result = await signer.async_pre_call_hook(
user_api_key_dict=user_dict,
cache=MagicMock(),
data=data,
call_type="list_mcp_tools",
)
assert isinstance(result, dict)
assert "extra_headers" in result
assert result["extra_headers"]["Authorization"].startswith("Bearer ")
token = result["extra_headers"]["Authorization"].removeprefix("Bearer ")
decoded = _decode_unverified(token)
scopes = set(decoded["scope"].split())
assert "mcp:tools/list" in scopes
# List-only JWTs must NOT carry mcp:tools/call — least-privilege
assert "mcp:tools/call" not in scopes
@pytest.mark.asyncio
async def test_signed_token_is_verifiable():
"""The JWT injected by the hook can be verified against the JWKS public key."""
@ -1128,3 +1156,116 @@ async def test_hook_raises_401_when_jwt_verification_fails():
)
assert exc_info.value.status_code == 401
# --- _build_scope branches: call_mcp_tool with empty tool name, list_mcp_tools ---
def test_build_scope_call_type_call_mcp_tool_without_tool_name():
"""call_mcp_tool with empty tool name emits a generic mcp:tools/call only."""
signer = _make_signer()
scope = signer._build_scope("", call_type="call_mcp_tool")
scopes = set(scope.split())
assert scopes == {"mcp:tools/call"}
def test_build_scope_call_type_list_mcp_tools_only_list():
"""list_mcp_tools (no tool) emits only mcp:tools/list, never tools/call."""
signer = _make_signer()
scope = signer._build_scope("", call_type="list_mcp_tools")
scopes = set(scope.split())
assert scopes == {"mcp:tools/list"}
def test_build_scope_default_is_list_only_when_no_call_type():
"""No call_type and no tool falls through to tools/list (least-privilege default)."""
signer = _make_signer()
scope = signer._build_scope("")
scopes = set(scope.split())
assert "mcp:tools/list" in scopes
assert "mcp:tools/call" not in scopes
# --- inject_mcp_jwt_headers_for_upstream ---
@pytest.mark.asyncio
async def test_inject_mcp_jwt_returns_unchanged_when_signer_not_configured():
"""No signer configured -> return a fresh copy of extra_headers untouched."""
import litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer as mod
from litellm.proxy._types import UserAPIKeyAuth
mod._mcp_jwt_signer_instance = None
headers = {"X-Trace-Id": "abc"}
user_dict = UserAPIKeyAuth(api_key="sk-test", user_id="alice")
result = await mod.inject_mcp_jwt_headers_for_upstream(
user_api_key_dict=user_dict,
extra_headers=headers,
)
assert result == headers
assert result is not headers # must be a copy
@pytest.mark.asyncio
async def test_inject_mcp_jwt_returns_unchanged_when_user_dict_none():
"""No user_api_key_dict -> short-circuit without invoking the signer."""
from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
inject_mcp_jwt_headers_for_upstream,
)
_make_signer() # ensure instance is created
result = await inject_mcp_jwt_headers_for_upstream(
user_api_key_dict=None,
extra_headers={"X-Trace-Id": "abc"},
)
assert result == {"X-Trace-Id": "abc"}
@pytest.mark.asyncio
async def test_inject_mcp_jwt_signs_for_list_tools_path():
"""When for_list_tools=True, signer is invoked with list_mcp_tools call_type."""
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
inject_mcp_jwt_headers_for_upstream,
)
_make_signer(issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300)
user_dict = UserAPIKeyAuth(api_key="sk-test", user_id="alice")
result = await inject_mcp_jwt_headers_for_upstream(
user_api_key_dict=user_dict,
extra_headers={"X-Trace": "1"},
raw_headers={"Authorization": "Bearer incoming.opaque.token"},
for_list_tools=True,
)
assert result["X-Trace"] == "1"
assert result["Authorization"].startswith("Bearer ")
token = result["Authorization"].removeprefix("Bearer ")
decoded = _decode_unverified(token)
scopes = set(decoded["scope"].split())
assert scopes == {"mcp:tools/list"}
@pytest.mark.asyncio
async def test_inject_mcp_jwt_signs_for_tool_call_path():
"""for_list_tools=False with a tool name signs a call_mcp_tool JWT."""
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
inject_mcp_jwt_headers_for_upstream,
)
_make_signer(issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300)
user_dict = UserAPIKeyAuth(api_key="sk-test", user_id="alice")
result = await inject_mcp_jwt_headers_for_upstream(
user_api_key_dict=user_dict,
for_list_tools=False,
mcp_tool_name="search_web",
)
assert result["Authorization"].startswith("Bearer ")
token = result["Authorization"].removeprefix("Bearer ")
decoded = _decode_unverified(token)
scopes = set(decoded["scope"].split())
assert "mcp:tools/call" in scopes
assert "mcp:tools/search_web:call" in scopes

View file

@ -2218,6 +2218,7 @@ class TestCLIKeyRegenerationFlow:
# Mock request
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://internal-proxy.local/"
# Test data
session_key = "cli-session-4567890"
@ -2242,11 +2243,14 @@ class TestCLIKeyRegenerationFlow:
"user_code_verified": False,
"session_data": None,
}
mock_request.url_for.return_value = (
"https://test.example.com/sso/cli/complete/cli-session-4567890"
)
with (
patch.dict(
os.environ,
{
"PROXY_BASE_URL": "https://test.example.com",
"SERVER_ROOT_PATH": "",
},
),
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
return_value=mock_user_info,
@ -2290,6 +2294,10 @@ class TestCLIKeyRegenerationFlow:
assert result.status_code == 200
# Verify response contains success message (response is HTML)
assert result.body is not None
assert (
'action="https://test.example.com/sso/cli/complete/cli-session-4567890"'
in result.body.decode()
)
@pytest.mark.asyncio
async def test_cli_poll_key_returns_teams_for_selection(self):

View file

@ -483,6 +483,136 @@ class TestProxyInitializationHelpers:
assert appended_params["connection_limit"] == 5
assert appended_params["pool_timeout"] == expected_timeout
def test_build_db_connection_url_params_defaults(self):
from litellm.proxy.proxy_cli import _build_db_connection_url_params
params = _build_db_connection_url_params(connection_limit=10, pool_timeout=60)
assert params == {"connection_limit": 10, "pool_timeout": 60}
def test_build_db_connection_url_params_omits_none_timeouts(self):
from litellm.proxy.proxy_cli import _build_db_connection_url_params
params = _build_db_connection_url_params(
connection_limit=10,
pool_timeout=60,
connect_timeout=None,
socket_timeout=None,
)
assert "connect_timeout" not in params
assert "socket_timeout" not in params
def test_build_db_connection_url_params_includes_optional_timeouts(self):
from litellm.proxy.proxy_cli import _build_db_connection_url_params
params = _build_db_connection_url_params(
connection_limit=10,
pool_timeout=60,
connect_timeout=15,
socket_timeout=120,
)
assert params["connect_timeout"] == 15
assert params["socket_timeout"] == 120
def test_build_db_connection_url_params_extras_override_defaults(self):
from litellm.proxy.proxy_cli import _build_db_connection_url_params
params = _build_db_connection_url_params(
connection_limit=10,
pool_timeout=60,
extra_params={
"pgbouncer": "true",
"statement_cache_size": 0,
"pool_timeout": 5,
},
)
assert params["pgbouncer"] == "true"
assert params["statement_cache_size"] == 0
assert params["pool_timeout"] == 5
@patch("subprocess.run")
@patch("atexit.register")
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
@patch(
"litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False
)
def test_db_connection_extra_params_forwarded_to_url(
self,
mock_should_update,
mock_setup_db,
mock_atexit_register,
mock_subprocess_run,
):
from click.testing import CliRunner
from litellm.proxy.proxy_cli import run_server
runner = CliRunner()
mock_subprocess_run.return_value = MagicMock(returncode=0)
mock_proxy_module = MagicMock(
app=MagicMock(),
ProxyConfig=MagicMock(),
KeyManagementSettings=MagicMock(),
save_worker_config=MagicMock(),
)
mock_proxy_module.ProxyConfig.return_value.get_config = AsyncMock(
return_value={
"general_settings": {
"database_url": "postgresql://test:test@localhost:5432/test",
"database_connect_timeout": 15,
"database_socket_timeout": 120,
"database_extra_connection_params": {
"pgbouncer": "true",
"statement_cache_size": 0,
},
}
}
)
clean_env = {
k: v
for k, v in os.environ.items()
if k not in ("DATABASE_URL", "DIRECT_URL")
}
with (
patch.dict(os.environ, clean_env, clear=True),
patch.dict(
"sys.modules",
{
"proxy_server": mock_proxy_module,
"litellm.proxy.proxy_server": mock_proxy_module,
},
),
patch(
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args"
) as mock_get_args,
patch(
"litellm.proxy.proxy_cli.append_query_params",
side_effect=lambda url, params: str(url),
) as mock_append_query_params,
):
mock_get_args.return_value = {
"app": "litellm.proxy.proxy_server:app",
"host": "localhost",
"port": 8000,
}
result = runner.invoke(
run_server,
["--local", "--config", "test-config.yaml", "--skip_server_startup"],
)
assert (
result.exit_code == 0
), f"exit_code={result.exit_code}, output={result.output}"
mock_append_query_params.assert_called()
appended_params = mock_append_query_params.call_args.args[1]
assert appended_params["connect_timeout"] == 15
assert appended_params["socket_timeout"] == 120
assert appended_params["pgbouncer"] == "true"
assert appended_params["statement_cache_size"] == 0
@patch("uvicorn.run")
@patch("atexit.register")
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")

View file

@ -5708,6 +5708,7 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
fake_redis = AsyncMock()
fake_redis.async_increment = AsyncMock(side_effect=record_increment)
fake_redis.async_get_cache = AsyncMock(return_value=None) # counter missing
fake_redis.async_set_cache = AsyncMock(return_value=True) # SET NX wins
counter_cache.redis_cache = fake_redis
# Prisma returns spend=42.0 (authoritative) while the stale cached
@ -5744,16 +5745,131 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(
where={"team_id": "team-9"}
)
# Two increments keyed on the counter: seed ($42) then request ($1.50).
# Seed uses SET NX with db_spend (42) — cross-pod safe, no INCR of 42.
# Only the per-request delta (1.5) goes through INCRBYFLOAT.
fake_redis.async_set_cache.assert_awaited_once_with(
key="spend:team:team-9", value=42.0, nx=True
)
writes = [(c["key"], c["value"]) for c in recorded_increments]
assert ("spend:team:team-9", 42.0) in writes
assert ("spend:team:team-9", 1.5) in writes
assert writes == [("spend:team:team-9", 1.5)]
finally:
ps.user_api_key_cache = orig_user
ps.spend_counter_cache = orig_counter
ps.prisma_client = orig_prisma
@pytest.mark.asyncio
async def test_primary_spend_counter_redis_concurrent_seed_does_not_double_seed():
"""Two pods both observing a missing Redis counter must not both
INCRBYFLOAT the full DB spend. SpendCounterReseed.coalesced uses SET NX
so the loser reads the winner's value; final Redis = db_spend, not
2 * db_spend.
The per-counter asyncio.Lock is per-process, so it does NOT coordinate
across pods. We simulate two pods by patching _get_lock to return a
fresh lock per call (each "pod" has its own lock registry in real life).
"""
from litellm.caching.dual_cache import DualCache
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
counter_key = "spend:team:team-concurrent-seed"
redis_store: dict = {}
db_read_count = 0
set_results: list = []
get_after_set_count = 0
set_completed_count = 0
async def redis_set_cache(key, value, nx=False, **_):
# Yield BEFORE the membership check so two concurrent callers
# interleave the way real atomic Redis SET NX does: the first
# to resume runs check + write atomically and wins; the second
# resumes after the key exists and loses. Yielding *after* the
# check would let both callers pass the empty-store check before
# either writes, so neither would ever lose.
await asyncio.sleep(0)
if nx and key in redis_store:
set_results.append(False)
return False
redis_store[key] = float(value)
set_results.append(True)
nonlocal set_completed_count
set_completed_count += 1
return True
async def redis_get_cache(key):
# Track reads that happen after at least one SET NX has completed
# — those are the loser-path fallback reads we want to verify.
if set_completed_count > 0:
nonlocal get_after_set_count
get_after_set_count += 1
return redis_store.get(key)
fake_redis = AsyncMock()
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get_cache)
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
async def slow_find_unique(**_):
nonlocal db_read_count
db_read_count += 1
# Both pods read DB before either's SET NX lands.
await asyncio.sleep(0)
row = MagicMock()
row.spend = 506.0
return row
fake_prisma = MagicMock()
fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(
side_effect=slow_find_unique
)
pod_a = DualCache()
pod_a.redis_cache = fake_redis
pod_b = DualCache()
pod_b.redis_cache = fake_redis
# Each "pod" has its own per-process lock registry. Patch _get_lock to
# always return a fresh lock so the two coalesced calls do not serialize
# via one in-process lock (which is what would happen across pods).
async def fresh_lock(_counter_key):
return asyncio.Lock()
with patch.object(SpendCounterReseed, "_get_lock", side_effect=fresh_lock):
results = await asyncio.gather(
SpendCounterReseed.coalesced(
prisma_client=fake_prisma,
spend_counter_cache=pod_a,
counter_key=counter_key,
),
SpendCounterReseed.coalesced(
prisma_client=fake_prisma,
spend_counter_cache=pod_b,
counter_key=counter_key,
),
)
assert all(r == 506.0 for r in results), results
assert redis_store[counter_key] == pytest.approx(506.0), redis_store
# Both pods read the DB and both attempted SET NX; exactly one wrote
# (winner) and one was rejected (loser).
assert db_read_count == 2
assert fake_redis.async_set_cache.await_count == 2
nx_writes = [
call
for call in fake_redis.async_set_cache.await_args_list
if call.kwargs.get("nx") is True
]
assert len(nx_writes) == 2
assert sorted(set_results) == [False, True], (
f"expected exactly one SET NX winner and one loser, got {set_results}"
)
# Loser path executed: after the winner's SET NX returned True, the
# losing coalesced() call falls back to async_get_cache to read the
# winner's value rather than re-seeding.
assert get_after_set_count >= 1, (
"loser branch (else: read back winner's value) was never exercised"
)
@pytest.mark.asyncio
async def test_reseed_spend_from_db_user_and_org_prefixes():
"""User and org counters reseed from their own DB tables.
@ -5877,9 +5993,16 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
redis_store[key] = (redis_store.get(key) or 0.0) + value
return redis_store[key]
async def redis_set_cache(key, value, nx=False, **_):
if nx and key in redis_store:
return False
redis_store[key] = float(value)
return True
fake_redis = AsyncMock()
fake_redis.async_get_cache = AsyncMock(return_value=None)
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
counter_cache.redis_cache = fake_redis
db_row = MagicMock()
@ -5907,6 +6030,7 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(
where={"team_id": "team-stale-local"}
)
# Seed via SET NX (42) + delta via INCRBYFLOAT (1.5) = 43.5.
assert redis_store[counter_key] == pytest.approx(43.5)
assert counter_cache.in_memory_cache.get_cache(
key=counter_key
@ -6297,14 +6421,14 @@ async def test_get_current_spend_reseeds_from_db_when_counter_missing():
from litellm.proxy.proxy_server import get_current_spend
counter_cache = DualCache()
recorded_warms: list = []
recorded_seeds: list = []
async def record_increment(key, value, ttl=None, **kwargs):
recorded_warms.append({"key": key, "value": value})
return value
async def record_set_cache(key, value, nx=False, **kwargs):
recorded_seeds.append({"key": key, "value": value, "nx": nx})
return True
fake_redis = AsyncMock()
fake_redis.async_increment = AsyncMock(side_effect=record_increment)
fake_redis.async_set_cache = AsyncMock(side_effect=record_set_cache)
fake_redis.async_get_cache = AsyncMock(return_value=None)
counter_cache.redis_cache = fake_redis
@ -6329,9 +6453,9 @@ async def test_get_current_spend_reseeds_from_db_when_counter_missing():
f"expected DB reseed to return 362.0, got {spend} "
f"(fallback would have returned 30.0 and caused bypass)"
)
# Counter warmed so subsequent reads are fast
assert ("spend:team_member:user-1:team-1", 362.0) in [
(w["key"], w["value"]) for w in recorded_warms
# Counter warmed via SET NX so subsequent reads are fast.
assert ("spend:team_member:user-1:team-1", 362.0, True) in [
(s["key"], s["value"], s["nx"]) for s in recorded_seeds
]
assert counter_cache.in_memory_cache.get_cache(
key="spend:team_member:user-1:team-1"
@ -6408,8 +6532,15 @@ async def test_get_current_spend_coalesces_concurrent_reseeds():
redis_store[key] = (redis_store.get(key) or 0.0) + value
return redis_store[key]
async def redis_set_cache(key, value, nx=False, **_):
if nx and key in redis_store:
return False
redis_store[key] = float(value)
return True
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get)
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
counter_cache.redis_cache = fake_redis
fake_prisma = MagicMock()
@ -6516,9 +6647,16 @@ async def test_concurrent_read_and_write_paths_share_one_db_query():
redis_store[key] = (redis_store.get(key) or 0.0) + value
return redis_store[key]
async def redis_set_cache(key, value, nx=False, **_):
if nx and key in redis_store:
return False
redis_store[key] = float(value)
return True
fake_redis = AsyncMock()
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get)
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
counter_cache.redis_cache = fake_redis
fake_prisma = MagicMock()
@ -6621,9 +6759,16 @@ async def test_reseed_warms_cache_even_on_zero_db_spend():
redis_store[key] = (redis_store.get(key) or 0.0) + value
return redis_store[key]
async def redis_set_cache(key, value, nx=False, **_):
if nx and key in redis_store:
return False
redis_store[key] = float(value)
return True
fake_redis = AsyncMock()
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get)
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
counter_cache.redis_cache = fake_redis
db_call_count = 0

View file

@ -2390,3 +2390,34 @@ def test_custom_pricing_without_cache_keys_preserves_legacy_behavior():
expected = 1000 * 0.0000025 + 100 * 0.000015
assert cost == pytest.approx(expected)
def test_openrouter_gemini_3_1_flash_lite_stable_pricing():
"""
Test that openrouter/google/gemini-3.1-flash-lite (stable, no -preview suffix)
has a pricing entry.
Google promoted gemini-3.1-flash-lite to GA on 2026-05-07. PR #27933 added the
stable pricing for the bare, gemini/, and vertex_ai/ prefixes but missed the
openrouter/google/ variant — every other Gemini family in the file has an
openrouter/google/ sibling (2.0-flash-001, 2.5-flash, 2.5-pro, 3-flash-preview,
3-pro-preview, 3.1-flash-lite-preview, 3.1-pro-preview), so the gap is a
consistency issue, not a design choice. Same shape as the preview-variant gap
fixed in PR #25610.
Pricing matches the existing -preview entry one-for-one (input $0.25/M, output
$1.50/M, cache-read $0.025/M) — Google did not change costs at the GA cutover.
"""
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model_name = "openrouter/google/gemini-3.1-flash-lite"
model_info = litellm.model_cost.get(model_name)
assert model_info is not None, f"Missing model pricing entry: {model_name}"
assert model_info["litellm_provider"] == "openrouter"
assert model_info["input_cost_per_token"] == 2.5e-07
assert model_info["output_cost_per_token"] == 1.5e-06
assert model_info["cache_read_input_token_cost"] == 2.5e-08
assert model_info["max_input_tokens"] == 1048576
assert model_info["max_output_tokens"] == 65536

View file

@ -1741,6 +1741,362 @@ async def test_acompletion_streaming_iterator_pre_first_chunk_skips_continuation
assert fallback_kwargs["messages"] == messages
# ---------------------------------------------------------------------------
# Shared helpers for the _aresponses_streaming_iterator test suite.
# ---------------------------------------------------------------------------
def _make_responses_iterator(
*,
chunks=(),
error=None,
bridge=False,
model="gpt-4",
hidden_params=None,
chat_chunks=None,
):
"""Build a minimal mock Responses-API streaming iterator.
Bypasses BaseResponsesAPIStreamingIterator.__init__ but mirrors every
attribute production code reads. Yields *chunks*, then raises *error*
(or StopAsyncIteration). Set bridge=True to inherit from
LiteLLMCompletionStreamingIterator so the wrapper's bridge-path
isinstance check (used by usage extraction) matches.
"""
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
LiteLLMCompletionStreamingIterator,
)
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
)
base = (
LiteLLMCompletionStreamingIterator
if bridge
else BaseResponsesAPIStreamingIterator
)
class _Iter(base):
def __init__(self):
self._chunks = list(chunks)
self._idx = 0
self._hidden_params = hidden_params or {}
self.model = model
self.custom_llm_provider = "anthropic"
self.logging_obj = MagicMock()
self.litellm_metadata = None
self.responses_api_provider_config = None
self.finished = False
self.completed_response = None
self.response = None
self.start_time = None
self.request_data = {}
self.call_type = None
if chat_chunks is not None:
self.collected_chat_completion_chunks = chat_chunks
def __aiter__(self):
return self
async def __anext__(self):
if self._idx < len(self._chunks):
self._idx += 1
return self._chunks[self._idx - 1]
if error is not None:
raise error
raise StopAsyncIteration
return _Iter()
class _AsyncList:
"""Generic async iterator over a list — used as the fallback response."""
def __init__(self, items=()):
self._items = list(items)
self._idx = 0
def __aiter__(self):
return self
async def __anext__(self):
if self._idx >= len(self._items):
raise StopAsyncIteration
item = self._items[self._idx]
self._idx += 1
return item
def _make_router_with_fallback(primary="gpt-4", secondary="gpt-3.5-turbo"):
return litellm.Router(
model_list=[
{
"model_name": primary,
"litellm_params": {"model": primary, "api_key": "k1"},
},
{
"model_name": secondary,
"litellm_params": {"model": secondary, "api_key": "k2"},
},
],
fallbacks=[{primary: [secondary]}],
)
@pytest.mark.asyncio
async def test_aresponses_streaming_iterator_fallback():
"""Catches MidStreamFallbackError, re-enters the fallback chain via
async_function_with_fallbacks_common_utils with the per-attempt helper
and original_generic_function preserved. Mirrors
test_acompletion_streaming_iterator for the aresponses path."""
from litellm.exceptions import MidStreamFallbackError
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
)
router = _make_router_with_fallback(
"anthropic/claude-sonnet-4-6", "vertex_ai/claude-sonnet-4-6"
)
src = _make_responses_iterator(
chunks=[MagicMock(type="response.created")],
error=MidStreamFallbackError(
message="anthropic socket timeout",
model="anthropic/claude-sonnet-4-6",
llm_provider="anthropic",
is_pre_first_chunk=False,
generated_content="",
),
model="anthropic/claude-sonnet-4-6",
hidden_params={"model_id": "src-deployment-1"},
)
fallback_chunks = [
MagicMock(type="response.output_text.delta"),
MagicMock(type="response.completed"),
]
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
return_value=_AsyncList(fallback_chunks),
) as mock_fallback_utils:
wrapped = await router._aresponses_streaming_iterator(
response=src,
initial_kwargs={
"model": "anthropic/claude-sonnet-4-6",
"stream": True,
"input": "Hi",
"original_generic_function": litellm.aresponses,
},
)
assert isinstance(wrapped, BaseResponsesAPIStreamingIterator)
assert wrapped._hidden_params.get("model_id") == "src-deployment-1"
collected = [c async for c in wrapped]
assert len(collected) == 3 # 1 primary chunk + 2 fallback chunks
call_kwargs = mock_fallback_utils.call_args.kwargs
fbk = call_kwargs["kwargs"]
# Bound methods compare equal when they share the same instance + __func__.
assert fbk["original_function"] == router._ageneric_api_call_with_fallbacks_helper
assert fbk["original_generic_function"] is litellm.aresponses
assert call_kwargs["model_group"] == "anthropic/claude-sonnet-4-6"
assert call_kwargs["disable_fallbacks"] is False
@pytest.mark.asyncio
async def test_aresponses_streaming_iterator_writes_litellm_metadata_on_fallback():
"""Regression: model_group must land under "litellm_metadata" (the key
litellm.aresponses reads), not the default "metadata"."""
from litellm.exceptions import MidStreamFallbackError
router = _make_router_with_fallback()
src = _make_responses_iterator(
error=MidStreamFallbackError(
message="boom",
model="gpt-4",
llm_provider="anthropic",
is_pre_first_chunk=True,
generated_content="",
)
)
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
return_value=_AsyncList(),
) as mock_fallback_utils:
wrapped = await router._aresponses_streaming_iterator(
response=src,
initial_kwargs={
"model": "gpt-4",
"stream": True,
"input": "Hello",
"original_generic_function": litellm.aresponses,
},
)
async for _ in wrapped:
pass
fbk = mock_fallback_utils.call_args.kwargs["kwargs"]
assert "litellm_metadata" in fbk, "wrong metadata_variable_name"
assert fbk["litellm_metadata"]["model_group"] == "gpt-4"
assert "model_group" not in fbk.get(
"metadata", {}
), "model_group leaked into 'metadata' instead of 'litellm_metadata'"
@pytest.mark.asyncio
async def test_aresponses_streaming_iterator_pre_first_chunk_skips_continuation():
"""Pre-first-chunk error: original input is preserved unchanged."""
from litellm.exceptions import MidStreamFallbackError
router = _make_router_with_fallback()
src = _make_responses_iterator(
error=MidStreamFallbackError(
message="socket timeout before first chunk",
model="gpt-4",
llm_provider="anthropic",
is_pre_first_chunk=True,
generated_content="",
)
)
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
return_value=_AsyncList(),
) as mock_fallback_utils:
wrapped = await router._aresponses_streaming_iterator(
response=src,
initial_kwargs={
"model": "gpt-4",
"stream": True,
"input": "Hello",
"original_generic_function": litellm.aresponses,
},
)
async for _ in wrapped:
pass
fbk = mock_fallback_utils.call_args.kwargs["kwargs"]
assert fbk["input"] == "Hello" # original input, no continuation messages
@pytest.mark.asyncio
async def test_aresponses_streaming_iterator_partial_content_injects_continuation():
"""Mid-stream error: input is rewritten to include user prompt +
developer instruction + prior assistant message with partial output."""
from litellm.exceptions import MidStreamFallbackError
router = _make_router_with_fallback()
src = _make_responses_iterator(
chunks=[MagicMock(type="response.output_text.delta")],
error=MidStreamFallbackError(
message="socket reset mid-stream",
model="gpt-4",
llm_provider="anthropic",
is_pre_first_chunk=False,
generated_content="The capital of France is",
),
)
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
return_value=_AsyncList(),
) as mock_fallback_utils:
wrapped = await router._aresponses_streaming_iterator(
response=src,
initial_kwargs={
"model": "gpt-4",
"stream": True,
"input": "What's the capital of France?",
"original_generic_function": litellm.aresponses,
},
)
async for _ in wrapped:
pass
new_input = mock_fallback_utils.call_args.kwargs["kwargs"]["input"]
assert isinstance(new_input, list)
assert new_input[0]["role"] == "user"
assert new_input[0]["content"][0]["text"] == "What's the capital of France?"
assert new_input[1]["role"] == "developer"
assert "do not repeat" in new_input[1]["content"][0]["text"].lower()
assert new_input[2]["role"] == "assistant"
assert new_input[2]["content"][0]["type"] == "output_text"
assert new_input[2]["content"][0]["text"] == "The capital of France is"
@pytest.mark.asyncio
async def test_aresponses_streaming_iterator_combines_partial_usage():
"""Partial usage from the bridge path is normalized to ResponseAPIUsage
and summed onto the fallback's response.completed event — no token-name
split, clean ResponseAPIUsage on output."""
from types import SimpleNamespace
from litellm.exceptions import MidStreamFallbackError
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
router = _make_router_with_fallback()
src = _make_responses_iterator(
bridge=True,
chat_chunks=[MagicMock()],
chunks=[MagicMock(type="response.output_text.delta")],
error=MidStreamFallbackError(
message="boom",
model="gpt-4",
llm_provider="anthropic",
is_pre_first_chunk=False,
generated_content="hello",
),
)
fallback_response_object = ResponsesAPIResponse(
id="resp_test", created_at=0, model="gpt-4", object="response", output=[]
)
fallback_response_object.usage = ResponseAPIUsage(
input_tokens=20, output_tokens=15, total_tokens=35
)
fallback_event = ResponseCompletedEvent(
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
response=fallback_response_object,
)
with (
patch(
"litellm.main.stream_chunk_builder",
return_value=SimpleNamespace(
usage=SimpleNamespace(prompt_tokens=10, completion_tokens=4)
),
),
patch.object(
router,
"async_function_with_fallbacks_common_utils",
return_value=_AsyncList([fallback_event]),
),
):
wrapped = await router._aresponses_streaming_iterator(
response=src,
initial_kwargs={
"model": "gpt-4",
"stream": True,
"input": "hi",
"original_generic_function": litellm.aresponses,
},
)
async for _ in wrapped:
pass
merged = fallback_response_object.usage
assert isinstance(merged, ResponseAPIUsage)
assert merged.input_tokens == 30 # 10 (translated from prompt_tokens) + 20
assert merged.output_tokens == 19 # 4 (translated from completion_tokens) + 15
assert merged.total_tokens == 49
@pytest.mark.asyncio
async def test_async_function_with_fallbacks_common_utils():
"""Test the async_function_with_fallbacks_common_utils method"""
@ -3863,7 +4219,15 @@ def test_is_deployment_blocked_static_helper_reflects_blocked_flag():
# No model_info on deployment object → treated as not blocked
assert litellm.Router._is_deployment_blocked(object()) is False
missing_blocked = types.SimpleNamespace()
assert litellm.Router._is_deployment_blocked(types.SimpleNamespace(model_info=missing_blocked)) is False
assert litellm.Router._is_deployment_blocked(
types.SimpleNamespace(model_info=types.SimpleNamespace(blocked=True))
) is True
assert (
litellm.Router._is_deployment_blocked(
types.SimpleNamespace(model_info=missing_blocked)
)
is False
)
assert (
litellm.Router._is_deployment_blocked(
types.SimpleNamespace(model_info=types.SimpleNamespace(blocked=True))
)
is True
)

View file

@ -100,6 +100,9 @@ async def get_spend_logs(session, request_id=None, api_key=None):
return await response.json()
@pytest.mark.skip(
reason="Flaky in CI: /spend/logs?request_id=... returns 500 even after a 20s wait for the spend log to be written. Spend-log accuracy is covered by tests/test_litellm/proxy/spend_tracking/ and the proxy_spend_accuracy_tests CircleCI job."
)
@pytest.mark.asyncio
async def test_spend_logs():
"""

View file

@ -137,9 +137,7 @@ def test_add_single_member(api_client, new_team):
@pytest.mark.skip(
reason="Flaky on CI: /team/info intermittently returns 404 immediately after a "
"successful /team/new + /team/member_add. The single-member happy path is covered "
"by test_add_single_member; we lose only the 10-iteration stress loop here."
reason="Flaky in CI: /team/info?team_id=... intermittently returns 404/400 mid-loop after add_team_member calls. Single-member coverage in test_add_single_member is sufficient; team-member CRUD is also covered by tests/test_litellm/proxy/management_endpoints/."
)
def test_add_multiple_members(api_client, new_team):
"""Test adding multiple members to a new team"""

View file

@ -28,6 +28,11 @@ export default defineConfig({
/* Action timeout for clicks, fills, waitForSelector, etc. */
actionTimeout: 15 * 1000,
navigationTimeout: 30 * 1000,
/* Slow down actions when SLOWMO=<ms> is set, useful for headed local debugging */
launchOptions: {
slowMo: process.env.SLOWMO ? (parseInt(process.env.SLOWMO, 10) || 0) : 0,
},
},
/* Configure projects for major browsers */

View file

@ -15,7 +15,7 @@ set -euo pipefail
# In CI (CI=true), expects:
# - PostgreSQL already running on 127.0.0.1:5432
# - DATABASE_URL already set
# - Python/Poetry already installed
# - Python/uv already installed
# - Node.js/npx already available
# ================================================================
@ -48,7 +48,7 @@ cleanup() {
trap cleanup EXIT INT TERM
# --- Pre-flight checks ---
for cmd in python3 npx poetry; do
for cmd in python3 npx uv; do
command -v "$cmd" >/dev/null 2>&1 || { echo "Error: $cmd not found."; exit 1; }
done
@ -117,19 +117,15 @@ echo "UI build copied and restructured"
# --- Python environment ---
echo "=== Setting up Python environment ==="
cd "$REPO_ROOT"
if ! poetry run python3 -c "import prisma" 2>/dev/null; then
echo "Installing Python dependencies (first run)..."
poetry install --with dev,proxy-dev --extras "proxy" --quiet
poetry run pip install nodejs-wheel-binaries 2>/dev/null || true
poetry run prisma generate --schema litellm/proxy/schema.prisma
fi
uv sync --group dev --group proxy-dev --extra proxy --frozen --quiet
uv run --no-sync python -m prisma generate --schema litellm/proxy/schema.prisma
echo "=== Pushing Prisma schema to database ==="
poetry run prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
uv run --no-sync python -m prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
# --- Mock LLM server ---
echo "=== Starting mock LLM server ==="
poetry run python3 "$SCRIPT_DIR/fixtures/mock_llm_server/server.py" &
uv run --no-sync python "$SCRIPT_DIR/fixtures/mock_llm_server/server.py" &
MOCK_PID=$!
for i in $(seq 1 15); do
@ -140,7 +136,7 @@ done
# --- LiteLLM proxy ---
echo "=== Starting LiteLLM proxy ==="
cd "$REPO_ROOT"
poetry run python3 -m litellm.proxy.proxy_cli \
uv run --no-sync python -m litellm.proxy.proxy_cli \
--config "$SCRIPT_DIR/fixtures/config.yml" \
--port 4000 &
PROXY_PID=$!

View file

@ -126,4 +126,84 @@ test.describe("Proxy Admin - Keys", () => {
await expect(page.getByText(E2E_INTERNAL_USER_KEY_ALIAS)).toBeVisible({ timeout: 10_000 });
});
test("Create a key with All Proxy Models (no team)", async ({ page }) => {
await navigateToPage(page, Page.ApiKeys);
await dismissFeedbackPopup(page);
await page.getByRole("button", { name: /Create New Key/i }).click();
await expect(page.getByText("Key Ownership")).toBeVisible({ timeout: 10_000 });
const keyName = `e2e-admin-allproxy-${Date.now()}`;
await page.getByTestId("base-input").fill(keyName);
// No team selection — leave team dropdown empty so the key is owned by the admin user
// Select models — open the multi-select and pick the all-models meta-option.
// The Create Key modal labels this "All Team Models" even when no team is selected
// (see src/components/organisms/create_key_button.tsx:944), unlike the team/user
// settings screens which use "All Proxy Models".
await page.locator(".ant-select-selection-overflow").click();
await page.locator(".ant-select-dropdown:visible").getByText("All Team Models").click();
await page.keyboard.press("Escape");
await page.getByRole("button", { name: "Create Key", exact: true }).click();
await expect(page.getByText("Save your Key")).toBeVisible({ timeout: 10_000 });
await page.keyboard.press("Escape");
await expect(page.getByText(keyName)).toBeVisible({ timeout: 10_000 });
});
test("Create a key with a specific proxy model (no team)", async ({ page }) => {
await navigateToPage(page, Page.ApiKeys);
await dismissFeedbackPopup(page);
await page.getByRole("button", { name: /Create New Key/i }).click();
await expect(page.getByText("Key Ownership")).toBeVisible({ timeout: 10_000 });
const keyName = `e2e-admin-specific-${Date.now()}`;
await page.getByTestId("base-input").fill(keyName);
// Open the model multi-select and pick a single specific model. Use
// getByRole("option", ...) to avoid the strict-mode collision between
// the option container and its inner text node.
const modelName = "fake-openai-gpt-4";
await page.locator(".ant-select-selection-overflow").click();
const option = page.locator(".ant-select-dropdown:visible").getByRole("option", { name: modelName, exact: true });
await option.waitFor({ state: "attached" });
// Dispatch the click via the DOM — antd's dropdown can render the option
// off-viewport during the open animation, which trips Playwright's
// visibility/stability checks. The click handler fires regardless.
await option.evaluate((el: HTMLElement) => el.click());
await page.keyboard.press("Escape");
await page.getByRole("button", { name: "Create Key", exact: true }).click();
await expect(page.getByText("Save your Key")).toBeVisible({ timeout: 10_000 });
// Grab the new key from the success modal (rendered inside a <pre>) and
// verify it can call /chat/completions for the model it was scoped to.
// The mock LLM server (fixtures/mock_llm_server/server.py) replies with
// a fixed "This is a mock response." body.
const apiKey = (await page.locator(".ant-modal:visible pre").innerText()).trim();
expect(apiKey).toMatch(/^sk-/);
const response = await page.request.post("/chat/completions", {
headers: { Authorization: `Bearer ${apiKey}` },
data: {
model: modelName,
messages: [{ role: "user", content: "ping" }],
},
});
expect(response.status()).toBe(200);
const body = await response.json();
expect(body.choices?.[0]?.message?.content).toBe("This is a mock response.");
await page.keyboard.press("Escape");
await expect(page.getByText(keyName)).toBeVisible({ timeout: 10_000 });
});
});

View file

@ -7,7 +7,7 @@ import { columns } from "@/components/molecules/models/columns";
import { getDisplayModelName } from "@/components/view_model/model_name_display";
import DeleteResourceModal from "@/components/common_components/DeleteResourceModal";
import NotificationsManager from "@/components/molecules/notifications_manager";
import { modelDeleteCall } from "@/components/networking";
import { modelDeleteCall, modelPatchUpdateCall } from "@/components/networking";
import { InfoCircleOutlined, SettingOutlined } from "@ant-design/icons";
import { PaginationState, SortingState } from "@tanstack/react-table";
import { useQueryClient } from "@tanstack/react-query";
@ -220,6 +220,25 @@ const AllModelsTab = ({
}
};
const [pausingModelId, setPausingModelId] = useState<string | null>(null);
const handleTogglePause = async (modelId: string, blocked: boolean) => {
if (!accessToken) return;
try {
setPausingModelId(modelId);
await modelPatchUpdateCall(accessToken, { blocked }, modelId);
NotificationsManager.success(blocked ? "Model paused" : "Model resumed");
// invalidateQueries already schedules a refetch for active observers
// on this key — no need to also call refetchModels() (would double-fetch).
queryClient.invalidateQueries({ queryKey: ["models", "list"] });
} catch (error) {
console.error("Error toggling model pause state:", error);
NotificationsManager.fromBackend(error);
} finally {
setPausingModelId(null);
}
};
return (
<TabPanel>
<Grid>
@ -536,6 +555,8 @@ const AllModelsTab = ({
expandedRows,
setExpandedRows,
setDeleteModalModelId,
handleTogglePause,
pausingModelId,
)}
data={filteredData}
isLoading={isLoadingModelsInfo}

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