Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/silly-wright-1b8559

This commit is contained in:
Yuneng Jiang 2026-05-20 19:10:00 -07:00
commit 63295a5ff5
No known key found for this signature in database
32 changed files with 2165 additions and 75 deletions

View file

@ -1611,11 +1611,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
"masked_entity_count", safe_dumps(masked_entity_count)
)
self.safe_set_attribute(
span=guardrail_span,
key="guardrail_response",
value=guardrail_information.get("guardrail_response"),
)
guardrail_response = guardrail_information.get("guardrail_response")
if guardrail_response is not None:
guardrail_span.set_attribute(
"guardrail_response", safe_dumps(guardrail_response)
)
self._set_team_attributes_from_kwargs(guardrail_span, kwargs)

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

@ -27296,6 +27296,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

@ -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

@ -12,7 +12,7 @@ import fnmatch
import re
import secrets
from datetime import datetime, timezone
from typing import Any, Iterator, List, Optional, Tuple, Union, cast
from typing import Any, Dict, Iterator, List, Optional, Tuple, Union, cast
import fastapi
from fastapi import HTTPException, Request, WebSocket, status
@ -333,8 +333,22 @@ def _apply_budget_limits_to_end_user_params(
async def user_api_key_auth_websocket(websocket: WebSocket):
# Accept the WebSocket connection
scope_headers = list(websocket.scope.get("headers") or [])
request = Request(scope={"type": "http", "headers": scope_headers})
ws_scope = websocket.scope or {}
scope_headers = list(ws_scope.get("headers") or [])
# ``get_request_route`` falls back to ``request.url.path`` when
# ``scope["path"]`` is absent. On WebSockets that fallback reads
# ``websocket.url``, which Starlette reconstructs from the (poisonable)
# Host header. Carry the ASGI scope's path / root_path so the lookup
# never reaches the fallback.
synthetic_scope: Dict[str, Any] = {
"type": "http",
"headers": scope_headers,
"path": ws_scope.get("path", ""),
}
for key in ("root_path", "app_root_path"):
if key in ws_scope:
synthetic_scope[key] = ws_scope[key]
request = Request(scope=synthetic_scope)
request._url = websocket.url

View file

@ -13,6 +13,34 @@ from litellm.proxy.common_utils.callback_utils import (
from litellm.types.router import Deployment
_FORM_CONTENT_TYPES: frozenset[str] = frozenset(
{"application/x-www-form-urlencoded", "multipart/form-data"}
)
def _normalize_media_type(content_type: str) -> str:
"""Return the bare media type per RFC 7231: strip params, trim, lowercase."""
if not content_type:
return ""
return content_type.split(";", 1)[0].strip().lower()
def _is_form_content_type(content_type: str) -> bool:
"""
True iff Starlette's ``request.form()`` will actually parse this body.
Substring matching ``"form"`` is unsafe: ``request.form()`` returns empty
``FormData`` for non-canonical types without consuming the body, leaving
the auth-time pre-read and the handler's read seeing different payloads.
"""
return _normalize_media_type(content_type) in _FORM_CONTENT_TYPES
def _is_json_content_type(content_type: str) -> bool:
"""True iff the body should be parsed as JSON."""
return _normalize_media_type(content_type) == "application/json"
async def _read_request_body(request: Optional[Request]) -> Dict:
"""
Safely read the request body and parse it as JSON.
@ -37,8 +65,24 @@ async def _read_request_body(request: Optional[Request]) -> Dict:
_request_headers: dict = _safe_get_request_headers(request=request)
content_type = _request_headers.get("content-type", "")
if "form" in content_type:
parsed_body = dict(await request.form())
if _is_form_content_type(content_type):
try:
form_data = await request.form()
except Exception as e:
# ``request.form()`` raises on malformed multipart (missing
# boundary, malformed chunk encoding, …). Surface as 400 so
# the auth-time pre-read does not silently cache ``{}`` while
# a later raw-body re-read sees the original payload —
# banned-param checks must see the same body the handler
# acts on.
verbose_proxy_logger.error(f"Invalid form payload: {e}")
raise ProxyException(
message=f"Invalid form payload: {e}",
type="invalid_request_error",
param="request_body",
code=status.HTTP_400_BAD_REQUEST,
)
parsed_body = dict(form_data)
if "metadata" in parsed_body and isinstance(parsed_body["metadata"], str):
parsed_body["metadata"] = json.loads(parsed_body["metadata"])
else:
@ -306,18 +350,13 @@ async def get_request_body(request: Request) -> Dict[str, Any]:
Read the request body and parse it as JSON.
"""
if request.method == "POST":
if request.headers.get("content-type", "") == "application/json":
content_type = request.headers.get("content-type", "")
if _is_json_content_type(content_type):
return await _read_request_body(request)
elif "multipart/form-data" in request.headers.get(
"content-type", ""
) or "application/x-www-form-urlencoded" in request.headers.get(
"content-type", ""
):
elif _is_form_content_type(content_type):
return await get_form_data(request)
else:
raise ValueError(
f"Unsupported content type: {request.headers.get('content-type')}"
)
raise ValueError(f"Unsupported content type: {content_type}")
return {}

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

@ -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

@ -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",

View file

@ -27296,6 +27296,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,
@ -28105,10 +28157,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,
@ -28118,7 +28170,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

@ -79,7 +79,7 @@ class RealTimeWebSocketClient:
def _is_initial_event(self, msg_type: str) -> bool:
"""Check if message type is an initial connection event"""
# OpenAI sends "session.created", xAI sends "conversation.created"
# OpenAI and xAI send "session.created"; some providers send "conversation.created"
return msg_type in ["session.created", "conversation.created"]
async def receive_text(self):

View file

@ -19,8 +19,8 @@ class TestXAIRealtime(BaseRealtimeTest):
"""
E2E tests for xAI Realtime API.
xAI's Grok Voice Agent API is OpenAI-compatible but uses:
- Different initial event: "conversation.created" instead of "session.created"
xAI's Grok Voice Agent API is OpenAI-compatible:
- Initial event: "session.created" (matches OpenAI)
- Different endpoint: wss://api.x.ai/v1/realtime
- Model: grok-4-1-fast-non-reasoning
"""
@ -32,4 +32,4 @@ class TestXAIRealtime(BaseRealtimeTest):
return "XAI_API_KEY"
def get_initial_event_type(self) -> str:
return "conversation.created"
return "session.created"

View file

@ -915,6 +915,36 @@ async def test_user_api_key_auth_websocket():
)
@pytest.mark.asyncio
async def test_user_api_key_auth_websocket_carries_asgi_path():
"""
The synthetic Request must carry the ASGI scope's ``path`` so
``get_request_route`` returns the real WebSocket path, not a value
reconstructed from the (Host-poisonable) ``websocket.url``.
"""
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_websocket
mock_websocket = MagicMock(spec=WebSocket)
mock_websocket.query_params = {"model": "some_model"}
mock_websocket.headers = {"authorization": "Bearer some_api_key"}
mock_websocket.scope = {
"type": "websocket",
"path": "/v1/realtime",
"root_path": "",
"headers": [(b"authorization", b"Bearer some_api_key")],
}
mock_websocket.url = URL(url="/v1/realtime")
with patch(
"litellm.proxy.auth.user_api_key_auth.user_api_key_auth", autospec=True
) as mock_user_api_key_auth:
await user_api_key_auth_websocket(mock_websocket)
request_arg = mock_user_api_key_auth.call_args.kwargs["request"]
assert request_arg.scope.get("path") == "/v1/realtime"
assert request_arg.scope.get("root_path") == ""
@pytest.mark.parametrize("enforce_rbac", [True, False])
@pytest.mark.asyncio
async def test_jwt_user_api_key_auth_builder_enforce_rbac(enforce_rbac, monkeypatch):

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

@ -66,7 +66,7 @@ class TestOpenTelemetryGuardrails(unittest.TestCase):
mock_span.set_attribute.assert_any_call("guardrail_name", "test_guardrail")
mock_span.set_attribute.assert_any_call("guardrail_mode", "input")
mock_span.set_attribute.assert_any_call(
"guardrail_response", "filtered_content"
"guardrail_response", safe_dumps("filtered_content")
)
mock_span.set_attribute.assert_any_call(
"masked_entity_count", safe_dumps({"CREDIT_CARD": 2})
@ -87,6 +87,65 @@ class TestOpenTelemetryGuardrails(unittest.TestCase):
# Verify that start_span was never called
otel.tracer.start_span.assert_not_called()
@patch("litellm.integrations.opentelemetry.datetime")
def test_guardrail_response_dict_is_json_serialized(self, mock_datetime):
"""Dict guardrail_response (e.g. OpenAI moderation result) must reach
the span as a JSON string so downstream pipelines can parse it for
metric extraction — this is the bug the PR fixes."""
otel = OpenTelemetry()
otel.tracer = MagicMock()
mock_span = MagicMock()
otel.tracer.start_span.return_value = mock_span
moderation_payload = {
"id": "modr-7740",
"model": "omni-moderation-latest",
"results": [{"categories": {"harassment": False}}],
}
guardrail_info = {
"guardrail_name": "test_guardrail",
"guardrail_mode": "input",
"guardrail_response": moderation_payload,
"start_time": 1609459200.0,
"end_time": 1609459201.0,
}
kwargs = {
"standard_logging_object": {"guardrail_information": [guardrail_info]}
}
otel._create_guardrail_span(kwargs=kwargs, context=None)
mock_span.set_attribute.assert_any_call(
"guardrail_response", safe_dumps(moderation_payload)
)
@patch("litellm.integrations.opentelemetry.datetime")
def test_guardrail_response_none_is_skipped(self, mock_datetime):
"""When guardrail_response is None, the attribute must not be set —
guards against round-tripping ``"null"`` into traces."""
otel = OpenTelemetry()
otel.tracer = MagicMock()
mock_span = MagicMock()
otel.tracer.start_span.return_value = mock_span
guardrail_info = {
"guardrail_name": "test_guardrail",
"guardrail_mode": "input",
"guardrail_response": None,
"start_time": 1609459200.0,
"end_time": 1609459201.0,
}
kwargs = {
"standard_logging_object": {"guardrail_information": [guardrail_info]}
}
otel._create_guardrail_span(kwargs=kwargs, context=None)
attribute_keys = [
call.args[0] for call in mock_span.set_attribute.call_args_list
]
self.assertNotIn("guardrail_response", attribute_keys)
class TestOpenTelemetryTeamAttributesOnChildSpans(unittest.TestCase):
"""team_id / team_alias must land on every child span of a
@ -1169,7 +1228,7 @@ class TestOpenTelemetry(unittest.TestCase):
mock_span.set_attribute.assert_any_call("guardrail_name", "test_guardrail")
mock_span.set_attribute.assert_any_call("guardrail_mode", "input")
mock_span.set_attribute.assert_any_call(
"guardrail_response", "filtered_content"
"guardrail_response", safe_dumps("filtered_content")
)
mock_span.set_attribute.assert_any_call(
"masked_entity_count", safe_dumps({"CREDIT_CARD": 2})

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

@ -16,6 +16,7 @@ sys.path.insert(
import litellm
from litellm.proxy._types import ProxyException
from litellm.proxy.common_utils.http_parsing_utils import (
_is_form_content_type,
_read_request_body,
_safe_get_request_headers,
_safe_get_request_parsed_body,
@ -853,3 +854,145 @@ class TestGetTagsFromRequestBodyStringCoerce:
tags = get_tags_from_request_body({"metadata": {"tags": ["x"]}})
assert tags == ["x"]
class TestIsFormContentType:
@pytest.mark.parametrize(
"content_type",
[
"application/x-www-form-urlencoded",
"multipart/form-data",
"multipart/form-data; boundary=----WebKitFormBoundary",
"Application/X-WWW-Form-Urlencoded",
" multipart/form-data ",
"application/x-www-form-urlencoded; charset=utf-8",
],
)
def test_form_types_match(self, content_type):
assert _is_form_content_type(content_type) is True
@pytest.mark.parametrize(
"content_type",
[
"",
"application/json",
"application/json; charset=utf-8",
"application/form-json",
"multiform/anything",
"application/json; xform=1",
"application/xml-with-form-data-but-not-actually",
"text/plain",
"form",
],
)
def test_non_form_types_rejected(self, content_type):
assert _is_form_content_type(content_type) is False
class TestReadRequestBodyNonCanonicalContentType:
"""A JSON body with a ``"form"``-substring Content-Type must parse as JSON."""
@pytest.mark.asyncio
@pytest.mark.parametrize(
"content_type",
[
"application/form-json",
"application/json; xform=1",
"multiform/anything",
],
)
async def test_json_body_with_formlike_content_type_parses_as_json(
self, content_type
):
payload = {"user_config": {"model_list": []}, "model": "x"}
mock_request = MagicMock()
mock_request.body = AsyncMock(return_value=orjson.dumps(payload))
mock_request.form = AsyncMock(return_value={})
mock_request.headers = {"content-type": content_type}
mock_request.scope = {}
result = await _read_request_body(mock_request)
assert result == payload
mock_request.form.assert_not_called()
@pytest.mark.asyncio
async def test_real_form_post_still_parsed_as_form(self):
mock_request = MagicMock()
mock_request.form = AsyncMock(return_value={"k": "v"})
mock_request.body = AsyncMock(return_value=b"")
mock_request.headers = {"content-type": "application/x-www-form-urlencoded"}
mock_request.scope = {}
result = await _read_request_body(mock_request)
assert result == {"k": "v"}
mock_request.form.assert_awaited_once()
class TestReadRequestBodyFormParseFailure:
"""
A failed ``request.form()`` parse (e.g. multipart with missing boundary)
must surface as a 400, not silently return ``{}`` — otherwise the
auth-time pre-read sees an empty body while a later raw-body re-read
sees the original payload, defeating every banned-param check.
"""
@pytest.mark.asyncio
@pytest.mark.parametrize(
"raised_exception",
[
ValueError("Missing boundary in multipart."),
AssertionError("malformed chunk"),
RuntimeError("form parser exploded"),
],
)
async def test_form_parse_failure_raises_400(self, raised_exception):
mock_request = MagicMock()
mock_request.form = AsyncMock(side_effect=raised_exception)
mock_request.headers = {"content-type": "multipart/form-data"}
mock_request.scope = {}
with pytest.raises(ProxyException) as exc_info:
await _read_request_body(mock_request)
assert str(exc_info.value.code) == "400"
class TestGetRequestBody:
@pytest.mark.asyncio
async def test_json_with_charset_param_parses_as_json(self):
payload = {"k": "v"}
mock_request = MagicMock()
mock_request.method = "POST"
mock_request.body = AsyncMock(return_value=orjson.dumps(payload))
mock_request.headers = {"content-type": "application/json; charset=utf-8"}
mock_request.scope = {}
result = await get_request_body(mock_request)
assert result == payload
@pytest.mark.asyncio
async def test_form_post_routes_to_form_data(self):
mock_request = MagicMock()
mock_request.method = "POST"
mock_request.headers = {"content-type": "multipart/form-data; boundary=x"}
mock_request.form = AsyncMock(return_value={"k": "v"})
mock_request.scope = {}
result = await get_request_body(mock_request)
assert result == {"k": "v"}
@pytest.mark.asyncio
async def test_substring_match_no_longer_accepted(self):
mock_request = MagicMock()
mock_request.method = "POST"
mock_request.headers = {"content-type": "application/form-json"}
mock_request.scope = {}
with pytest.raises(ValueError, match="Unsupported content type"):
await get_request_body(mock_request)
@pytest.mark.asyncio
async def test_non_post_returns_empty(self):
mock_request = MagicMock()
mock_request.method = "GET"
assert await get_request_body(mock_request) == {}

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

@ -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

@ -136,6 +136,9 @@ def test_add_single_member(api_client, new_team):
), f"Team size did not increase by 1 (was {initial_size}, now {updated_size})"
@pytest.mark.skip(
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"""
# Get initial team size

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}

View file

@ -182,16 +182,18 @@ export function ToolTestPanel({
Object.entries(values).forEach(([key, value]) => {
const prop = schemaToUse.properties?.[key];
if (prop && value !== null && value !== undefined && value !== "") {
// Strip leading/trailing whitespace from string inputs before submitting
const normalizedValue = typeof value === "string" ? value.trim() : value;
if (prop && normalizedValue !== null && normalizedValue !== undefined && normalizedValue !== "") {
switch (prop.type) {
case "boolean":
convertedValues[key] = value === "true" || value === true;
convertedValues[key] = normalizedValue === "true" || normalizedValue === true;
break;
case "number":
case "integer": {
const numericValue = Number(value);
const numericValue = Number(normalizedValue);
convertedValues[key] = Number.isNaN(numericValue)
? value
? normalizedValue
: prop.type === "integer"
? Math.trunc(numericValue)
: numericValue;
@ -200,28 +202,28 @@ export function ToolTestPanel({
case "object":
case "array": {
try {
const parsed = typeof value === "string" ? JSON.parse(value) : value;
const parsed = typeof normalizedValue === "string" ? JSON.parse(normalizedValue) : normalizedValue;
const isValidObject =
prop.type === "object" && parsed !== null && typeof parsed === "object" && !Array.isArray(parsed);
const isValidArray = prop.type === "array" && Array.isArray(parsed);
if ((prop.type === "object" && isValidObject) || (prop.type === "array" && isValidArray)) {
convertedValues[key] = parsed;
} else {
convertedValues[key] = value;
convertedValues[key] = normalizedValue;
}
} catch (err) {
convertedValues[key] = value;
convertedValues[key] = normalizedValue;
}
break;
}
case "string":
convertedValues[key] = String(value);
convertedValues[key] = String(normalizedValue);
break;
default:
convertedValues[key] = value;
convertedValues[key] = normalizedValue;
}
} else if (value !== null && value !== undefined && value !== "") {
convertedValues[key] = value;
} else if (normalizedValue !== null && normalizedValue !== undefined && normalizedValue !== "") {
convertedValues[key] = normalizedValue;
}
});

View file

@ -6,6 +6,7 @@ export interface ModelInfo {
team_id: string;
db_model: boolean;
access_groups: string[] | null;
blocked?: boolean;
}
export interface LiteLLMParams {

View file

@ -944,4 +944,108 @@ describe("columns", () => {
expect(screen.getByText("Out: $0.03")).toBeInTheDocument();
expect(screen.queryByText(/In:/)).not.toBeInTheDocument();
});
describe("pause/resume toggle", () => {
const renderWithToggle = (
overrides: Partial<ReturnType<typeof createMockModel>["model_info"]> = {},
togglePauseHandler?: ReturnType<typeof vi.fn>,
userRole: string = "Admin",
) => {
const handler = togglePauseHandler ?? vi.fn();
const cols = columns(
userRole,
defaultProps.userID,
defaultProps.premiumUser,
defaultProps.setSelectedModelId,
defaultProps.setSelectedTeamId,
defaultProps.getDisplayModelName,
defaultProps.handleEditClick,
defaultProps.handleRefreshClick,
defaultProps.expandedRows,
defaultProps.setExpandedRows,
vi.fn(),
handler,
);
const model = createMockModel({
model_info: { ...createMockModel().model_info, ...overrides },
});
render(<TestTable data={[model]} columns={cols} />);
return { handler };
};
it("renders the toggle ON for a db_model that is not blocked", () => {
renderWithToggle({ db_model: true, blocked: false });
const toggle = screen.getByRole("switch", { name: /pause model/i });
expect(toggle).toBeEnabled();
expect(toggle).toHaveAttribute("aria-checked", "true");
});
it("renders the toggle OFF for a db_model that is blocked", () => {
renderWithToggle({ db_model: true, blocked: true });
const toggle = screen.getByRole("switch", { name: /resume model/i });
expect(toggle).toBeEnabled();
expect(toggle).toHaveAttribute("aria-checked", "false");
});
it("calls the handler with blocked=true when an admin flips an active toggle off", async () => {
const handler = vi.fn();
renderWithToggle({ db_model: true, blocked: false }, handler);
await userEvent.click(screen.getByRole("switch", { name: /pause model/i }));
expect(handler).toHaveBeenCalledWith("test-model-id", true);
});
it("calls the handler with blocked=false when an admin flips a paused toggle on", async () => {
const handler = vi.fn();
renderWithToggle({ db_model: true, blocked: true }, handler);
await userEvent.click(screen.getByRole("switch", { name: /resume model/i }));
expect(handler).toHaveBeenCalledWith("test-model-id", false);
});
it("disables the toggle for non-admin users", () => {
const handler = vi.fn();
renderWithToggle({ db_model: true, blocked: false }, handler, "User");
const toggle = screen.getByRole("switch", { name: /pause model/i });
expect(toggle).toBeDisabled();
});
it("disables the toggle for config models", () => {
const handler = vi.fn();
renderWithToggle({ db_model: false, blocked: false }, handler, "Admin");
const toggle = screen.getByRole("switch", { name: /pause model/i });
expect(toggle).toBeDisabled();
});
it("disables the toggle while a PATCH for the same row is in-flight", () => {
// Regression for Greptile P1 on PR #28151 — antd's `loading` prop is
// visual only and does not prevent click events, so the row needs to
// be explicitly disabled while its PATCH is pending to avoid
// racing/conflicting PATCH calls on double-click.
const handler = vi.fn();
const model = createMockModel({
model_info: {
...createMockModel().model_info,
db_model: true,
blocked: false,
},
});
const cols = columns(
"Admin",
defaultProps.userID,
defaultProps.premiumUser,
defaultProps.setSelectedModelId,
defaultProps.setSelectedTeamId,
defaultProps.getDisplayModelName,
defaultProps.handleEditClick,
defaultProps.handleRefreshClick,
defaultProps.expandedRows,
defaultProps.setExpandedRows,
vi.fn(),
handler,
model.model_info.id, // pausingModelId matches this row
);
render(<TestTable data={[model]} columns={cols} />);
const toggle = screen.getByRole("switch", { name: /pause model/i });
expect(toggle).toBeDisabled();
});
});
});

View file

@ -2,7 +2,7 @@ import { EditOutlined, InfoCircleOutlined, SyncOutlined } from "@ant-design/icon
import { TrashIcon } from "@heroicons/react/outline";
import { ColumnDef } from "@tanstack/react-table";
import { Badge, Button, Icon } from "@tremor/react";
import { Divider, Flex, Popover, Space, Tooltip, Typography } from "antd";
import { Divider, Flex, Popover, Space, Switch, Tooltip, Typography } from "antd";
import { ModelData } from "../../model_dashboard/types";
import { ProviderLogo } from "./ProviderLogo";
@ -53,6 +53,8 @@ export const columns = (
expandedRows: Set<string>,
setExpandedRows: (expandedRows: Set<string>) => void,
onDeleteClick?: (modelId: string) => void,
onTogglePauseClick?: (modelId: string, blocked: boolean) => void | Promise<void>,
pausingModelId?: string | null,
): ColumnDef<ModelData>[] => [
{
header: () => <span className="text-sm font-semibold">Model ID</span>,
@ -398,15 +400,48 @@ export const columns = (
{
id: "actions",
header: () => <span className="text-sm font-semibold">Actions</span>,
size: 60,
minSize: 40,
size: 100,
minSize: 80,
enableResizing: false,
cell: ({ row }) => {
const model = row.original;
const canEditModel = userRole === "Admin" || model.model_info?.created_by === userID;
const isConfigModel = !model.model_info?.db_model;
const isAdmin = userRole === "Admin";
const isBlocked = model.model_info?.blocked === true;
const isPauseToggleable = !isConfigModel && isAdmin && Boolean(onTogglePauseClick);
const pauseTooltip = isConfigModel
? "Config models cannot be paused from the dashboard. Pause is DB-backed."
: !isAdmin
? "Only proxy admins can pause or resume a model."
: isBlocked
? "Resume model — restore normal routing."
: "Pause model — stop routing requests until resumed.";
// antd's `loading` prop on Switch is purely cosmetic — it does not block
// clicks. Pair `loading` with `disabled` derived from the same condition
// so a double-click during a pending PATCH cannot send a second,
// conflicting `blocked` value.
const isPausing = pausingModelId === model.model_info?.id;
return (
<div className="flex items-center justify-end gap-2 pr-4">
<Tooltip title={pauseTooltip}>
<Switch
size="small"
checked={!isBlocked}
disabled={!isPauseToggleable || isPausing}
loading={isPausing}
aria-label={isBlocked ? "Resume model" : "Pause model"}
onClick={(_, e) => {
e.stopPropagation();
}}
onChange={(nextChecked) => {
const modelId = model.model_info?.id;
if (isPauseToggleable && onTogglePauseClick && modelId) {
void onTogglePauseClick(modelId, !nextChecked);
}
}}
/>
</Tooltip>
{isConfigModel ? (
<Tooltip title="Config model cannot be deleted on the dashboard. Please delete it from the config file.">
<Icon icon={TrashIcon} size="sm" className="opacity-50 cursor-not-allowed" />