Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_bedrock_openai_xhigh_flags

This commit is contained in:
mateo 2026-09-11 19:02:54 +00:00
commit 7b25c6a29e
71 changed files with 9673 additions and 933 deletions

View file

@ -106,11 +106,6 @@ dockerfiles:
and lint workflows already exercise that output, so building the image adds no signal about it
paths:
- ui/Dockerfile
- reason: >-
The Rust gateway ships as its own chart and package with a separate release pipeline, so its
image is not part of this repo's Python image set
paths:
- litellm-rust/crates/ai-gateway/Dockerfile
- reason: >-
An example image under cookbook/ that is documentation rather than a shipped artifact
paths:

73
.github/workflows/ai-gateway-image.yml vendored Normal file
View file

@ -0,0 +1,73 @@
name: ai-gateway image
on:
push:
paths:
- "litellm-rust/**"
- "litellm/**"
- "enterprise/**"
- "litellm-proxy-extras/**"
- "pyproject.toml"
- "rust-toolchain.toml"
- ".github/workflows/ai-gateway-image.yml"
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "litellm-rust/**"
- "litellm/**"
- "enterprise/**"
- "litellm-proxy-extras/**"
- "pyproject.toml"
- "rust-toolchain.toml"
- ".github/workflows/ai-gateway-image.yml"
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
ai-gateway-image:
name: ai-gateway release image
runs-on: ubuntu-latest
timeout-minutes: 60
permissions:
contents: read
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Build the release image
run: docker build -f litellm-rust/crates/ai-gateway/Dockerfile -t litellm-ai-gateway:${{ github.sha }} .
- name: Start the gateway and wait for readiness
env:
IMAGE: litellm-ai-gateway:${{ github.sha }}
run: |
docker run -d --name ai-gateway -p 4001:4001 \
-e LITELLM_MASTER_KEY=sk-ci-not-a-real-key \
-e OPENAI_API_KEY=sk-ci-not-a-real-key \
"$IMAGE"
for _ in $(seq 1 60); do
if curl -fsS http://127.0.0.1:4001/health/readiness; then
echo "gateway is serving readiness"
exit 0
fi
sleep 2
done
echo "gateway never became ready" >&2
docker logs ai-gateway >&2
exit 1
- name: Assert the gateway loaded the baked config
run: |
docker logs ai-gateway 2>&1 | tee gateway.log
grep 'via python config reader' gateway.log
- name: Stop the gateway
if: always()
run: docker rm -f ai-gateway || true

View file

@ -262,6 +262,8 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
}
```
For MCP OAuth, an upstream may advertise dynamic client registration but refuse requests with HTTP 401 or 403. If the provider requires a pre-registered OAuth app, configure its `credentials.client_id` and, when required, `credentials.client_secret` on the MCP server. This skips dynamic registration in the gateway sign-in flow. The provider must approve the app for MCP access; reaching its authorization page does not establish that login or tool calls will succeed
[**Docs: MCP Gateway**](https://docs.litellm.ai/docs/mcp)
</details>

View file

@ -14,15 +14,20 @@
# ---- Chef -------------------------------------------------------------------
# cargo-chef caches the dependency build so only the gateway crate recompiles on
# a source-only change. python3-dev is present in every rust stage because the
# `python-config` feature links libpython via pyo3 (even in the cook step).
FROM rust:1.90-slim-bookworm AS chef
# `python-config` feature links libpython via pyo3 (even in the cook step), and
# python3-pip builds the litellm wheel in the builder stage.
FROM rust:1.98-slim-bookworm AS chef
ENV PYO3_PYTHON=python3.11
# rustup reads rust-toolchain.toml from any parent of the working directory, so
# copying it in is what keeps every cargo call below on the repo's pinned
# channel rather than on whatever the base image happens to ship.
COPY rust-toolchain.toml /build/rust-toolchain.toml
WORKDIR /build/litellm-rust
RUN apt-get update \
&& apt-get install -y --no-install-recommends \
python3 python3-dev pkg-config libssl-dev clang \
python3 python3-dev python3-pip pkg-config libssl-dev clang \
&& rm -rf /var/lib/apt/lists/* \
&& cargo install cargo-chef --locked --version 0.1.77
WORKDIR /build/litellm-rust
# ---- Planner ----------------------------------------------------------------
# Produce the dependency recipe from the rust workspace manifests + Cargo.lock.
@ -43,6 +48,19 @@ RUN cargo chef cook --locked --release \
COPY litellm-rust/ .
RUN cargo build --locked --release -p litellm-ai-gateway --bin litellm-ai-gateway --features server,python-config
# The root pyproject builds with maturin against litellm-rust/crates/python-bridge,
# so the wheel is built here, next to the crate sources and the cargo toolchain,
# and the runtime stage installs the artifact instead of compiling anything.
# litellm[proxy] pins litellm-enterprise and litellm-proxy-extras to the versions
# in this repo, and those hit PyPI hours after every version bump merges, so both
# wheels are built from the repo too instead of being resolved from PyPI.
COPY pyproject.toml README.md LICENSE /build/
COPY litellm/ /build/litellm/
COPY enterprise/ /build/enterprise/
COPY litellm-proxy-extras/ /build/litellm-proxy-extras/
RUN pip3 wheel --no-cache-dir --no-deps --wheel-dir /build/dist \
/build /build/enterprise /build/litellm-proxy-extras
# ---- Runtime ----------------------------------------------------------------
# python:3.11-slim-bookworm ships libpython3.11, matching the builder's PyO3
# 3.11 ABI so the embedded interpreter links and imports cleanly.
@ -56,11 +74,16 @@ RUN apt-get update \
WORKDIR /app
# Install litellm (with proxy extras) FROM THIS REPO'S SOURCE so
# `import litellm.proxy.read_model_list` works — it is not on PyPI yet. Copy the
# package + packaging metadata, then pip install the proxy extra.
COPY pyproject.toml README.md LICENSE ./
COPY litellm/ ./litellm/
RUN pip install --no-cache-dir ".[proxy]"
# `import litellm.proxy.read_model_list` works — it is not on PyPI yet. The two
# sibling wheels come from the builder as well, so the pins in litellm[proxy]
# resolve against them and never wait on a PyPI publish.
COPY --from=builder /build/dist/*.whl /tmp/wheels/
RUN wheel="$(ls /tmp/wheels/litellm-*.whl)" \
&& pip install --no-cache-dir \
/tmp/wheels/litellm_enterprise-*.whl \
/tmp/wheels/litellm_proxy_extras-*.whl \
"${wheel}[proxy]" \
&& rm -rf /tmp/wheels
# The compiled gateway binary (pure-Rust realtime hot path; Python is load-time
# only).

View file

@ -9,19 +9,28 @@
# Strategy: ignore everything, then re-include only what the build needs:
# - litellm/ (pip install . needs the full package + proxy reader)
# - litellm-rust/ (the rust workspace; Cargo.lock + crate sources)
# - pyproject.toml / README.md / LICENSE (packaging metadata for pip install)
# - enterprise/ (litellm/proxy/enterprise symlinks into it; maturin walks it)
# - litellm-proxy-extras/ (built into a wheel alongside enterprise/ for litellm[proxy])
# - pyproject.toml / README.md / LICENSE (packaging metadata for the wheel build)
# - rust-toolchain.toml (the pinned channel every cargo call in the build uses)
*
# --- re-include the build inputs ---
!litellm/
!litellm-rust/
!enterprise/
!litellm-proxy-extras/
!pyproject.toml
!rust-toolchain.toml
!README.md
!LICENSE
# --- prune heavy / irrelevant subpaths back out of the re-included trees ---
# Rust build artifacts (huge; regenerated in the builder).
**/target/
# Committed python distribution artifacts; the wheel build does not read them.
enterprise/dist/
litellm-proxy-extras/dist/
# Python caches and compiled bytecode.
**/__pycache__/
**/*.pyc

View file

@ -1370,6 +1370,7 @@ from .exceptions import (
InvalidRequestError,
BadRequestError,
ImageFetchError,
VectorStoreSearchError,
NotFoundError,
PermissionDeniedError,
RateLimitError,
@ -1471,9 +1472,11 @@ from .vector_stores.vector_store_registry import (
VectorStoreRegistry,
VectorStoreIndexRegistry,
)
from .types.vector_stores import VectorStoreSearchFailureMode
vector_store_registry: Optional[VectorStoreRegistry] = None
vector_store_index_registry: Optional[VectorStoreIndexRegistry] = None
vector_store_search_failure_mode: VectorStoreSearchFailureMode = "annotate"
### RAG ###
from . import rag

View file

@ -22,7 +22,7 @@ from litellm._logging import print_verbose, verbose_logger
from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE
from .base_cache import BaseCache
from .in_memory_cache import InMemoryCache
from .in_memory_cache import DEFAULT_MAX_SIZE_IN_MEMORY, InMemoryCache
from .redis_cache import RedisCache, RedisCircuitBreakerOpenError, log_redis_failure
if TYPE_CHECKING:
@ -83,6 +83,9 @@ class DualCache(BaseCache):
if default_redis_ttl is not None:
self.default_redis_ttl = default_redis_ttl
def update_in_memory_max_size(self, max_size: int | None) -> None:
self.in_memory_cache.max_size_in_memory = DEFAULT_MAX_SIZE_IN_MEMORY if max_size is None else max_size
def attach_redis_cache(
self,
redis_cache: RedisCache | None = None,
@ -376,7 +379,9 @@ class DualCache(BaseCache):
)
# async_batch_set_cache
async def async_set_cache_pipeline(self, cache_list: list, local_only: bool = False, **kwargs):
async def async_set_cache_pipeline(
self, cache_list: Sequence[tuple[str, object]], local_only: bool = False, **kwargs
):
"""
Batch write values to the cache
"""

View file

@ -24,11 +24,13 @@ from litellm.constants import MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB
from .base_cache import BaseCache
DEFAULT_MAX_SIZE_IN_MEMORY: Final = 200
class InMemoryCache(BaseCache):
def __init__(
self,
max_size_in_memory: int | None = 200,
max_size_in_memory: int | None = DEFAULT_MAX_SIZE_IN_MEMORY,
default_ttl: int
| None = 600, # default ttl is 10 minutes. At maximum litellm rate limiting logic requires objects to be in memory for 1 minute
max_size_per_item: int | None = 1024, # 1MB = 1024KB
@ -37,7 +39,7 @@ class InMemoryCache(BaseCache):
max_size_in_memory [int]: Maximum number of items in cache. done to prevent memory leaks. Use 200 items as a default
"""
self.max_size_in_memory = (
max_size_in_memory if max_size_in_memory is not None else 200
max_size_in_memory if max_size_in_memory is not None else DEFAULT_MAX_SIZE_IN_MEMORY
) # set an upper bound of 200 items in-memory
self.default_ttl = default_ttl or 600
self.max_size_per_item = max_size_per_item or MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB # 1MB = 1024KB

View file

@ -10,12 +10,14 @@
## LiteLLM versions of the OpenAI Exception Types
import enum
from collections.abc import Sequence
from typing import Any, Final
import httpx
import openai
from litellm.types.utils import LiteLLMCommonStrings
from litellm.types.vector_stores import VectorStoreSearchFailure
class RateLimitErrorCategory(str, enum.Enum):
@ -288,6 +290,29 @@ class ImageFetchError(BadRequestError):
)
VECTOR_STORE_SEARCH_FAILED_CODE: Final = "vector_store_search_failed"
class VectorStoreSearchError(BadRequestError):
def __init__(
self,
failures: Sequence[VectorStoreSearchFailure],
model: str | None = None,
llm_provider: str | None = None,
) -> None:
self.failures: Final[tuple[VectorStoreSearchFailure, ...]] = tuple(failures)
detail: Final = "; ".join(f"{failure['vector_store_id']}: {failure['error']}" for failure in self.failures)
super().__init__(
message=(
"The request could not be grounded in every configured vector store. "
f"{len(self.failures)} vector store search(es) failed: {detail}"
),
model=model,
llm_provider=llm_provider,
body={"type": "invalid_request_error", "code": VECTOR_STORE_SEARCH_FAILED_CODE},
)
class UnprocessableEntityError(openai.UnprocessableEntityError):
def __init__(
self,

View file

@ -10,6 +10,7 @@ from contextlib import AbstractAsyncContextManager
from datetime import timedelta
from functools import partial
from importlib import metadata
from types import MappingProxyType
from typing import Any, Final, Protocol, TypeAlias, TypeVar
import httpx
@ -77,6 +78,7 @@ from litellm._logging import verbose_logger
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR, MCP_TOOL_LISTING_TIMEOUT
from litellm.experimental_mcp_client.tools import list_tools_with_pagination
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response
from litellm.types.llms.custom_http import VerifyTypes
from litellm.types.mcp import (
MCPAuth,
@ -631,7 +633,9 @@ class MCPClient:
auth=effective_auth,
verify=ssl_config,
follow_redirects=True,
event_hooks={"request": [guard]} if guard else {},
event_hooks=MappingProxyType(
{"response": [capture_upstream_error_response], "request": [guard] if guard else []}
), # mutable-ok: httpx types require lists of hooks
)
return factory

View file

@ -5,20 +5,29 @@ This hook is called before making an LLM request when a vector store is configur
It searches the vector store for relevant context and appends it to the messages.
"""
from collections.abc import Awaitable, Callable
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final, Protocol, cast
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, get_args
from pydantic import TypeAdapter, ValidationError
from typing_extensions import assert_never
import litellm
import litellm.vector_stores
from litellm._logging import verbose_logger
from litellm.exceptions import VectorStoreSearchError
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionUserMessage,
ResponsesAPIResponse,
)
from litellm.types.prompts.init_prompts import PromptSpec
from litellm.types.utils import CallTypes, StandardCallbackDynamicParams
from litellm.types.vector_stores import (
LiteLLM_ManagedVectorStore,
VectorStoreResultContent,
VectorStoreSearchFailure,
VectorStoreSearchFailureMode,
VectorStoreSearchResponse,
VectorStoreSearchResult,
)
@ -30,6 +39,10 @@ if TYPE_CHECKING:
else:
LiteLLMLoggingObj = Any
SEARCH_FAILURES_FIELD: Final = "vector_store_search_failures"
_DEFAULT_FAILURE_MODE: Final[VectorStoreSearchFailureMode] = "annotate"
_FAILURE_MODE_ADAPTER: Final = TypeAdapter(VectorStoreSearchFailureMode)
class ProxyRuntime(Protocol):
def llm_router(self) -> "Router | None": ...
@ -54,11 +67,31 @@ class ProxyServerRuntime:
return prisma_client
@dataclass(frozen=True, slots=True)
class SearchSucceeded:
response: VectorStoreSearchResponse
@dataclass(frozen=True, slots=True)
class SearchFailed:
failure: VectorStoreSearchFailure
SearchOutcome = SearchSucceeded | SearchFailed
@dataclass(frozen=True, slots=True)
class VectorStoreAugmentation:
messages: tuple[AllMessageValues, ...]
search_results: tuple[VectorStoreSearchResponse, ...]
failures: tuple[VectorStoreSearchFailure, ...]
class VectorStorePreCallHook(CustomLogger):
CONTENT_PREFIX_STRING = "Context:\n\n"
"""
Custom logger that handles vector store searches before LLM calls.
When a vector store is configured, this hook:
1. Extracts the query from the last user message
2. Calls litellm.vector_stores.search() to get relevant context
@ -101,100 +134,153 @@ class VectorStorePreCallHook(CustomLogger):
Returns:
Tuple of (model, modified_messages, non_default_params)
"""
requested_vector_store_ids: Final = _requested_vector_store_ids(non_default_params)
try:
# Check if vector store is configured
if litellm.vector_store_registry is None:
return model, messages, non_default_params
prisma_client: Final = self.proxy_runtime.prisma_client()
llm_router: Final = self.proxy_runtime.llm_router()
# Use database fallback to ensure synchronization across instances
vector_stores_to_run: list[
LiteLLM_ManagedVectorStore
] = await litellm.vector_store_registry.pop_vector_stores_to_run_with_db_fallback(
augmentation: VectorStoreAugmentation | None = await self._augment_messages(
messages=messages,
non_default_params=non_default_params,
tools=tools,
prisma_client=prisma_client,
litellm_logging_obj=litellm_logging_obj,
)
if not vector_stores_to_run:
return model, messages, non_default_params
# Extract the query from the last user message
query: Final = self._extract_query_from_messages(messages)
if not query:
verbose_logger.debug("No query found in messages for vector store search")
return model, messages, non_default_params
modified_messages: list[AllMessageValues] = messages.copy()
all_search_results: Final[list[VectorStoreSearchResponse]] = []
for vector_store_to_run in vector_stores_to_run:
# Get vector store id from the vector store config
vector_store_id = vector_store_to_run.get("vector_store_id", "")
custom_llm_provider = vector_store_to_run.get("custom_llm_provider")
litellm_params_for_vector_store = vector_store_to_run.get("litellm_params", {}) or {}
request_litellm_params = litellm_logging_obj.model_call_details.get("litellm_params", {})
request_metadata = (
request_litellm_params.get("metadata", {}) if isinstance(request_litellm_params, dict) else {}
)
if llm_router is not None:
search_function = cast( # cast-ok: normalize router search callable
Callable[..., Awaitable[VectorStoreSearchResponse]],
llm_router.avector_store_search,
)
else:
search_function = cast( # cast-ok: normalize SDK search callable
Callable[..., Awaitable[VectorStoreSearchResponse]],
litellm.vector_stores.asearch,
)
try:
search_response = await search_function(
**{
"vector_store_id": vector_store_id,
"query": query,
"custom_llm_provider": custom_llm_provider,
"metadata": request_metadata,
**litellm_params_for_vector_store,
},
)
except Exception as search_error:
verbose_logger.warning(
"Vector store search failed for vector_store_id=%s, continuing without its context: %s",
vector_store_id,
search_error,
)
continue
verbose_logger.debug("search_response: %s", search_response)
# Store search results for later use in citations
all_search_results.append(search_response)
# Process search results and append as context
modified_messages = self._append_search_results_to_messages(
messages=modified_messages, search_response=search_response
)
# Get the number of results for logging
num_results = 0
num_results = len(search_response.get("data", []) or [])
verbose_logger.debug("Vector store search completed. Added context from %s results", num_results)
# Store search results as-is (already in OpenAI-compatible format)
if litellm_logging_obj and all_search_results:
litellm_logging_obj.model_call_details["search_results"] = all_search_results
return model, modified_messages, non_default_params
except Exception as e:
verbose_logger.exception("Error in VectorStorePreCallHook: %s", e)
# Return original parameters on error
verbose_logger.exception(
"Error in VectorStorePreCallHook for vector_store_ids=%s: %s",
requested_vector_store_ids,
e,
)
return model, messages, non_default_params
def _extract_query_from_messages(self, messages: list[AllMessageValues]) -> str | None:
if augmentation is None:
return model, messages, non_default_params
for detail, value in (
("search_results", list(augmentation.search_results)),
(SEARCH_FAILURES_FIELD, augmentation.failures),
):
if value:
litellm_logging_obj.model_call_details[detail] = value
if augmentation.failures:
failure_mode: Final = _configured_failure_mode()
match failure_mode:
case "error":
raise VectorStoreSearchError(failures=augmentation.failures, model=model)
case "annotate":
pass
case _:
assert_never(failure_mode)
return model, list(augmentation.messages), non_default_params
async def _augment_messages(
self,
messages: Sequence[AllMessageValues],
non_default_params: dict,
tools: list[dict] | None,
litellm_logging_obj: LiteLLMLoggingObj,
) -> VectorStoreAugmentation | None:
if litellm.vector_store_registry is None:
return None
prisma_client: Final = self.proxy_runtime.prisma_client()
llm_router: Final = self.proxy_runtime.llm_router()
# Use database fallback to ensure synchronization across instances
vector_stores_to_run: Final[
Sequence[LiteLLM_ManagedVectorStore]
] = await litellm.vector_store_registry.pop_vector_stores_to_run_with_db_fallback(
non_default_params=non_default_params,
tools=tools,
prisma_client=prisma_client,
)
if not vector_stores_to_run:
return None
query: Final = self._extract_query_from_messages(messages)
if not query:
verbose_logger.debug("No query found in messages for vector store search")
return None
request_litellm_params: Final = litellm_logging_obj.model_call_details.get("litellm_params", {})
request_metadata: Final = (
request_litellm_params.get("metadata", {}) if isinstance(request_litellm_params, dict) else {}
)
search_function: Final = (
cast( # cast-ok: normalize router search callable
Callable[..., Awaitable[VectorStoreSearchResponse]],
llm_router.avector_store_search,
)
if llm_router is not None
else cast( # cast-ok: normalize SDK search callable
Callable[..., Awaitable[VectorStoreSearchResponse]],
litellm.vector_stores.asearch,
)
)
outcomes: Final = tuple(
[
await self._search_one(
vector_store=vector_store_to_run,
query=query,
request_metadata=request_metadata,
search_function=search_function,
)
for vector_store_to_run in vector_stores_to_run
]
)
search_results: Final = tuple(outcome.response for outcome in outcomes if isinstance(outcome, SearchSucceeded))
failures: Final = tuple(outcome.failure for outcome in outcomes if isinstance(outcome, SearchFailed))
return VectorStoreAugmentation(
messages=self._messages_with_context(messages=messages, search_results=search_results),
search_results=search_results,
failures=failures,
)
async def _search_one(
self,
vector_store: LiteLLM_ManagedVectorStore,
query: str,
request_metadata: Mapping[str, object],
search_function: Callable[..., Awaitable[VectorStoreSearchResponse]],
) -> SearchOutcome:
vector_store_id: Final = vector_store.get("vector_store_id", "")
custom_llm_provider: Final = vector_store.get("custom_llm_provider")
litellm_params_for_vector_store: Final = vector_store.get("litellm_params", {}) or {}
try:
search_response: Final = await search_function(
**{
"vector_store_id": vector_store_id,
"query": query,
"custom_llm_provider": custom_llm_provider,
"metadata": request_metadata,
**litellm_params_for_vector_store,
},
)
except Exception as search_error:
verbose_logger.warning(
"Vector store search failed for vector_store_id=%s, continuing without its context: %s",
vector_store_id,
search_error,
)
return SearchFailed(
failure=VectorStoreSearchFailure(
vector_store_id=vector_store_id,
custom_llm_provider=custom_llm_provider,
error=str(search_error),
)
)
verbose_logger.debug(
"Vector store search completed for vector_store_id=%s. Added context from %s results",
vector_store_id,
len(search_response.get("data") or ()),
)
return SearchSucceeded(response=search_response)
def _extract_query_from_messages(self, messages: Sequence[AllMessageValues]) -> str | None:
"""
Extract the query from the last user message.
@ -223,48 +309,40 @@ class VectorStorePreCallHook(CustomLogger):
return None
def _append_search_results_to_messages(
def _messages_with_context(
self,
messages: list[AllMessageValues],
search_response: VectorStoreSearchResponse,
) -> list[AllMessageValues]:
"""
Append search results as context to the messages.
messages: Sequence[AllMessageValues],
search_results: Sequence[VectorStoreSearchResponse],
) -> tuple[AllMessageValues, ...]:
context_messages: Final = tuple(
context_message
for search_response in search_results
if (context_message := self._context_message(search_response)) is not None
)
if not context_messages:
return tuple(messages)
return (*messages[:-1], *context_messages, *messages[-1:])
Args:
messages: Original list of messages
search_response: Response from vector store search
Returns:
Modified list of messages with context appended
"""
search_response_data: Final[list[VectorStoreSearchResult] | None] = search_response.get("data")
def _context_message(self, search_response: VectorStoreSearchResponse) -> AllMessageValues | None:
"""Build the context message for one vector store's results, or None when it returned nothing usable."""
search_response_data: Final[Sequence[VectorStoreSearchResult] | None] = search_response.get("data")
if not search_response_data:
return messages
return None
context_content = self.CONTENT_PREFIX_STRING
context_texts: Final = tuple(
content_text
for result in search_response_data
for content_item in (result.get("content") or ())
if (content_text := content_item.get("text"))
)
if not context_texts:
return None
for result in search_response_data:
result_content: list[VectorStoreResultContent] | None = result.get("content")
if result_content:
for content_item in result_content:
content_text: str | None = content_item.get("text")
if content_text:
context_content += content_text + "\n\n"
# Only add context if we found any content
if context_content != "Context:\n\n":
# Create a copy of messages to avoid modifying the original
modified_messages: Final = messages.copy()
# Add context as a new message before the last user message
context_message: Final[ChatCompletionUserMessage] = {
"role": "user",
"content": context_content,
}
modified_messages.insert(-1, cast(AllMessageValues, context_message))
return modified_messages
return messages
context_message: Final[ChatCompletionUserMessage] = {
"role": "user",
"content": self.CONTENT_PREFIX_STRING + "".join(f"{text}\n\n" for text in context_texts),
}
return cast(AllMessageValues, context_message)
async def async_post_call_success_deployment_hook(
self,
@ -287,34 +365,34 @@ class VectorStorePreCallHook(CustomLogger):
verbose_logger.debug("No litellm_logging_obj in request_data")
return None
verbose_logger.debug("model_call_details keys: %s", list(litellm_logging_obj.model_call_details.keys()))
# Get search results from model_call_details (already in OpenAI format)
search_results: Final[list[VectorStoreSearchResponse] | None] = litellm_logging_obj.model_call_details.get(
"search_results"
search_results: Final[Sequence[VectorStoreSearchResponse] | None] = (
litellm_logging_obj.model_call_details.get("search_results")
)
search_failures: Final[Sequence[VectorStoreSearchFailure] | None] = (
litellm_logging_obj.model_call_details.get(SEARCH_FAILURES_FIELD)
)
verbose_logger.debug("Search results found: %s", search_results is not None)
if not search_results:
verbose_logger.debug("No search results found")
if not search_results and not search_failures:
verbose_logger.debug("No search results or search failures found")
return None
if isinstance(response, ResponsesAPIResponse):
if search_failures:
setattr(response, SEARCH_FAILURES_FIELD, list(search_failures))
return response
# Add search results to response object
if hasattr(response, "choices") and response.choices:
for choice in response.choices:
if hasattr(choice, "message") and choice.message:
# Get existing provider_specific_fields or create new dict
provider_fields = getattr(choice.message, "provider_specific_fields", None) or {}
# Add search results (already in OpenAI-compatible format)
provider_fields["search_results"] = search_results
# Set the provider_specific_fields
if search_results:
provider_fields["search_results"] = search_results
if search_failures:
provider_fields[SEARCH_FAILURES_FIELD] = search_failures
setattr(choice.message, "provider_specific_fields", provider_fields)
verbose_logger.debug("Added %s search results to response", len(search_results))
# Return modified response
return response
@ -339,29 +417,24 @@ class VectorStorePreCallHook(CustomLogger):
verbose_logger.debug("VectorStorePreCallHook.async_post_call_streaming_deployment_hook called")
# Get search results from model_call_details (already in OpenAI format)
search_results: Final[list[VectorStoreSearchResponse] | None] = request_data.get("search_results")
search_results: Final[Sequence[VectorStoreSearchResponse] | None] = request_data.get("search_results")
search_failures: Final[Sequence[VectorStoreSearchFailure] | None] = request_data.get(SEARCH_FAILURES_FIELD)
verbose_logger.debug("Search results found for streaming chunk: %s", search_results is not None)
if not search_results:
verbose_logger.debug("No search results found for streaming chunk")
if not search_results and not search_failures:
verbose_logger.debug("No search results or search failures found for streaming chunk")
return response_chunk
# Add search results to streaming chunk
if hasattr(response_chunk, "choices") and response_chunk.choices:
for choice in response_chunk.choices:
if hasattr(choice, "delta") and choice.delta:
# Get existing provider_specific_fields or create new dict
provider_fields = getattr(choice.delta, "provider_specific_fields", None) or {}
# Add search results (already in OpenAI-compatible format)
provider_fields["search_results"] = search_results
# Set the provider_specific_fields
if search_results:
provider_fields["search_results"] = search_results
if search_failures:
provider_fields[SEARCH_FAILURES_FIELD] = search_failures
choice.delta.provider_specific_fields = provider_fields
verbose_logger.debug("Added %s search results to streaming chunk", len(search_results))
# Return modified chunk
return response_chunk
@ -369,3 +442,23 @@ class VectorStorePreCallHook(CustomLogger):
verbose_logger.exception("Error adding search results to streaming chunk: %s", e)
# Don't fail the request if search results fail to be added
return response_chunk
def _requested_vector_store_ids(non_default_params: Mapping[str, object]) -> tuple[str, ...]:
requested: Final = non_default_params.get("vector_store_ids")
if not isinstance(requested, (list, tuple)):
return ()
return tuple(str(vector_store_id) for vector_store_id in requested)
def _configured_failure_mode() -> VectorStoreSearchFailureMode:
try:
return _FAILURE_MODE_ADAPTER.validate_python(litellm.vector_store_search_failure_mode)
except ValidationError:
verbose_logger.warning(
"Unsupported vector_store_search_failure_mode=%r, falling back to %r. Supported modes: %s",
litellm.vector_store_search_failure_mode,
_DEFAULT_FAILURE_MODE,
", ".join(get_args(VectorStoreSearchFailureMode)),
)
return _DEFAULT_FAILURE_MODE

View file

@ -14,7 +14,7 @@ from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from datetime import datetime as dt_object
from functools import lru_cache
from types import MappingProxyType, TracebackType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Union, cast
from httpx import Response
from pydantic import BaseModel, JsonValue
@ -211,7 +211,7 @@ if TYPE_CHECKING:
from litellm.integrations.otel.logger import OpenTelemetryV2
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
from litellm.litellm_core_utils.llm_cost_calc.utils import BilledTokenRates
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, LoggedRelayResponse
from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector
try:
from litellm_enterprise.enterprise_callbacks.callback_controls import (
EnterpriseCallbackControls,
@ -2396,52 +2396,17 @@ class Logging(LiteLLMLoggingBaseClass):
for scope in [key for key in spans_logged if isinstance(key, tuple) and key[-1:] == ("success",)]:
del spans_logged[scope]
def _flush_passthrough_collected_chunks_helper(
self,
raw_bytes: list[bytes],
provider_config: "BasePassthroughConfig",
) -> Optional["LoggedRelayResponse"]:
all_chunks: Final = provider_config._convert_raw_bytes_to_str_lines(raw_bytes)
complete_streaming_response: Final = provider_config.handle_logging_collected_chunks(
all_chunks=all_chunks,
litellm_logging_obj=self,
model=self.model,
custom_llm_provider=self.model_call_details.get("custom_llm_provider", ""),
endpoint=self.model_call_details.get("endpoint", ""),
)
return complete_streaming_response
def flush_passthrough_collected_chunks(
self,
raw_bytes: list[bytes],
provider_config: "BasePassthroughConfig",
):
def flush_passthrough_collected_chunks(self, collector: "PassthroughStreamCollector"):
"""
Flush collected chunks from the logging object
This is used to log the collected chunks once streaming is done on passthrough endpoints
1. Decode the raw bytes to string lines
2. Get the complete streaming response from the provider config
3. Log the complete streaming response (trigger success handler)
This is used for passthrough endpoints
Log the response a passthrough stream collector assembled once streaming is done (trigger success handler)
"""
complete_streaming_response: Final = self._flush_passthrough_collected_chunks_helper(
raw_bytes=raw_bytes,
provider_config=provider_config,
)
complete_streaming_response: Final = collector.build_logged_response(litellm_logging_obj=self)
if complete_streaming_response is not None:
self.success_handler(result=complete_streaming_response)
async def async_flush_passthrough_collected_chunks(
self,
raw_bytes: list[bytes],
provider_config: "BasePassthroughConfig",
):
complete_streaming_response: Final = self._flush_passthrough_collected_chunks_helper(
raw_bytes=raw_bytes,
provider_config=provider_config,
)
async def async_flush_passthrough_collected_chunks(self, collector: "PassthroughStreamCollector"):
complete_streaming_response: Final = collector.build_logged_response(litellm_logging_obj=self)
if complete_streaming_response is not None:
await self.async_success_handler(result=complete_streaming_response)
@ -6505,7 +6470,7 @@ def _get_traceback_str_for_error(error_str: str) -> str:
from decimal import Decimal
# used for unit testing
from typing import Any, Optional, Union
from typing import Any, Union
def create_dummy_standard_logging_payload() -> StandardLoggingPayload:

View file

@ -724,10 +724,11 @@ class ModelResponseIterator:
content_block: Final = ContentBlockDelta(**chunk)
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] = []
self.content_blocks.append(content_block)
if "text" in content_block["delta"]:
text = content_block["delta"]["text"]
elif "partial_json" in content_block["delta"]:
return text, tool_use, thinking_blocks, provider_specific_fields, reasoning_content
self.content_blocks.append(content_block)
if "partial_json" in content_block["delta"]:
# Only emit tool calls if we're in a tool_use or server_tool_use block
# web_search_tool_result blocks also have input_json_delta but should not be treated as tool calls
# See: https://github.com/BerriAI/litellm/issues/17254

View file

@ -4,7 +4,7 @@ import re
from abc import abstractmethod
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Final, TypeAlias
from typing import TYPE_CHECKING, Final, Protocol, TypeAlias
from pydantic import TypeAdapter, ValidationError
@ -80,6 +80,38 @@ def logged_relay_shape(
return parsed
class PassthroughStreamCollector(Protocol):
"""Consumes relayed stream bytes as they arrive and builds the response logged for spend tracking."""
def add(self, chunk: bytes) -> None: ...
def build_logged_response(self, litellm_logging_obj: LiteLLMLoggingObj) -> LoggedRelayResponse | None: ...
class RawBytesStreamCollector:
def __init__(
self, provider_config: BasePassthroughConfig, model: str, custom_llm_provider: str, endpoint: str
) -> None:
self._provider_config = provider_config
self._model = model
self._custom_llm_provider = custom_llm_provider
self._endpoint = endpoint
self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks
def add(self, chunk: bytes) -> None:
self._raw_bytes.append(chunk)
def build_logged_response(self, litellm_logging_obj: LiteLLMLoggingObj) -> LoggedRelayResponse | None:
all_chunks: Final = self._provider_config._convert_raw_bytes_to_str_lines(self._raw_bytes)
return self._provider_config.handle_logging_collected_chunks(
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=self._model,
custom_llm_provider=self._custom_llm_provider,
endpoint=self._endpoint,
)
class BasePassthroughConfig(BaseLLMModelInfo):
@abstractmethod
def is_streaming_request(self, endpoint: str, request_data: dict) -> bool:
@ -182,6 +214,13 @@ class BasePassthroughConfig(BaseLLMModelInfo):
) -> LoggedRelayResponse | None:
return None
def create_stream_collector(
self, model: str, custom_llm_provider: str, endpoint: str
) -> PassthroughStreamCollector:
return RawBytesStreamCollector(
provider_config=self, model=model, custom_llm_provider=custom_llm_provider, endpoint=endpoint
)
def _convert_raw_bytes_to_str_lines(self, raw_bytes: list[bytes]) -> list[str]:
"""
Converts a list of raw bytes into a list of string lines, similar to aiter_lines()

View file

@ -490,10 +490,10 @@ class AWSEventStreamDecoder:
reasoning_content: str | None = None
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
self.content_blocks.append(delta_obj)
if "text" in delta_obj:
text = delta_obj["text"]
elif "toolUse" in delta_obj:
self.content_blocks.append(delta_obj)
# When json_mode is True and this is the internal json_tool_call,
# convert tool input to text content instead of tool call arguments
if self.json_mode is True and self._current_tool_name == RESPONSE_FORMAT_TOOL_NAME:

View file

@ -1,23 +1,129 @@
import json
from collections.abc import Mapping
from collections.abc import Callable, Mapping, Sequence
from typing import TYPE_CHECKING, Final, Optional, cast
import httpx
from httpx import Response
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, PassthroughStreamCollector
from litellm.types.utils import ModelResponseStream
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import BedrockError, BedrockEventStreamDecoderBase, BedrockModelInfo
if TYPE_CHECKING:
from botocore.eventstream import EventStreamMessage
from httpx import URL
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder
from litellm.types.utils import CostResponseTypes
_TEXT_ONLY_DELTA_FIELDS: Final = frozenset({"content", "role"})
def _plain_text_delta(chunk: ModelResponseStream) -> str | None:
"""Return the delta text when the chunk carries nothing else that stream_chunk_builder reads."""
if chunk.get("usage") is not None or chunk.provider_specific_fields or len(chunk.choices) != 1:
return None
choice: Final = chunk.choices[0]
if choice.finish_reason or choice.logprobs is not None:
return None
populated: Final = frozenset(key for key, value in choice.delta.model_dump().items() if value is not None)
if not populated <= _TEXT_ONLY_DELTA_FIELDS:
return None
content: Final = choice.delta.get("content")
return content if isinstance(content, str) else None
class _CoalescedChunks:
"""Retains translated chunks with consecutive text deltas folded into one, so memory tracks the response text,
not the event count."""
def __init__(self) -> None:
self._chunks: list[ModelResponseStream] = [] # mutable-ok: instance accumulator for streaming chunks
self._open_text_parts: list[str] = [] # mutable-ok: text deltas pending fold into self._chunks[-1]
def add(self, chunk: ModelResponseStream) -> None:
text: Final = _plain_text_delta(chunk)
if text is not None and self._open_text_parts:
self._open_text_parts.append(text)
return
self._seal_text_run()
self._chunks.append(chunk)
if text is not None:
self._open_text_parts.append(text)
def _seal_text_run(self) -> None:
if len(self._open_text_parts) > 1:
self._chunks[-1].choices[0].delta.content = "".join(self._open_text_parts)
self._open_text_parts.clear()
def chunks(self) -> Sequence[ModelResponseStream]:
self._seal_text_run()
return self._chunks
def _translate_message(decoder: "AWSEventStreamDecoder", message: str) -> ModelResponseStream | None:
from litellm.litellm_core_utils.streaming_handler import (
convert_generic_chunk_to_model_response_stream,
generic_chunk_has_all_required_fields,
)
from litellm.types.utils import GenericStreamingChunk
translated_chunk: Final = decoder._chunk_parser(chunk_data=json.loads(message))
if isinstance(translated_chunk, ModelResponseStream):
return translated_chunk
if generic_chunk_has_all_required_fields(cast(dict, translated_chunk)):
return convert_generic_chunk_to_model_response_stream(cast(GenericStreamingChunk, translated_chunk))
return None
def _build_logged_response(
chunks: Sequence[ModelResponseStream], litellm_logging_obj: "LiteLLMLoggingObj"
) -> Optional["CostResponseTypes"]:
from litellm.main import stream_chunk_builder
if len(chunks) == 0:
return None
return stream_chunk_builder(chunks=list(chunks), logging_obj=litellm_logging_obj)
class BedrockEventStreamCollector:
"""Decodes and translates Bedrock event-stream frames as they are relayed instead of buffering the stream."""
def __init__(
self,
parse_event: Callable[["EventStreamMessage"], str | None],
decoder: Optional["AWSEventStreamDecoder"],
) -> None:
from botocore.eventstream import EventStreamBuffer
self._parse_event = parse_event
self._decoder = decoder
self._event_stream_buffer: Final[EventStreamBuffer] = EventStreamBuffer()
self._chunks: Final = _CoalescedChunks()
def add(self, chunk: bytes) -> None:
if self._decoder is None:
return
self._event_stream_buffer.add_data(chunk)
for event in self._event_stream_buffer:
self._add_event(self._decoder, event)
def _add_event(self, decoder: "AWSEventStreamDecoder", event: "EventStreamMessage") -> None:
message: Final = self._parse_event(event)
translated: Final = _translate_message(decoder, message) if message is not None else None
if translated is not None:
self._chunks.add(translated)
def build_logged_response(self, litellm_logging_obj: "LiteLLMLoggingObj") -> Optional["CostResponseTypes"]:
return _build_logged_response(self._chunks.chunks(), litellm_logging_obj)
class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamDecoderBase, BasePassthroughConfig):
def get_error_class(
self,
@ -168,87 +274,32 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD
return litellm_model_response
def _convert_raw_bytes_to_str_lines(self, raw_bytes: list[bytes]) -> list[str]:
from botocore.eventstream import EventStreamBuffer
all_chunks: Final = []
event_stream_buffer: Final = EventStreamBuffer()
for chunk in raw_bytes:
event_stream_buffer.add_data(chunk)
for event in event_stream_buffer:
message = self._parse_message_from_event(event)
if message is not None:
all_chunks.append(message)
return all_chunks
def handle_logging_collected_chunks(
self,
all_chunks: list[str],
litellm_logging_obj: "LiteLLMLoggingObj",
model: str,
custom_llm_provider: str,
endpoint: str,
) -> Optional["CostResponseTypes"]:
"""
1. Convert all_chunks to a ModelResponseStream
2. combine model_response_stream to model_response
3. Return the model_response
"""
from litellm.litellm_core_utils.streaming_handler import (
convert_generic_chunk_to_model_response_stream,
generic_chunk_has_all_required_fields,
def create_stream_collector(
self, model: str, custom_llm_provider: str, endpoint: str
) -> PassthroughStreamCollector:
return BedrockEventStreamCollector(
parse_event=self._parse_message_from_event,
decoder=self._get_event_stream_decoder(model=model, endpoint=endpoint),
)
def _get_event_stream_decoder(self, model: str, endpoint: str) -> Optional["AWSEventStreamDecoder"]:
from litellm.llms.bedrock.chat import get_bedrock_event_stream_decoder
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
AmazonInvokeConfig,
)
from litellm.main import stream_chunk_builder
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
all_translated_chunks: Final = []
if "invoke" in endpoint:
invoke_provider: Final = AmazonInvokeConfig.get_bedrock_invoke_provider(model)
if invoke_provider is None:
raise ValueError(f"Invalid invoke provider: {invoke_provider}, for model: {model}")
obj = get_bedrock_event_stream_decoder(
invoke_provider=invoke_provider,
model=model,
sync_stream=True,
json_mode=False,
)
elif "converse" in endpoint:
obj = get_bedrock_event_stream_decoder(
invoke_provider=None,
model=model,
sync_stream=True,
json_mode=False,
)
else:
return None
for chunk in all_chunks:
message = json.loads(chunk)
translated_chunk = obj._chunk_parser(chunk_data=message)
if isinstance(translated_chunk, dict) and generic_chunk_has_all_required_fields(
cast(dict, translated_chunk)
):
chunk_obj = convert_generic_chunk_to_model_response_stream(
cast(GenericStreamingChunk, translated_chunk)
verbose_logger.warning(
"Bedrock passthrough spend tracking skipped: no invoke provider for model %s", model
)
elif isinstance(translated_chunk, ModelResponseStream):
chunk_obj = translated_chunk
else:
continue
all_translated_chunks.append(chunk_obj)
if len(all_translated_chunks) > 0:
model_response: Final = stream_chunk_builder(
chunks=all_translated_chunks,
logging_obj=litellm_logging_obj,
return None
return get_bedrock_event_stream_decoder(
invoke_provider=invoke_provider, model=model, sync_stream=True, json_mode=False
)
if "converse" in endpoint:
return get_bedrock_event_stream_decoder(
invoke_provider=None, model=model, sync_stream=True, json_mode=False
)
return model_response
return None

View file

@ -745,13 +745,25 @@ class OCIStreamWrapper(CustomStreamWrapper):
# single-event case (terminal chunk carries the only copy of the text).
self._cohere_text_emitted = False
def chunk_creator(self, chunk: Any) -> ModelResponseStream:
def _emit_chunk(self, parsed: ModelResponseStream) -> ModelResponseStream:
for choice in parsed.choices:
if getattr(choice.delta, "tool_calls", None):
self.tool_call = True
if choice.finish_reason is not None:
self.received_finish_reason = choice.finish_reason
self.sent_last_chunk = True
return self.model_response_creator(chunk={"choices": parsed.choices})
def chunk_creator(self, chunk: Any) -> ModelResponseStream | None:
if not isinstance(chunk, str):
raise ValueError(f"Chunk is not a string: {chunk}")
if not chunk.startswith("data:"):
raise ValueError(f"Chunk does not start with 'data:': {chunk}")
payload: Final = chunk[5:].strip()
if payload == "[DONE]":
return None
try:
dict_chunk: Final = json.loads(chunk[5:])
dict_chunk: Final = json.loads(payload)
except json.JSONDecodeError as e:
raise OCIError(
status_code=500,
@ -774,8 +786,8 @@ class OCIStreamWrapper(CustomStreamWrapper):
if getattr(choice.delta, "content", None):
self._cohere_text_emitted = True
break
return result
return handle_generic_stream_chunk(dict_chunk)
return self._emit_chunk(result)
return self._emit_chunk(handle_generic_stream_chunk(dict_chunk))
__all__ = [

View file

@ -17,7 +17,7 @@ from httpx._types import CookieTypes, QueryParamTypes, RequestContent, RequestFi
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, PassthroughStreamCollector
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.passthrough.utils import CommonUtils
@ -36,6 +36,35 @@ def _as_generator(iterable: Iterator[bytes]) -> Generator[bytes, bytes, None]:
yield from iterable
class _SpendCollection:
"""Feeds relayed chunks to the provider's stream collector without letting spend tracking break the relay."""
def __init__(self, provider_config: BasePassthroughConfig, litellm_logging_obj: LiteLLMLoggingObj) -> None:
self.collector: Final[PassthroughStreamCollector] = provider_config.create_stream_collector(
model=litellm_logging_obj.model,
custom_llm_provider=litellm_logging_obj.model_call_details.get("custom_llm_provider", ""),
endpoint=litellm_logging_obj.model_call_details.get("endpoint", ""),
)
self.chunk_count = 0
self._failed = False
def add(self, chunk: bytes) -> None:
self.chunk_count += 1
if self._failed:
return
try:
self.collector.add(chunk)
except Exception as e: # noqa: BLE001 # Safe catch-all: spend tracking must never break the relayed stream
self._failed = True
verbose_logger.exception(
"Passthrough spend-tracking collector failed; spend dropped for this stream: %s", e
)
@property
def should_flush(self) -> bool:
return self.chunk_count > 0 and not self._failed
class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]):
def __init__(
self,
@ -50,8 +79,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]):
self._response: httpx.Response
self._iterator: AsyncGenerator[bytes, bytes]
self._litellm_logging_obj = litellm_logging_obj
self._provider_config = provider_config
self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks
self._spend = _SpendCollection(provider_config, litellm_logging_obj)
self._flush_scheduled = False
self._background_tasks: set[asyncio.Task] = set() # mutable-ok: instance set for background task tracking
self._hidden_params: dict[str, object] = {} # mutable-ok: router attaches response headers here in place
@ -101,16 +129,13 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]):
return _init().__await__()
def _start_flush(self) -> None:
if self._flush_scheduled or not self._raw_bytes:
if self._flush_scheduled or not self._spend.should_flush:
return
self._flush_scheduled = True
try:
task: Final = asyncio.create_task(
self._litellm_logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=self._raw_bytes,
provider_config=self._provider_config,
)
self._litellm_logging_obj.async_flush_passthrough_collected_chunks(collector=self._spend.collector)
)
self._background_tasks.add(task)
@ -118,8 +143,8 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]):
task.add_done_callback(self._background_tasks.discard)
except Exception as e: # noqa: BLE001 # Safe catch-all for verbose logging
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush; %d buffered chunks dropped: %s",
len(self._raw_bytes),
"Failed to schedule passthrough spend-tracking flush; %d collected chunks dropped: %s",
self._spend.chunk_count,
e,
)
@ -134,7 +159,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]):
await self # pyright: ignore[reportGeneralTypeIssues] # structural type check misses __await__
try:
chunk: Final = await anext(self._iterator)
self._raw_bytes.append(chunk)
self._spend.add(chunk)
except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic
self._start_flush()
try:
@ -181,13 +206,12 @@ class PassthroughStreamingResponse(Generator[bytes, bytes, None]):
self.headers = response.headers
self.status_code = response.status_code
self._litellm_logging_obj = litellm_logging_obj
self._provider_config = provider_config
self._iterator: Generator[bytes, bytes, None] = _as_generator(response.iter_bytes())
self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks
self._spend = _SpendCollection(provider_config, litellm_logging_obj)
self._flush_scheduled = False
def _start_flush(self) -> None:
if self._flush_scheduled or not self._raw_bytes:
if self._flush_scheduled or not self._spend.should_flush:
return
self._flush_scheduled = True
@ -195,14 +219,12 @@ class PassthroughStreamingResponse(Generator[bytes, bytes, None]):
try:
executor.submit(
self._litellm_logging_obj.flush_passthrough_collected_chunks,
raw_bytes=self._raw_bytes,
provider_config=self._provider_config,
self._litellm_logging_obj.flush_passthrough_collected_chunks, collector=self._spend.collector
)
except Exception as e: # noqa: BLE001 # Safe catch-all for verbose logging
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush; %d buffered chunks dropped: %s",
len(self._raw_bytes),
"Failed to schedule passthrough spend-tracking flush; %d collected chunks dropped: %s",
self._spend.chunk_count,
e,
)
@ -212,7 +234,7 @@ class PassthroughStreamingResponse(Generator[bytes, bytes, None]):
def __next__(self) -> bytes:
try:
chunk: Final = next(self._iterator)
self._raw_bytes.append(chunk)
self._spend.add(chunk)
except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic
self._start_flush()
try:

View file

@ -22,6 +22,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import (
GatewayRejected,
UpstreamOAuthFault,
UpstreamProtocolFault,
UpstreamRegistrationRefused,
UpstreamReportedFault,
)
@ -31,6 +32,7 @@ __all__ = [
"GatewayRejected",
"UpstreamOAuthFault",
"UpstreamProtocolFault",
"UpstreamRegistrationRefused",
"UpstreamReportedFault",
"classify_upstream_dcr_rejection",
"classify_upstream_token_rejection",

View file

@ -21,6 +21,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import (
GatewayRejected,
UpstreamOAuthFault,
UpstreamProtocolFault,
UpstreamRegistrationRefused,
UpstreamReportedFault,
)
@ -122,11 +123,13 @@ def classify_upstream_dcr_rejection(response: httpx.Response, log_context: str)
"""Classify a dynamic-client-registration rejection. RFC 7591 §3.2.2 errors carry
``error`` / ``error_description`` and go through the same blame assignment as token errors
(registration sends no client credentials, so credential codes stay caller-actionable); anything
without a usable ``error`` field is an upstream protocol fault."""
without a usable ``error`` field is a registration refusal for 401/403 and a protocol fault otherwise."""
parsed: Final = _safe_json(response)
fields: Final = parsed if isinstance(parsed, dict) else {}
code: Final = _bounded_field(fields.get("error"))
if code is None:
if response.status_code == 401 or response.status_code == 403:
return UpstreamRegistrationRefused(status_code=response.status_code)
_log_out_of_contract("registration", response, log_context)
return UpstreamProtocolFault(note=f"upstream registration failed with HTTP {response.status_code}")
return _classify_oauth_error_code(

View file

@ -11,7 +11,7 @@ from typing import Final
from fastapi.responses import JSONResponse
from typing_extensions import assert_never
from litellm.proxy._experimental.mcp_server.faults.types import UpstreamOAuthFault
from litellm.proxy._experimental.mcp_server.faults.types import CallerRejected, UpstreamOAuthFault
from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS
@ -35,6 +35,24 @@ def _upstream_reported_status_and_description(code: str) -> tuple[int, str]:
return 502, "the upstream authorization server reported an internal error"
def _registration_refused_description(status_code: int) -> str:
return (
f"the upstream authorization server refused dynamic client registration (HTTP {status_code}). "
"This provider may require a pre-registered OAuth client. Configure client_id and, if required "
"by the provider, client_secret for this MCP server to skip dynamic registration"
)
def _render_caller_rejected(fault: CallerRejected) -> JSONResponse:
content: Final = {
"error": fault.code,
**({"error_description": fault.description} if fault.description else {}),
**({"error_uri": fault.error_uri} if fault.error_uri else {}),
}
status_code: Final = 401 if fault.code == "invalid_client" else 400
return JSONResponse(status_code=status_code, content=content, headers=TOKEN_NO_CACHE_HEADERS)
def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse:
"""RFC 6749 §5.2 response for a token-endpoint fault. Caller-actionable rejections relay the
upstream's code on the status that code implies (401 for invalid_client per §5.2, else 400);
@ -42,13 +60,7 @@ def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse:
blamed for, or shown the internals of, a failure only the operator can fix."""
match fault.tag:
case "caller_rejected":
content: Final = {
"error": fault.code,
**({"error_description": fault.description} if fault.description else {}),
**({"error_uri": fault.error_uri} if fault.error_uri else {}),
}
status_code = 401 if fault.code == "invalid_client" else 400
return JSONResponse(status_code=status_code, content=content, headers=TOKEN_NO_CACHE_HEADERS)
return _render_caller_rejected(fault)
case "gateway_rejected":
return JSONResponse(
status_code=502,
@ -65,6 +77,13 @@ def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse:
content={"error": fault.code, "error_description": description},
headers=TOKEN_NO_CACHE_HEADERS,
)
case "upstream_registration_refused":
return _render_caller_rejected(
CallerRejected(
code="unauthorized_client",
description=_registration_refused_description(fault.status_code),
)
)
case "upstream_protocol_fault":
return JSONResponse(
status_code=502,
@ -78,7 +97,7 @@ def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse:
def dcr_fault_detail(fault: UpstreamOAuthFault) -> tuple[int, str]:
"""Status and detail string for a registration fault, raised as HTTPException by the caller.
RFC 7591 §3.2.2 defines registration errors as 400, so a contract-conformant rejection is 400
regardless of the status the upstream chose; everything else is a 502 upstream fault."""
regardless of the upstream status; a bare 401/403 is a registration refusal rendered as 403."""
match fault.tag:
case "caller_rejected":
detail: Final = f"{fault.code}: {fault.description}" if fault.description else fault.code
@ -87,6 +106,8 @@ def dcr_fault_detail(fault: UpstreamOAuthFault) -> tuple[int, str]:
return 502, _gateway_rejected_description(fault.code)
case "upstream_reported_fault":
return _upstream_reported_status_and_description(fault.code)
case "upstream_registration_refused":
return 403, _registration_refused_description(fault.status_code)
case "upstream_protocol_fault":
return 502, fault.note
case _:

View file

@ -77,4 +77,12 @@ class UpstreamProtocolFault(BaseModel):
note: str
UpstreamOAuthFault: TypeAlias = CallerRejected | GatewayRejected | UpstreamReportedFault | UpstreamProtocolFault
class UpstreamRegistrationRefused(BaseModel):
model_config = ConfigDict(frozen=True)
tag: Literal["upstream_registration_refused"] = "upstream_registration_refused"
status_code: Literal[401, 403]
UpstreamOAuthFault: TypeAlias = (
CallerRejected | GatewayRejected | UpstreamReportedFault | UpstreamProtocolFault | UpstreamRegistrationRefused
)

View file

@ -100,14 +100,24 @@ Usage with curl::
http://localhost:4000/mcp/atlassian_mcp
"""
import asyncio
import base64
import io
import json
from collections.abc import Callable, Mapping
import re
from collections.abc import AsyncIterator, Callable, Mapping
from http.cookies import CookieError, SimpleCookie
from itertools import islice
from types import MappingProxyType
from typing import Final
from urllib.parse import parse_qsl, quote, quote_plus, unquote_plus, urlencode
import httpx
from pydantic import JsonValue, TypeAdapter
from starlette.requests import HTTPConnection
from starlette.types import Message, Send
from litellm.litellm_core_utils.secret_redaction import REDACTED, redact_string
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import AuthResolution
@ -221,6 +231,10 @@ class MCPDebug:
@staticmethod
def _mask(value: str | None) -> str:
"""Mask a single value for safe display in headers."""
return MCPDebug.mask_secret(value)
@staticmethod
def mask_secret(value: str | None) -> str:
if not value:
return "(none)"
return MCPDebug._masker._mask_value(value)
@ -378,3 +392,230 @@ class MCPDebug:
server_url=server_url,
server_auth_type=server_auth_type,
)
_BODY_PREVIEW_CHARS: Final = 512
_BODY_CAPTURE_BYTES: Final = 16384
_CAPTURE_TIMEOUT_SECONDS: Final = 1.0
_CAPTURE_EXTENSION: Final = "litellm_mcp_error_preview"
_SAFE_HEADER_NAMES: Final = frozenset({"content-type", "content-length", "accept"})
_PUBLIC_HEADER_NAMES: Final = _SAFE_HEADER_NAMES | frozenset(("host", "user-agent", "accept-encoding", "connection"))
_JSON_BODY: Final = TypeAdapter(JsonValue)
_LOG_MASKER: Final = SensitiveDataMasker(visible_prefix=0, visible_suffix=0)
def _safe_text(value: str, limit: int = _BODY_PREVIEW_CHARS) -> str:
escaped: Final = "".join(json.dumps(char)[1:-1] if ord(char) < 32 or ord(char) == 127 else char for char in value)
return escaped if len(escaped) <= limit else f"{escaped[:limit]}...(truncated)"
def safe_upstream_url(url: httpx.URL) -> str:
return _safe_text(str(url.copy_with(username="", password="", path="/", query=None, fragment=None)))
def _sensitive_field(key: str) -> bool:
normalized: Final = re.sub(r"[^a-z0-9]", "", key.casefold())
return normalized in ("code", "clientassertion") or any(
pattern in normalized for pattern in _LOG_MASKER.sensitive_patterns
)
def _redact_object(
fields: Mapping[str, JsonValue],
) -> dict[str, JsonValue]: # mutable-ok: the standard JSON encoder requires dict objects
return { # mutable-ok: construct the JSON object once for the standard parser and encoder
key: REDACTED if _sensitive_field(key) else value for key, value in fields.items()
}
def _header_secret_values(name: str, value: str) -> tuple[str, ...]:
if name == "cookie":
cookie: Final = SimpleCookie[str]()
try:
cookie.load(value)
except CookieError:
return (value,)
return (value, *(item.value for item in cookie.values()))
if name not in ("authorization", "proxy-authorization"):
return (value,)
scheme, _, credential = value.partition(" ")
if scheme.lower() != "basic":
return (value, credential)
try:
decoded: Final = base64.b64decode(credential, validate=True).decode("utf-8")
except ValueError:
return (value, credential)
password: Final = decoded.partition(":")[2]
return (value, credential, decoded, password, unquote_plus(password))
def _body_secret_values(request: httpx.Request) -> tuple[str, ...] | None:
try:
raw: Final = request.content
except httpx.RequestNotRead:
return None
if not raw:
return ()
if len(raw) > _BODY_CAPTURE_BYTES:
return None
if request.headers.get("content-type", "").split(";", 1)[0].strip().lower() == "application/x-www-form-urlencoded":
return tuple(value for key, value in parse_qsl(raw.decode("utf-8", errors="replace")) if _sensitive_field(key))
try:
body: Final = _JSON_BODY.validate_json(raw)
except ValueError:
return None
from litellm.proxy._experimental.mcp_server.utils import ( # noqa: PLC0415 # MCP utils imports clients; inspect bodies only after initialization
json_string_leaves,
)
leaves: Final = json_string_leaves(body)
if leaves is None:
return None
return tuple(
value
for path, value in leaves
if not path or any(isinstance(part, str) and _sensitive_field(part) for part in path)
)
def _request_secret_values(request: httpx.Request) -> tuple[str, ...] | None:
body_values: Final = _body_secret_values(request)
if body_values is None:
return None
values: Final = (
*body_values,
request.url.password,
*(value for _, value in request.url.params.multi_items()),
*(
secret
for name, value in request.headers.items()
if name not in _PUBLIC_HEADER_NAMES
for secret in _header_secret_values(name, value)
),
)
return tuple(sorted(frozenset(value for value in values if value), key=len, reverse=True))
def _mask_known_values(value: str, secrets: tuple[str, ...]) -> str:
variants: Final = tuple(
sorted(
frozenset(
variant
for secret in secrets
for variant in (secret, json.dumps(secret)[1:-1], quote(secret, safe=""), quote_plus(secret))
),
key=len,
reverse=True,
)
)
return re.sub("|".join(re.escape(secret) for secret in variants), REDACTED, value) if variants else value
def _preview(raw: bytes, content_type: str = "", secrets: tuple[str, ...] = ()) -> str:
if not raw:
return "(empty)"
if len(raw) > _BODY_CAPTURE_BYTES:
return "(omitted: body exceeds capture limit)"
try:
parsed: Final = _JSON_BODY.validate_python(json.loads(raw, object_hook=_redact_object))
except (ValueError, RecursionError):
text: Final = raw.decode("utf-8", errors="replace")
if (
content_type.split(";", 1)[0].strip().lower() != "application/x-www-form-urlencoded"
or "=" not in text
or any(char in text for char in "<>\n\r")
):
return "(omitted: unstructured body)"
fields: Final = parse_qsl(text, keep_blank_values=True)
return _safe_text(
_mask_known_values(
urlencode(tuple((key, REDACTED if _sensitive_field(key) else value) for key, value in fields)), secrets
)
)
if not isinstance(parsed, (dict, list)):
return "(omitted: unstructured body)"
return _safe_text(redact_string(_mask_known_values(json.dumps(parsed, separators=(",", ":")), secrets)))
def _masked_headers(headers: httpx.Headers) -> str:
return _safe_text(", ".join(f"{name}={value}" for name, value in headers.items() if name in _SAFE_HEADER_NAMES))
def _request_body_preview(request: httpx.Request, secrets: tuple[str, ...] | None) -> str:
try:
return _preview(request.content, request.headers.get("content-type", ""), secrets or ())
except httpx.RequestNotRead:
return "(streamed, not captured)"
def _response_body_preview(response: httpx.Response, secrets: tuple[str, ...] | None) -> str:
if secrets is None:
return "(omitted: request credentials unavailable)"
captured: Final = response.extensions.get(_CAPTURE_EXTENSION)
if isinstance(captured, str):
return captured
try:
return _preview(response.content, response.headers.get("content-type", ""), secrets)
except httpx.ResponseNotRead:
return "(not read)"
async def _read_error_prefix(chunks: AsyncIterator[bytes], limit: int) -> bytes:
buffer: Final = io.BytesIO()
async for chunk in chunks:
buffer.write(chunk[: limit - buffer.tell()])
if buffer.tell() >= limit:
break
return buffer.getvalue()
async def capture_upstream_error_response(response: httpx.Response) -> None:
if not response.is_error:
return
try:
prefix: Final = await asyncio.wait_for(
_read_error_prefix(response.aiter_bytes(chunk_size=4096), _BODY_CAPTURE_BYTES + 1),
timeout=_CAPTURE_TIMEOUT_SECONDS,
)
response._content = prefix # pyright: ignore[reportPrivateUsage] # rebind-ok: httpx has no public setter to retain consumed bytes for auth retries
secrets: Final = _request_secret_values(response.request)
preview: Final = (
_preview(prefix, response.headers.get("content-type", ""), secrets)
if secrets is not None
else "(omitted: request credentials unavailable)"
)
except (asyncio.TimeoutError, httpx.HTTPError, httpx.StreamError):
response._content = b"" # pyright: ignore[reportPrivateUsage] # rebind-ok: httpx auth retries must survive diagnostic read failures
response.extensions[_CAPTURE_EXTENSION] = (
"(unavailable: error body read failed)" # rebind-ok: httpx response hooks communicate through extensions
)
return
response.extensions[_CAPTURE_EXTENSION] = preview # rebind-ok: httpx response hooks communicate through extensions
def describe_upstream_response(response: httpx.Response) -> str:
try:
request: Final = response.request
except RuntimeError:
return f"HTTP {response.status_code} | request unavailable"
secrets: Final = _request_secret_values(request)
return (
f"{_safe_text(request.method)} {safe_upstream_url(request.url)} -> HTTP {response.status_code}"
f" | request headers: {_masked_headers(request.headers)}"
f" | request body: {_request_body_preview(request, secrets)}"
f" | response body: {_response_body_preview(response, secrets)}"
)
def describe_upstream_http_failure(exc: BaseException) -> str | None:
from litellm.proxy._experimental.mcp_server.faults.traversal import ( # noqa: PLC0415 # fault package initialization imports the credential resolver
iter_exception_tree,
)
lines: Final = tuple(
describe_upstream_response(response)
for current in islice(iter_exception_tree(exc), 16)
for response in (getattr(current, "response", None),)
if isinstance(response, httpx.Response)
)
return " | ".join(lines) or None

View file

@ -81,7 +81,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
raise_classified_list_failure,
upstream_auth_challenge,
)
from litellm.proxy._experimental.mcp_server.mcp_debug import record_auth_resolution
from litellm.proxy._experimental.mcp_server.mcp_debug import describe_upstream_http_failure, record_auth_resolution
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
MCPPerUserTokenCache,
mcp_per_user_token_cache,
@ -255,6 +255,7 @@ _TRUE_ENV_VALUES: Final = frozenset(("1", "true", "yes", "on"))
_OAUTH_DISCOVERY_RETRY_DELAYS_SECONDS: Final = (0.05, 0.15)
_OAUTH_DISCOVERY_RETRY_BASE_SECONDS: Final = 30.0
_OAUTH_DISCOVERY_RETRY_MAX_SECONDS: Final = 900.0
_OAUTH_TEMPORARY_DISCOVERY_TTL_SECONDS: Final = 300.0
def _oauth_discovery_now() -> float:
@ -1404,6 +1405,11 @@ def _extract_upstream_auth_failure(
return upstream_auth_challenge(exc)
def _upstream_failure_suffix(exc: BaseException) -> str:
detail: Final = describe_upstream_http_failure(exc)
return f"\n upstream exchange: {detail}" if detail else ""
def _obo_retry_applies(server: MCPServer, subject_token: str | None) -> bool:
"""Whether an upstream 401/403 should invalidate the minted credential and retry once.
@ -1947,6 +1953,10 @@ class MCPServerManager:
slot: Final = self._oauth_discovery_slot(server_id)
return slot is not None and slot.generation == generation
def _expire_temporary_oauth_discovery(self, server_id: str, generation: int) -> None:
if self._oauth_discovery_slot_is_current(server_id, generation):
self._remove_oauth_discovery_slot(server_id)
def _publish_resolved_oauth_server(
self,
server: MCPServer,
@ -1959,7 +1969,13 @@ class MCPServerManager:
elif server.server_id in self.config_mcp_servers:
self.config_mcp_servers[server.server_id] = server
else:
return None
asyncio.get_running_loop().call_later(
_OAUTH_TEMPORARY_DISCOVERY_TTL_SECONDS,
self._expire_temporary_oauth_discovery,
server.server_id,
generation,
)
return server
self._remove_oauth_discovery_slot(server.server_id)
return server
@ -2051,6 +2067,12 @@ class MCPServerManager:
if slot.task is not None:
if not slot.task.done() or _oauth_discovery_now() < slot.retry_not_before:
return slot.task, slot.generation
if (
not slot.task.cancelled()
and slot.task.exception() is None
and isinstance(slot.task.result(), _OAuthDiscoveryResolved)
):
return slot.task, slot.generation
task: Final = asyncio.create_task(
self._run_oauth_metadata_resolution(self._registered_server(server), slot.generation)
)
@ -2080,7 +2102,7 @@ class MCPServerManager:
if should_defer != has_slot:
self._set_oauth_discovery_deferred(server.server_id, should_defer)
async def ensure_oauth_metadata_discovered(self, server: MCPServer) -> MCPServer:
async def ensure_oauth_metadata_discovered(self, server: MCPServer, *, _retry_stale: bool = True) -> MCPServer:
"""Join the bounded discovery task and return the resolved server.
Concurrent callers share one task per server. A failed attempt remains
@ -2107,13 +2129,13 @@ class MCPServerManager:
outcome: Final = await asyncio.shield(task)
except asyncio.CancelledError:
if task.cancelled() and not self._oauth_discovery_slot_is_current(server.server_id, generation):
return await self.ensure_oauth_metadata_discovered(server)
return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale)
raise
match outcome:
case _OAuthDiscoveryResolved(resolved_server):
return resolved_server
case _OAuthDiscoveryStale():
return await self.ensure_oauth_metadata_discovered(server)
return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale)
case _OAuthDiscoveryFailed(timed_out=timed_out):
current: Final = self._registered_server(server)
if current.is_client_forwarded_token:
@ -2125,6 +2147,14 @@ class MCPServerManager:
detail=f"OAuth metadata discovery {reason} for MCP server {server_ref!r}",
)
async def _rejoin_oauth_metadata_discovery(self, server: MCPServer, *, retry_stale: bool) -> MCPServer:
if retry_stale:
return await self.ensure_oauth_metadata_discovered(server, _retry_stale=False)
current: Final = self._registered_server(server)
if not _oauth_endpoints_unresolved(current) or current.is_client_forwarded_token:
return current
raise HTTPException(status_code=503, detail="OAuth metadata discovery changed repeatedly; retry shortly")
def _remember_upstream_initialize_instructions(self, server: MCPServer, client: MCPClient) -> None:
raw: Final[str | None] = getattr(client, "_last_initialize_instructions", None)
if raw and str(raw).strip():
@ -4340,7 +4370,9 @@ class MCPServerManager:
except MCPServerListError:
raise
except Exception as e:
verbose_logger.warning("Failed to get tools from server %s: %s", server.name, e)
verbose_logger.warning(
"Failed to get tools from server %s: %s%s", server.name, type(e).__name__, _upstream_failure_suffix(e)
)
raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge)
async def get_prompts_from_server(
@ -5085,7 +5117,9 @@ class MCPServerManager:
verbose_logger.warning("Connection error while listing tools from %s: %s", server_name, e)
raise MCPServerListError(ServerListFault(tag="unreachable"), server_name) from e
except Exception as e:
verbose_logger.warning("Error listing tools from %s: %s", server_name, e)
verbose_logger.warning(
"Error listing tools from %s: %s%s", server_name, type(e).__name__, _upstream_failure_suffix(e)
)
raise_classified_list_failure(e, server_name)
_SHORT_PREFIX_MAX_REHASH_ATTEMPTS = 1024

View file

@ -37,6 +37,7 @@ import httpx
from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError
from typing_extensions import assert_never
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
InMemoryTokenCacheBackend,
OAuthToken,
@ -101,6 +102,11 @@ async def post_client_credentials_grant(
from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415 # defer heavy handler import to call time
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # handler factory params are coarsely typed
)
from litellm.proxy._experimental.mcp_server.mcp_debug import ( # noqa: PLC0415 # diagnostics import credential enums through this package
describe_upstream_http_failure,
describe_upstream_response,
safe_upstream_url,
)
from litellm.types.llms.custom_http import httpxSpecialProvider # noqa: PLC0415 # deferred with the handler import
try:
@ -110,15 +116,28 @@ async def post_client_credentials_grant(
)
except httpx.HTTPStatusError as status_err:
status_code: Final = status_err.response.status_code
verbose_logger.warning(
"OAuth2 client_credentials token request denied:\n upstream exchange: %s",
describe_upstream_http_failure(status_err),
)
return TokenEndpointDenied(status_code=status_code, detail=f"token endpoint returned HTTP {status_code}")
except Exception as exc: # noqa: BLE001 # any transport failure is the same outcome: unreachable
return TokenEndpointUnreachable(detail=str(exc))
verbose_logger.warning(
"OAuth2 client_credentials POST %s failed: %s", safe_upstream_url(httpx.URL(url)), type(exc).__name__
)
return TokenEndpointUnreachable(detail=type(exc).__name__)
try:
body: Final = _TOKEN_BODY_ADAPTER.validate_json(response.content)
except ValidationError:
verbose_logger.warning("OAuth2 client_credentials invalid response: %s", describe_upstream_response(response))
return TokenEndpointDenied(
status_code=response.status_code, detail="token endpoint returned a non-JSON-object body"
)
access_token: Final = body.get("access_token")
if not isinstance(access_token, str) or not access_token:
verbose_logger.warning(
"OAuth2 client_credentials response has no access token | %s", describe_upstream_response(response)
)
return TokenEndpointSuccess(body=body)

View file

@ -965,7 +965,9 @@ if MCP_AVAILABLE:
apply_tool_filters=apply_tool_filters,
)
except Exception as e:
verbose_logger.exception("Error getting tools from %s: %s", server.name, e)
verbose_logger.warning(
"Error getting tools from %s: %s", server.name, classify_list_exception(e).tag
)
return (), classify_list_exception(e)
return tools_result, ServerListOk(tool_count=len(tools_result))

View file

@ -8,14 +8,17 @@ omits each feature's routes until the feature is warmed.
import asyncio
import importlib
from collections.abc import Callable
from collections.abc import Callable, Mapping, Sequence
from collections.abc import Set as AbstractSet
from dataclasses import dataclass, field
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from starlette.routing import BaseRoute, Match
from starlette.types import Receive, Scope, Send
from litellm._logging import verbose_proxy_logger
from litellm.proxy.route_priority import hot_routes_first
if TYPE_CHECKING:
from fastapi import APIRouter, FastAPI
@ -185,6 +188,31 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
module_path="litellm.proxy.management_endpoints.config_override_endpoints",
path_prefixes=("/config_overrides",),
),
LazyFeature(
name="llm_passthrough",
module_path="litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints",
path_prefixes=(
"/anthropic/",
"/assemblyai/",
"/azure/",
"/azure_ai/",
"/bedrock/",
"/cohere/",
"/comprehendmedical",
"/cursor/",
"/eu.assemblyai/",
"/gemini/",
"/gigachat/",
"/milvus/",
"/mistral/",
"/openai/",
"/openai_passthrough/",
"/vertex-ai/",
"/vertex_ai/",
"/vllm/",
"/watsonx/",
),
),
LazyFeature(
name="realtime",
module_path="litellm.proxy.realtime_endpoints.endpoints",
@ -308,14 +336,73 @@ class LazyFeatureMiddleware:
if root_path and path.startswith(root_path + "/"):
path = path[len(root_path) :] # rebind-ok: local strip after the boundary check above
for feat in self._features:
if feat.module_path in self._loaded:
if feat.module_path in self._loaded or not feat.matches(path):
continue
if feat.matches(path):
await _force_load(self._fastapi_app, feat)
if _eager_route_wins(self._fastapi_app, feat, scope):
continue
await _force_load(self._fastapi_app, feat, self._features)
await self.app(scope, receive, send)
async def _force_load(app: "FastAPI", feat: LazyFeature) -> bool:
def _lazy_slots(app: "FastAPI") -> Mapping[str, BaseRoute | None]:
return app.state.lazy_slots if hasattr(app.state, "lazy_slots") else MappingProxyType({})
def reserve_lazy_slot(app: "FastAPI", name: str, features: tuple[LazyFeature, ...] = LAZY_FEATURES) -> None:
"""Record the route the feature's router used to be included after, so its routes
are spliced back in there once it loads and keep the same precedence. Anchoring on
the route rather than its index survives later reordering of the table."""
feat: Final = next(f for f in features if f.name == name)
anchor: Final = app.router.routes[-1] if app.router.routes else None
app.state.lazy_slots = MappingProxyType({**_lazy_slots(app), feat.module_path: anchor})
def _slot_index(routes: Sequence[BaseRoute], anchor: BaseRoute | None) -> int:
if anchor is None:
return 0
return next((i + 1 for i, route in enumerate(routes) if route is anchor), len(routes))
def _eager_route_wins(app: "FastAPI", feat: LazyFeature, scope: Scope) -> bool:
"""Routes ahead of a feature's reserved slot beat its routes in Starlette's scan,
so a request one of them fully matches never needs the feature loaded."""
slots: Final = _lazy_slots(app)
if feat.module_path not in slots:
return False
ahead: Final = app.router.routes[: _slot_index(app.router.routes, slots[feat.module_path])]
return any(route.matches(scope)[0] is Match.FULL for route in ahead)
def _in_registry_order(
routes: Sequence[BaseRoute],
lazy_routes: Mapping[str, tuple[BaseRoute, ...]],
features: tuple[LazyFeature, ...],
slots: Mapping[str, BaseRoute | None],
) -> tuple[BaseRoute, ...]:
"""Lazy routers land in registry order, not first-request order, so overlapping
paths (/openai/{endpoint:path} vs /openai/v1/realtime/calls) resolve the same
way no matter which feature a deployment happens to hit first. Features with a
reserved slot go back where they were eagerly included; the rest follow every
eager route."""
rank: Final = MappingProxyType({f.module_path: i for i, f in enumerate(features)})
modules: Final = tuple(sorted(lazy_routes, key=lambda m: rank.get(m, len(rank))))
lazy_ids: Final = frozenset(id(route) for module_path in modules for route in lazy_routes[module_path])
eager: Final = tuple(route for route in routes if id(route) not in lazy_ids)
def slot_of(module_path: str) -> int:
return _slot_index(eager, slots[module_path]) if module_path in slots else len(eager)
return tuple(
route
for index in range(len(eager) + 1)
for route in (
*(r for module_path in modules if slot_of(module_path) == index for r in lazy_routes[module_path]),
*eager[index : index + 1],
)
)
async def _force_load(app: "FastAPI", feat: LazyFeature, features: tuple[LazyFeature, ...] = LAZY_FEATURES) -> bool:
"""Import + register a lazy feature exactly once per (app, module).
Shared by the middleware and the /lazy/warm endpoint."""
if not hasattr(app.state, "lazy_loaded"):
@ -330,7 +417,18 @@ async def _force_load(app: "FastAPI", feat: LazyFeature) -> bool:
# mutates app.router.routes, so it stays on the loop thread.
loop: Final = asyncio.get_running_loop()
module: Final = await loop.run_in_executor(None, importlib.import_module, feat.module_path)
before: Final = len(app.router.routes)
feat.register_fn(app, module)
previous: Final[Mapping[str, tuple[BaseRoute, ...]]] = (
app.state.lazy_routes if hasattr(app.state, "lazy_routes") else MappingProxyType({})
)
lazy_routes: Final[Mapping[str, tuple[BaseRoute, ...]]] = MappingProxyType(
{**previous, feat.module_path: tuple(app.router.routes[before:])}
)
app.state.lazy_routes = lazy_routes # rebind-ok: the app owns the record of which routes each feature added
app.router.routes[:] = hot_routes_first( # rebind-ok: the app owns its route table
_in_registry_order(app.router.routes, lazy_routes, features, _lazy_slots(app))
)
app.state.lazy_loaded.add(feat.module_path)
app.openapi_schema = None
verbose_proxy_logger.info(

File diff suppressed because it is too large Load diff

View file

@ -2602,6 +2602,15 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
global_max_parallel_requests: int | None = Field(
None, description="global max parallel requests to allow for a proxy instance."
)
user_api_key_cache_max_size: int | None = Field(
None,
gt=0,
description=(
"max number of entries (virtual keys, teams, users, end users, memberships, ...) each worker keeps in "
"its in-memory auth cache. Defaults to 200. Raise this if you have more active keys than that or auth "
"lookups keep hitting the DB"
),
)
max_request_size_mb: int | None = Field(
None,
description="max request size in MB, if a request is larger than this size it will be rejected",
@ -2848,6 +2857,25 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"UI username/password login. Default is False."
),
)
disable_responses_id_security: bool | None = Field(
None,
description=(
"If True, disables ownership enforcement on Responses API ids. "
"Keys may then retrieve, cancel, delete, and chain from any response id, "
"including ids belonging to another user or team and ids this proxy never issued. "
"WARNING: this removes tenant isolation on /v1/responses"
),
)
allow_unmanaged_response_ids: bool | None = Field(
None,
description=(
"If True, lets keys address Responses API ids that this proxy did not issue "
"(raw provider ids, or ids issued before response-id encryption was configured). "
"Such an id carries no owner, so no ownership check can run on it; ids this proxy "
"did issue keep full ownership enforcement. Off by default, in which case an "
"unrecognized response id is rejected with 403"
),
)
disable_env_credential_login: bool | None = Field(
None,
description=(

View file

@ -196,9 +196,7 @@ class AuthCacheInvalidationSubscriber:
for additional_cache in self._additional_in_memory_caches:
additional_cache.set_cache(parsed.cache_key, parsed.new_value, ttl=parsed.ttl)
return
in_memory_cache: Final = self._user_api_key_cache.in_memory_cache
if in_memory_cache is not None:
in_memory_cache.delete_cache(parsed.cache_key)
self._user_api_key_cache.in_memory_cache_for(parsed.cache_key).delete_cache(parsed.cache_key)
for additional_cache in self._additional_in_memory_caches:
additional_cache.delete_cache(parsed.cache_key)

View file

@ -147,8 +147,11 @@ async def memory_usage_in_mem_cache(
llm_router.cache.in_memory_cache.ttl_dict
)
num_items_in_user_api_key_cache: Final = len(user_api_key_cache.in_memory_cache.cache_dict) + len(
user_api_key_cache.in_memory_cache.ttl_dict
num_items_in_user_api_key_cache: Final = (
len(user_api_key_cache.in_memory_cache.cache_dict)
+ len(user_api_key_cache.in_memory_cache.ttl_dict)
+ len(user_api_key_cache.key_object_cache.in_memory_cache.cache_dict)
+ len(user_api_key_cache.key_object_cache.in_memory_cache.ttl_dict)
)
num_items_in_proxy_logging_obj_cache: Final = len(
@ -189,6 +192,8 @@ async def memory_usage_in_mem_cache_items(
return {
"user_api_key_cache": user_api_key_cache.in_memory_cache.cache_dict,
"user_api_key_ttl": user_api_key_cache.in_memory_cache.ttl_dict,
"user_key_object_cache": user_api_key_cache.key_object_cache.in_memory_cache.cache_dict,
"user_key_object_ttl": user_api_key_cache.key_object_cache.in_memory_cache.ttl_dict,
"llm_router_cache": llm_router_in_memory_cache_dict,
"llm_router_ttl": llm_router_in_memory_ttl_dict,
"proxy_logging_obj_cache": proxy_logging_obj.internal_usage_cache.dual_cache.in_memory_cache.cache_dict,
@ -294,7 +299,9 @@ async def get_memory_summary(
try:
# User API key cache
user_cache_items: Final = len(user_api_key_cache.in_memory_cache.cache_dict)
user_cache_items: Final = len(user_api_key_cache.in_memory_cache.cache_dict) + len(
user_api_key_cache.key_object_cache.in_memory_cache.cache_dict
)
total_cache_items += user_cache_items
caches["user_api_keys"] = {
"count": user_cache_items,
@ -429,10 +436,16 @@ def _get_cache_memory_stats(
cache_stats: Final[dict[str, object]] = {}
try:
# User API key cache
user_cache_size: Final = sys.getsizeof(user_api_key_cache.in_memory_cache.cache_dict)
user_ttl_size: Final = sys.getsizeof(user_api_key_cache.in_memory_cache.ttl_dict)
key_object_in_memory_cache: Final = user_api_key_cache.key_object_cache.in_memory_cache
user_cache_size: Final = sys.getsizeof(user_api_key_cache.in_memory_cache.cache_dict) + sys.getsizeof(
key_object_in_memory_cache.cache_dict
)
user_ttl_size: Final = sys.getsizeof(user_api_key_cache.in_memory_cache.ttl_dict) + sys.getsizeof(
key_object_in_memory_cache.ttl_dict
)
cache_stats["user_api_key_cache"] = {
"num_items": len(user_api_key_cache.in_memory_cache.cache_dict),
"num_items": len(user_api_key_cache.in_memory_cache.cache_dict)
+ len(key_object_in_memory_cache.cache_dict),
"cache_dict_size_bytes": user_cache_size,
"ttl_dict_size_bytes": user_ttl_size,
"total_size_mb": round((user_cache_size + user_ttl_size) / (1024 * 1024), 2),

View file

@ -1,11 +1,15 @@
from __future__ import annotations
import re
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any, Final, TypeVar, cast, overload
from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
from litellm.caching.dual_cache import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.caching.redis_cache import RedisCache
from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
@ -14,6 +18,13 @@ if TYPE_CHECKING:
T = TypeVar("T", bound=BaseModel)
_HASHED_TOKEN_CACHE_KEY: Final = re.compile(r"[0-9a-f]{64}")
def is_user_key_cache_key(key: str) -> bool:
"""Only user-key objects are cached under a bare ``hash_token`` digest; every other object uses a prefixed key."""
return _HASHED_TOKEN_CACHE_KEY.fullmatch(key) is not None
class UserApiKeyCache(DualCache):
"""
@ -36,10 +47,50 @@ class UserApiKeyCache(DualCache):
``async_set_cache_pipeline`` applies the same untyped Codec pass as omitting
``model_type`` on ``async_set_cache`` (so ``BaseModel`` rows are dumped before Redis).
User-key objects (see ``is_user_key_cache_key``) live in their own in-memory partition,
``key_object_cache``, so churn in the other management objects cannot evict them. Both
partitions share the same Redis backend and TTL settings.
``get_cache`` / ``async_get_cache`` overloads and implementations must be contiguous
(no other methods in between) so mypy resolves ``@overload`` + implementation correctly.
"""
def __init__(
self,
in_memory_cache: InMemoryCache | None = None,
redis_cache: RedisCache | None = None,
default_in_memory_ttl: float | None = None,
default_redis_ttl: float | None = None,
key_object_in_memory_cache: InMemoryCache | None = None,
) -> None:
super().__init__(
in_memory_cache=in_memory_cache,
redis_cache=redis_cache,
default_in_memory_ttl=default_in_memory_ttl,
default_redis_ttl=default_redis_ttl,
)
self.key_object_cache: Final = DualCache(
in_memory_cache=key_object_in_memory_cache or InMemoryCache(),
redis_cache=redis_cache,
default_in_memory_ttl=default_in_memory_ttl,
default_redis_ttl=default_redis_ttl,
)
def in_memory_cache_for(self, key: str) -> InMemoryCache:
return self.key_object_cache.in_memory_cache if is_user_key_cache_key(key) else self.in_memory_cache
def update_cache_ttl(self, default_in_memory_ttl: float | None, default_redis_ttl: float | None) -> None:
super().update_cache_ttl(default_in_memory_ttl=default_in_memory_ttl, default_redis_ttl=default_redis_ttl)
self.key_object_cache.update_cache_ttl(
default_in_memory_ttl=default_in_memory_ttl, default_redis_ttl=default_redis_ttl
)
def attach_redis_cache(
self, redis_cache: RedisCache | None = None, *, default_redis_ttl: float | None = None
) -> None:
super().attach_redis_cache(redis_cache, default_redis_ttl=default_redis_ttl)
self.key_object_cache.attach_redis_cache(redis_cache, default_redis_ttl=default_redis_ttl)
@overload
def get_cache(
self,
@ -71,7 +122,11 @@ class UserApiKeyCache(DualCache):
) -> object:
if model_type is None and "model_type" in kwargs:
model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
cached: Final = super().get_cache(key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs)
cached: Final = (
self.key_object_cache.get_cache(key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs)
if is_user_key_cache_key(key)
else super().get_cache(key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs)
)
if model_type is None:
return cached
if cached is None:
@ -117,8 +172,14 @@ class UserApiKeyCache(DualCache):
) -> object:
if model_type is None and "model_type" in kwargs:
model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
cached: Final = await super().async_get_cache(
key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs
cached: Final = (
await self.key_object_cache.async_get_cache(
key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs
)
if is_user_key_cache_key(key)
else await super().async_get_cache(
key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs
)
)
if model_type is None:
return cached
@ -137,20 +198,49 @@ class UserApiKeyCache(DualCache):
def set_cache(self, key: str | None, value: object, local_only: bool = False, **kwargs: object):
model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
if key is not None and is_user_key_cache_key(key):
return self.key_object_cache.set_cache(key=key, value=payload, local_only=local_only, **kwargs)
return super().set_cache(key=key, value=payload, local_only=local_only, **kwargs)
async def async_set_cache(self, key: str | None, value: object, local_only: bool = False, **kwargs: object):
model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
if key is not None and is_user_key_cache_key(key):
return await self.key_object_cache.async_set_cache(key=key, value=payload, local_only=local_only, **kwargs)
return await super().async_set_cache(key=key, value=payload, local_only=local_only, **kwargs)
async def async_set_cache_pipeline(self, cache_list: list, local_only: bool = False, **kwargs: object) -> None:
def delete_cache(self, key: str) -> None:
if is_user_key_cache_key(key):
self.key_object_cache.delete_cache(key)
return
super().delete_cache(key)
async def async_delete_cache(self, key: str) -> None:
if is_user_key_cache_key(key):
await self.key_object_cache.async_delete_cache(key)
return
await super().async_delete_cache(key)
def flush_cache(self) -> None:
super().flush_cache()
self.key_object_cache.in_memory_cache.flush_cache()
async def async_set_cache_pipeline(
self, cache_list: Sequence[tuple[str, object]], local_only: bool = False, **kwargs: object
) -> None:
"""
Batch writes with the same Codec boundary as ``async_set_cache`` without
``model_type``: ``BaseModel`` values become JSON-safe dicts; dicts/scalars unchanged.
"""
normalized: Final = [(key, CacheCodec.serialize(value, model_type=None)) for key, value in cache_list]
return await super().async_set_cache_pipeline(cache_list=normalized, local_only=local_only, **kwargs)
normalized: Final = tuple((key, CacheCodec.serialize(value, model_type=None)) for key, value in cache_list)
key_object_entries: Final = tuple(entry for entry in normalized if is_user_key_cache_key(entry[0]))
other_entries: Final = tuple(entry for entry in normalized if not is_user_key_cache_key(entry[0]))
if key_object_entries:
await self.key_object_cache.async_set_cache_pipeline(
cache_list=key_object_entries, local_only=local_only, **kwargs
)
if other_entries:
await super().async_set_cache_pipeline(cache_list=other_entries, local_only=local_only, **kwargs)
#: Value cached under ``user_object_permission_id_cache_key`` when the user links no permission row,

View file

@ -32,6 +32,29 @@ if TYPE_CHECKING:
_RESPONSES_API_PROVIDER_PREFIX: Final = "/openai"
_RESPONSES_API_CREATE_ROUTES: Final = frozenset({"/v1/responses", "/responses"})
_ADDRESSED_RESPONSE_ID_KEY: Final = "_litellm_addressed_response_id"
_UNMANAGED_RESPONSE_ID_DETAIL: Final = (
"Forbidden. This response id was not issued by this proxy, so the proxy cannot tell who owns it. "
"To let keys address responses this proxy did not issue, set "
"general_settings::allow_unmanaged_response_ids to True in the config.yaml file."
)
_PROXY_ADMIN_ROLES: Final = frozenset({LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN.value})
def _proxy_general_settings() -> Mapping[str, Any]:
from litellm.proxy.proxy_server import general_settings
return general_settings
def _proxy_signing_key() -> str | None:
import os
from litellm.proxy.proxy_server import master_key
salt_key: Final = os.getenv("LITELLM_SALT_KEY", None)
return master_key if salt_key is None else salt_key
_RESPONSE_PAYLOAD_ADAPTER: Final = TypeAdapter(Mapping[str, object])
@ -83,8 +106,13 @@ def _is_responses_api_create_route(request_route: str | None) -> bool:
class ResponsesIDSecurity(CustomLogger):
def __init__(self):
pass
def __init__(
self,
general_settings_reader: Callable[[], Mapping[str, Any]] = _proxy_general_settings,
signing_key_reader: Callable[[], str | None] = _proxy_signing_key,
) -> None:
self._general_settings_reader: Final = general_settings_reader
self._signing_key_reader: Final = signing_key_reader
async def async_pre_call_hook(
self,
@ -103,30 +131,51 @@ class ResponsesIDSecurity(CustomLogger):
}
if call_type not in responses_api_call_types:
return None
if call_type == "aresponses":
# check 'previous_response_id' if present in the data
previous_response_id: Final = data.get("previous_response_id")
if previous_response_id and self._is_encrypted_response_id(previous_response_id):
original_response_id, user_id, team_id = self._decrypt_response_id(previous_response_id)
self.check_user_access_to_response_id(user_id, team_id, user_api_key_dict)
data["previous_response_id"] = original_response_id
elif call_type in {"aget_responses", "adelete_responses", "acancel_responses", "alist_input_items"}:
response_id: Final = data.get("response_id")
if response_id and self._is_encrypted_response_id(response_id):
original_response_id, user_id, team_id = self._decrypt_response_id(response_id)
self.check_user_access_to_response_id(user_id, team_id, user_api_key_dict)
data["response_id"] = original_response_id
addressed_id_field: Final = "previous_response_id" if call_type == "aresponses" else "response_id"
retained_id: Final = data.get(_ADDRESSED_RESPONSE_ID_KEY)
addressed_id: Final = (
retained_id if isinstance(retained_id, str) and retained_id else data.get(addressed_id_field)
)
if not isinstance(addressed_id, str) or not addressed_id:
return data
authorized_id: Final = self._authorize_response_id(addressed_id, user_api_key_dict)
data[addressed_id_field] = authorized_id
data[_ADDRESSED_RESPONSE_ID_KEY] = addressed_id
return data
def _authorize_response_id(
self,
response_id: str,
user_api_key_dict: "UserAPIKeyAuth",
) -> str:
if self._is_encrypted_response_id(response_id):
original_response_id, user_id, team_id = self._decrypt_response_id(response_id)
self.check_user_access_to_response_id(user_id, team_id, user_api_key_dict)
return original_response_id
if self._unmanaged_response_ids_allowed(user_api_key_dict):
return response_id
raise HTTPException(status_code=403, detail=_UNMANAGED_RESPONSE_ID_DETAIL)
def _unmanaged_response_ids_allowed(self, user_api_key_dict: "UserAPIKeyAuth") -> bool:
general_settings: Final = self._general_settings_reader()
if general_settings.get("disable_responses_id_security", False):
return True
if general_settings.get("allow_unmanaged_response_ids", False):
return True
if self._get_signing_key() is None:
return True
return user_api_key_dict.user_role in _PROXY_ADMIN_ROLES
def check_user_access_to_response_id(
self,
response_id_user_id: str | None,
response_id_team_id: str | None,
user_api_key_dict: "UserAPIKeyAuth",
) -> bool:
from litellm.proxy.proxy_server import general_settings
general_settings: Final = self._general_settings_reader()
if (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
@ -219,15 +268,7 @@ class ResponsesIDSecurity(CustomLogger):
return response_id, None, None
def _get_signing_key(self) -> str | None:
"""Get the signing key for encryption/decryption."""
import os
from litellm.proxy.proxy_server import master_key
salt_key = os.getenv("LITELLM_SALT_KEY", None)
if salt_key is None:
salt_key = master_key
return salt_key
return self._signing_key_reader()
def _encrypt_response_id(
self,
@ -274,7 +315,7 @@ class ResponsesIDSecurity(CustomLogger):
This method adds response IDs to an in-memory queue, which are then
batch-processed by the DBSpendUpdateWriter during regular database update cycles.
"""
from litellm.proxy.proxy_server import general_settings
general_settings: Final = self._general_settings_reader()
if general_settings.get("disable_responses_id_security", False):
return response
@ -288,7 +329,7 @@ class ResponsesIDSecurity(CustomLogger):
async def async_post_call_streaming_iterator_hook(
self, user_api_key_dict: "UserAPIKeyAuth", response: Any, request_data: dict
) -> AsyncGenerator[BaseLiteLLMOpenAIResponseObject, None]:
from litellm.proxy.proxy_server import general_settings
general_settings: Final = self._general_settings_reader()
# Create a request-scoped cache for consistent encryption across streaming chunks.
request_encryption_cache: Final[dict[str, str]] = {}

View file

@ -3729,7 +3729,7 @@ async def delete_key_fn(
)
verbose_proxy_logger.debug(
"/keys/delete - cache after delete: %s", user_api_key_cache.in_memory_cache.cache_dict
"/keys/delete - cache after delete: %s", user_api_key_cache.key_object_cache.in_memory_cache.cache_dict
)
asyncio.create_task(

View file

@ -97,7 +97,6 @@ else:
vertex_llm_base: Final = VertexBase()
router: Final = APIRouter()
openai_passthrough_router: Final = APIRouter()
default_vertex_config: Final = None
passthrough_endpoint_router: Final = PassthroughEndpointRouter()
@ -2297,11 +2296,6 @@ async def vertex_proxy_route(
)
@openai_passthrough_router.api_route(
"/openai_passthrough/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
tags=["OpenAI Pass-through", "pass-through"],
)
@router.api_route(
"/openai/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],

View file

@ -0,0 +1,44 @@
"""/openai_passthrough must be matched ahead of the native /{provider}/v1/files and
/{provider}/v1/batches routes, so unlike the other provider passthrough routes it is
registered at startup and defers to the lazily loaded handler per call."""
from typing import Final
from fastapi import APIRouter, Depends, Request, Response
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
router: Final = APIRouter()
@router.api_route(
"/openai_passthrough/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
tags=["OpenAI Pass-through", "pass-through"],
)
async def openai_passthrough_route(
endpoint: str,
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> Response:
"""
Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native
implementations (e.g. the Responses API at /v1/responses).
Examples:
- /openai_passthrough/v1/responses
- /openai_passthrough/v1/responses/{response_id}
- /openai_passthrough/v1/responses/{response_id}/input_items
[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)
"""
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import openai_proxy_route
return await openai_proxy_route(
endpoint=endpoint,
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
)

View file

@ -304,7 +304,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import (
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._lazy_features import attach_lazy_features
from litellm.proxy._lazy_features import attach_lazy_features, reserve_lazy_slot
from litellm.proxy._types import *
from litellm.proxy.analytics_endpoints.analytics_endpoints import (
router as analytics_router,
@ -639,13 +639,8 @@ from litellm.proxy.openai_files_endpoints.files_endpoints import (
from litellm.proxy.openai_files_endpoints.files_endpoints import (
set_files_config,
)
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
openai_passthrough_router,
passthrough_endpoint_router,
vertex_ai_live_websocket_passthrough,
)
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
router as llm_passthrough_router,
from litellm.proxy.pass_through_endpoints.openai_passthrough_endpoints import (
router as openai_passthrough_router,
)
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
initialize_pass_through_endpoints,
@ -660,6 +655,7 @@ from litellm.proxy.rag_endpoints.endpoints import router as rag_router
from litellm.proxy.rerank_endpoints.endpoints import router as rerank_router
from litellm.proxy.response_api_endpoints.endpoints import router as response_router
from litellm.proxy.route_llm_request import route_request
from litellm.proxy.route_priority import hot_routes_first
from litellm.proxy.search_endpoints.endpoints import router as search_router
from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
@ -5813,6 +5809,16 @@ class ProxyConfig:
default_redis_ttl=ttl,
)
### USER API KEY CACHE MAX SIZE (in-memory tier shared by keys, teams, users, end users, ...) ###
if "user_api_key_cache_max_size" in general_settings:
user_api_key_cache.update_in_memory_max_size(
ConfigGeneralSettings.model_validate(
MappingProxyType(
{"user_api_key_cache_max_size": general_settings["user_api_key_cache_max_size"]}
)
).user_api_key_cache_max_size
)
### PKCE MULTI-INSTANCE PREREQUISITE CHECK ###
# PKCE verifiers are stored in redis_usage_cache when available so they can
# be read back by any instance (not just the one that started the auth flow).
@ -6047,6 +6053,10 @@ class ProxyConfig:
set_files_config(config=files_config)
## default config for vertex ai routes
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
passthrough_endpoint_router,
)
default_vertex_config: Final = config.get("default_vertex_config", None)
passthrough_endpoint_router.set_default_vertex_config(config=default_vertex_config)
@ -7059,6 +7069,23 @@ class ProxyConfig:
"enable_openai_websocket_passthrough"
)
if "user_api_key_cache_max_size" not in self._yaml_general_settings_keys:
db_cache_max_size: Final = _general_settings.get("user_api_key_cache_max_size")
try:
cache_max_size: Final = ConfigGeneralSettings.model_validate(
MappingProxyType({"user_api_key_cache_max_size": db_cache_max_size})
).user_api_key_cache_max_size
except ValidationError:
verbose_proxy_logger.warning(
"Ignoring invalid general_settings.user_api_key_cache_max_size=%r from the DB", db_cache_max_size
)
else:
if cache_max_size is None:
general_settings.pop("user_api_key_cache_max_size", None)
else:
general_settings["user_api_key_cache_max_size"] = cache_max_size
user_api_key_cache.update_in_memory_max_size(cache_max_size)
## STORE MODEL IN DB ##
if "store_model_in_db" in _general_settings:
value = _general_settings["store_model_in_db"]
@ -11763,6 +11790,10 @@ async def vertex_ai_live_passthrough_endpoint(
This endpoint delegates to the WebSocket function defined in llm_passthrough_endpoints.py
"""
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
vertex_ai_live_websocket_passthrough,
)
return await vertex_ai_live_websocket_passthrough(
websocket=websocket,
model=model,
@ -16967,6 +16998,7 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro
"cancel_on_disconnect": "Boolean",
"disable_auto_add_proxy_admin_to_teams": "Boolean",
"apply_user_budget_to_team_keys": "Boolean",
"user_api_key_cache_max_size": "Integer",
}
)
@ -18668,7 +18700,7 @@ app.include_router(credential_router)
app.include_router(openai_passthrough_router)
app.include_router(batches_router)
app.include_router(openai_files_router)
app.include_router(llm_passthrough_router)
reserve_lazy_slot(app, "llm_passthrough")
app.include_router(pass_through_router)
app.include_router(health_router)
app.include_router(key_management_router)
@ -18708,6 +18740,7 @@ app.include_router(ui_discovery_endpoints_router)
app.include_router(google_router)
attach_lazy_features(app)
app.router.routes = hot_routes_first(app.router.routes)
app.add_middleware(
RequestSizeLimitMiddleware,
get_max_request_size_mb=lambda: general_settings.get("max_request_size_mb"),

View file

@ -0,0 +1,24 @@
"""Starlette matches routes in registration order, so the routes that take the most traffic go first."""
from collections.abc import Sequence
from typing import Final
from starlette.routing import BaseRoute, Route
HOT_ROUTE_PATHS: Final[frozenset[str]] = frozenset(
(
"/health/liveliness",
"/health/liveness",
"/v1/chat/completions",
"/chat/completions",
"/v1/messages",
)
)
def _is_hot(route: BaseRoute) -> bool:
return isinstance(route, Route) and route.path in HOT_ROUTE_PATHS
def hot_routes_first(routes: Sequence[BaseRoute]) -> list[BaseRoute]: # mutable-ok: assigned to Router.routes, a list
return sorted(routes, key=lambda route: not _is_hot(route))

View file

@ -12669,6 +12669,7 @@ class Router:
input: str | list | None = None,
specific_deployment: bool | None = False,
parent_otel_span: Span | None = None,
health_check_probe: bool = False,
) -> list[dict] | dict:
"""
Get the healthy deployments for a model.
@ -12718,6 +12719,7 @@ class Router:
healthy_deployments = await self._async_filter_health_check_unhealthy_deployments(
healthy_deployments=healthy_deployments,
parent_otel_span=parent_otel_span,
health_check_probe=health_check_probe,
)
cooldown_deployments: Final = await _async_get_cooldown_deployments(
@ -14100,6 +14102,7 @@ class Router:
self,
healthy_deployments: list[dict],
parent_otel_span: Span | None = None,
health_check_probe: bool = False,
) -> list[dict]:
"""
Filter out deployments marked unhealthy by background health checks.
@ -14136,8 +14139,7 @@ class Router:
]
if not filtered:
verbose_router_logger.warning("All deployments marked unhealthy by health checks, bypassing health filter")
return healthy_deployments
return [] if health_check_probe else healthy_deployments # mutable-ok: empty list signals unavailable probe
return filtered

View file

@ -461,6 +461,8 @@ class AdaptiveRouter:
if d_alpha == 0 and d_beta == 0:
continue
cell_key = (attribution_type, target_model)
if cell_key not in self._cells:
continue
self._cells[cell_key] = apply_delta(
self._cells[cell_key],
d_alpha,

View file

@ -270,6 +270,18 @@ change or default takeover records `cause: modality_escalation` with the displac
pinned by session affinity, and by default a KEPT session pin bypasses the gate: a session pinned
to a text-only model keeps it even when an image arrives.
Context-window and modality recovery take priority over the default model. If a compatible tier
cannot serve, the router checks the remaining compatible recovery tiers before using `default_model`.
A capacity failure without those constraints tries the selected tier's peers, then the default
The default must fit the context and accept the request's modality. It cannot bypass routing plugins
or a plan-mode floor. Context fit uses the auto-router's existing buffer even when Router-wide pre-call
checks are off. Missing context metadata retains the existing unknown-window behavior
Health fallback records `cause: health_default_fallback` and `health_displaced:<MODEL>` in `signals`.
It does not replace the session's tier pin. Adaptive feedback retains the model that actually served,
but a default outside the adaptive candidate pool does not become a normal candidate
Add `modality_pin_override: true` to lift that last exemption. The image turn is then re-placed
the same way every other decision is, and records `cause: modality_pin_override` whether or not
the tier moved, since the model left the pin either way. The pin itself is untouched: the session

View file

@ -937,6 +937,7 @@ def _decision_is_pinnable(decision: StandardLoggingRoutingDecision | None) -> bo
"modality_escalation",
"modality_pin_override",
"health_failover",
"health_default_fallback",
)
and not decision.get("context_escalated")
and _CLASSIFIER_CIRCUIT_OPEN_SIGNAL not in (decision.get("signals") or ())
@ -1098,6 +1099,15 @@ def _group_provably_fits(facts: tuple[int | None, bool], needed: int, buffer: fl
return window is not None and not has_unknown and needed <= int(window * buffer)
class _RequestContextFit(NamedTuple):
facts: Mapping[str, tuple[int | None, bool]]
needed: int | None
buffer: float
def accepts(self, model: str) -> bool:
return self.needed is None or _window_can_hold(self.facts.get(model, (None, True))[0], self.needed, self.buffer)
class _ContextWindowPlacement(NamedTuple):
"""Where the context-window gate placed the request: the placement tier, the subset of its
pool the pick may use, and every configured group not provably misfit (the adaptive filter)."""
@ -2681,12 +2691,32 @@ class ComplexityRouter(CustomLogger):
verbose_router_logger.debug("ComplexityRouter: context-window token count failed. Got - %s", e)
return None
async def _request_context_fit(
self,
resolved_messages: Sequence[Mapping[str, object]] | None,
request_kwargs: Mapping[str, object],
) -> _RequestContextFit:
if not self.config.enable_context_window_escalation or not resolved_messages:
return _RequestContextFit(EMPTY_MAPPING, None, self.config.context_window_escalation_buffer)
names: Final = frozenset(model for pool in self._tier_pools().values() for model in pool) | frozenset(
(self.config.default_model,) if self.config.default_model else ()
)
facts: Final = MappingProxyType({name: self._group_window_facts(name) for name in names})
known: Final = tuple(window for window, _ in facts.values() if window is not None)
buffer: Final = self.config.context_window_escalation_buffer
needs_count: Final = known and self._request_byte_upper_bound(resolved_messages, request_kwargs) > int(
min(known) * buffer
)
needed: Final = await self._counted_request_tokens(resolved_messages, request_kwargs) if needs_count else None
return _RequestContextFit(facts=facts, needed=needed, buffer=buffer)
async def _context_window_placement(
self,
tier: ComplexityTier | str,
resolved_messages: Sequence[Mapping[str, object]] | None,
request_kwargs: Mapping[str, object],
pool_override: tuple[str, ...] | None = None,
context_fit: _RequestContextFit | None = None,
) -> _ContextWindowPlacement | None:
"""Correct a decided placement whose models provably cannot hold the prompt, or None
(the placement stands). Only a real tokenizer count ever moves a request, escalation
@ -2698,17 +2728,10 @@ class ComplexityRouter(CustomLogger):
pool: Final = pool_override if pool_override is not None else tuple(pools.get(_tier_name(tier), ()))
if not pool:
return None
facts: Final = MappingProxyType({group: self._group_window_facts(group) for group in pool})
known_windows: Final = tuple(window for window, _ in facts.values() if window is not None)
if not known_windows:
fit: Final = context_fit or await self._request_context_fit(resolved_messages, request_kwargs)
if fit.needed is None:
return None
buffer: Final = self.config.context_window_escalation_buffer
if self._request_byte_upper_bound(resolved_messages, request_kwargs) <= int(min(known_windows) * buffer):
return None
needed: Final = await self._counted_request_tokens(resolved_messages, request_kwargs)
if needed is None:
return None
return self._placement_for_tokens(tier=tier, pool=pool, pools=pools, facts=facts, needed=needed)
return self._placement_for_tokens(tier=tier, pool=pool, pools=pools, facts=fit.facts, needed=fit.needed)
def _placement_for_tokens(
self,
@ -2720,14 +2743,16 @@ class ComplexityRouter(CustomLogger):
needed: int,
) -> _ContextWindowPlacement | None:
buffer: Final = self.config.context_window_escalation_buffer
in_tier: Final = tuple(group for group in pool if _window_can_hold(facts[group][0], needed, buffer))
in_tier: Final = tuple(
group for group in pool if _window_can_hold(facts.get(group, (None, True))[0], needed, buffer)
)
if in_tier and len(in_tier) == len(pool):
return None
holdable: Final = frozenset(
group
for tier_pool in pools.values()
for group in tier_pool
if _window_can_hold(self._group_window_facts(group)[0], needed, buffer)
if _window_can_hold(facts.get(group, (None, True))[0], needed, buffer)
)
if in_tier:
return _ContextWindowPlacement(tier=tier, allowed_models=in_tier, holdable_models=holdable)
@ -2735,7 +2760,7 @@ class ComplexityRouter(CustomLogger):
proven = tuple(
group
for group in pools.get(name, ())
if _group_provably_fits(self._group_window_facts(group), needed, buffer)
if _group_provably_fits(facts.get(group, (None, True)), needed, buffer)
)
if proven:
return _ContextWindowPlacement(
@ -2881,6 +2906,7 @@ class ComplexityRouter(CustomLogger):
messages: list[dict[str, Any]] | None, # mutable-ok: forwarded verbatim to the list-typed re-pick
resolved_messages: Sequence[Mapping[str, object]] | None,
request_kwargs: dict, # mutable-ok: same shape the hook receives
context_fit: _RequestContextFit | None = None,
) -> PreRoutingHookResponse:
"""Replace a routed model that cannot accept this request's image input.
@ -2911,7 +2937,8 @@ class ComplexityRouter(CustomLogger):
or self._model_accepts_image_input(response.model)
):
return response
eligible: Final = self._modality_eligible_models()
fit: Final = context_fit or await self._request_context_fit(resolved_messages, request_kwargs)
eligible: Final = frozenset(name for name in self._modality_eligible_models() if fit.accepts(name))
names: Final = self.config.tier_names()
pools: Final = self._tier_pools()
decided: Final = decision.get("tier") if decision is not None else None
@ -3030,12 +3057,15 @@ class ComplexityRouter(CustomLogger):
Every way the owner says "nothing here can serve this" is a negative verdict: no healthy
deployment for the group at all (BadRequestError, which ContextWindowExceededError
subclasses), every deployment filtered out (RouterRateLimitError), and every deployment
over its RPM (RouterRateLimitErrorBasic). Anything else is unknown rather than negative,
so it reads as capacity: absent information must never decide the verdict.
subclasses), every deployment filtered out (RouterRateLimitError), every deployment over
its RPM (RouterRateLimitErrorBasic), and every deployment refused by a filter that reports
exhaustion as a bare ValueError naming a RouterErrors marker -- provider and deployment
budgets, and tag routing, which have no typed error of their own. Anything else is unknown
rather than negative, so it reads as capacity: absent information must never decide the
verdict.
"""
from litellm.exceptions import BadRequestError
from litellm.types.router import RouterRateLimitError, RouterRateLimitErrorBasic
from litellm.types.router import RouterErrors, RouterRateLimitError, RouterRateLimitErrorBasic
probe_kwargs: Final = dict(request_kwargs) # mutable-ok: the owner pops routing keys off the dict it is handed
try:
@ -3045,10 +3075,15 @@ class ComplexityRouter(CustomLogger):
messages=messages,
input=input,
parent_otel_span=_get_parent_otel_span_from_kwargs(request_kwargs),
health_check_probe=True,
)
except (RouterRateLimitError, RouterRateLimitErrorBasic, BadRequestError):
except (RouterRateLimitError, RouterRateLimitErrorBasic, BadRequestError) as exc:
verbose_router_logger.debug("health probe unavailable model=%s error=%s", model_name, type(exc).__name__)
return False
except Exception as exc: # noqa: BLE001 # a speculative eligibility read must fail open on unknown faults
if isinstance(exc, ValueError) and any(marker.value in str(exc) for marker in RouterErrors):
verbose_router_logger.debug("health probe exhausted model=%s error=%s", model_name, exc)
return False
verbose_router_logger.debug(
"ComplexityRouter: eligibility probe for %s failed, treating the group as live: %s", model_name, exc
)
@ -3062,76 +3097,124 @@ class ComplexityRouter(CustomLogger):
input: str | list | None, # mutable-ok: mirrors the owner's own input parameter, which this forwards verbatim
resolved_messages: Sequence[Mapping[str, object]] | None,
request_kwargs: dict, # mutable-ok: same shape the hook receives
context_fit: _RequestContextFit | None = None,
) -> PreRoutingHookResponse:
"""Replace a decided model group that has no serving capacity with a live peer in the same tier.
Applied to the decided response at the hook's exits, so every arm that can place a request
is covered by one owner: a fresh classification, a replayed or escalated session pin, a
plan-mode floor, a context-window escalation, an adaptive pick, and whatever arm is added
next. Peers come from the DECIDED tier only; climbing to another tier is deliberately not
done here, since a higher tier costs more than the classifier asked for.
Serving capacity is one question asked of one owner (`_model_group_can_serve`), so the
substitute is only ever a group the pipeline would actually accept for this request. The
pick then runs through `_pick_model_for_tier`, so routing plugins decide the substitute
exactly as they decided the original.
Fails open everywhere it cannot be sure: an unreadable eligibility view, a decision
carrying no tier (default_model), or a tier whose every peer is unusable too. It fails
CLOSED on a plugin that empties the pool, leaving the original decision to fail rather
than serving a model the plugin excluded.
"""
"""Try compatible tier recovery before the default, preserving request policy and fit."""
decision: Final = response.routing_decision
decided_tier: Final = decision.get("tier") if decision is not None else None
if decision is None or not isinstance(decided_tier, str):
return response
peers: Final = tuple(self._tier_pools().get(decided_tier, ()))
if len(peers) < 2:
return response
if await self._model_group_can_serve(response.model, messages, input, request_kwargs):
fit: Final = context_fit or await self._request_context_fit(resolved_messages, request_kwargs)
if fit.accepts(response.model) and await self._model_group_can_serve(
response.model, messages, input, request_kwargs
):
return response
eligible: Final = (
self._modality_eligible_models()
if self.config.modality_routing and resolved_messages and request_contains_image_content(resolved_messages)
else None
)
candidates: Final = tuple(
peer for peer in peers if peer != response.model and (eligible is None or peer in eligible)
pools: Final = self._tier_pools()
context_recovery: Final = bool(decision.get("context_escalated")) or any(
not fit.accepts(model) for model in pools.get(decided_tier, ())
)
if not candidates:
return response
servable: Final = await asyncio.gather(
*(self._model_group_can_serve(peer, messages, input, request_kwargs) for peer in candidates)
modality_recovery: Final = eligible is not None
names: Final = self.config.tier_names()
tiers: Final = (
tuple(names[names.index(decided_tier) :])
if (context_recovery or modality_recovery) and decided_tier in names
else (decided_tier,)
)
live: Final = tuple(peer for peer, can_serve in zip(candidates, servable) if can_serve)
if not live:
return response
repick_messages: Final = (
list(resolved_messages) if resolved_messages else None # mutable-ok: the pick's param is list-typed
)
try:
new_model: Final = await self._pick_model_for_tier(
decided_tier if self.config.has_custom_tiers else ComplexityTier(decided_tier),
messages,
repick_messages, # pyright: ignore[reportArgumentType] # hook-resolved message dicts; the pick only reads them
request_kwargs,
allowed_models=live,
async def recover_tier(candidate_tier: str) -> PreRoutingHookResponse | None:
peers: Final = tuple(
model
for model in pools.get(candidate_tier, ())
if not context_recovery
or candidate_tier == decided_tier
or fit.needed is None
or _group_provably_fits(fit.facts.get(model, (None, True)), fit.needed, fit.buffer)
)
except ValueError as exc:
verbose_router_logger.debug(
"ComplexityRouter: health failover found no candidate the routing plugins allow: %s", exc
candidates: Final = tuple(
peer
for peer in peers
if peer != response.model and fit.accepts(peer) and (eligible is None or peer in eligible)
)
servable: Final = await asyncio.gather(
*(self._model_group_can_serve(peer, messages, input, request_kwargs) for peer in candidates)
)
live: Final = tuple(peer for peer, can_serve in zip(candidates, servable) if can_serve)
if live:
repick_messages: Final = (
list(resolved_messages) if resolved_messages else None # mutable-ok: the pick's param is list-typed
)
try:
new_model: Final = await self._pick_model_for_tier(
candidate_tier if self.config.has_custom_tiers else ComplexityTier(candidate_tier),
messages,
repick_messages, # pyright: ignore[reportArgumentType] # hook-resolved message dicts; the pick only reads them
request_kwargs,
allowed_models=live,
)
except ValueError as exc:
verbose_router_logger.debug(
"ComplexityRouter: health failover found no candidate the routing plugins allow: %s", exc
)
else:
self._restamp_adaptive_choice(request_kwargs, response.model, new_model)
verbose_router_logger.info(
"ComplexityRouter: routing decision cause=health_failover, routed_model=%s, displaced=%s",
new_model,
response.model,
)
new_decision: Final = self._build_routing_decision(
routed_model=new_model,
cause="health_failover",
tier=candidate_tier,
score=decision.get("score"),
signals=(*(decision.get("signals") or ()), f"health_displaced:{response.model}"),
matched_keyword=decision.get("matched_keyword"),
escalation_keyword=decision.get("escalation_keyword"),
escalated=bool(decision.get("escalated", False)),
classifier_model=decision.get("classifier_model"),
classifier_cost=decision.get("classifier_cost"),
conversation_continuing=bool(decision.get("conversation_continuing", True)),
tier_litellm_params=self._litellm_params_for_model(candidate_tier, new_model),
context_escalation_original_tier=decision.get("context_escalation_original_tier"),
)
return response.model_copy(
update={ # mutable-ok: model_copy types update as a plain dict
"model": new_model,
"litellm_params": self._litellm_params_for_model(candidate_tier, new_model),
"routing_decision": new_decision,
}
)
return None
for candidate_tier in tiers:
if (recovered := await recover_tier(candidate_tier)) is not None:
return recovered
default_model: Final = self.config.default_model
plan_mode_active: Final = self._matched_plan_mode_signal(request_kwargs, resolved_messages) is not None
if (
plan_mode_active
or self.config.plugins
or not default_model
or default_model == response.model
or not fit.accepts(default_model)
or (eligible is not None and default_model not in eligible)
or not await self._model_group_can_serve(default_model, messages, input, request_kwargs)
):
return response
self._restamp_adaptive_choice(request_kwargs, response.model, new_model)
self._restamp_adaptive_choice(request_kwargs, response.model, default_model)
verbose_router_logger.info(
"ComplexityRouter: routing decision cause=health_failover, routed_model=%s, displaced=%s",
new_model,
"ComplexityRouter: routing decision cause=health_default_fallback, routed_model=%s, displaced=%s",
default_model,
response.model,
)
new_decision: Final = self._build_routing_decision(
routed_model=new_model,
cause="health_failover",
tier=decision.get("tier"),
default_decision: Final = self._build_routing_decision(
routed_model=default_model,
cause="health_default_fallback",
score=decision.get("score"),
signals=(*(decision.get("signals") or ()), f"health_displaced:{response.model}"),
matched_keyword=decision.get("matched_keyword"),
@ -3140,14 +3223,14 @@ class ComplexityRouter(CustomLogger):
classifier_model=decision.get("classifier_model"),
classifier_cost=decision.get("classifier_cost"),
conversation_continuing=bool(decision.get("conversation_continuing", True)),
tier_litellm_params=self._litellm_params_for_model(decided_tier, new_model),
tier_litellm_params=self._litellm_params_for_model(None, default_model),
context_escalation_original_tier=decision.get("context_escalation_original_tier"),
)
return response.model_copy(
update={ # mutable-ok: model_copy types update as a plain dict
"model": new_model,
"litellm_params": self._litellm_params_for_model(decided_tier, new_model),
"routing_decision": new_decision,
"model": default_model,
"litellm_params": self._litellm_params_for_model(None, default_model),
"routing_decision": default_decision,
}
)
@ -3462,6 +3545,7 @@ class ComplexityRouter(CustomLogger):
# chat-completions messages, so it is real work on every non-chat surface, and
# both the conversation shape and the classifier read the same list.
resolved_messages: Final = self._resolve_messages(messages, request_kwargs)
context_fit: Final = await self._request_context_fit(resolved_messages, request_kwargs)
marker_pairs: Final = self._reminder_markers_for_request(request_kwargs)
conversation_continuing: Final = _conversation_is_continuing(resolved_messages)
@ -3512,7 +3596,11 @@ class ComplexityRouter(CustomLogger):
pin_source_tier: Final = self._tier_for_model(routed_model)
pin_placement: Final = (
await self._context_window_placement(
pin_source_tier, resolved_messages, request_kwargs, pool_override=(routed_model,)
pin_source_tier,
resolved_messages,
request_kwargs,
pool_override=(routed_model,),
context_fit=context_fit,
)
if pin_source_tier is not None
else None
@ -3582,11 +3670,13 @@ class ComplexityRouter(CustomLogger):
messages,
resolved_messages,
request_kwargs,
context_fit,
),
messages,
input,
resolved_messages,
request_kwargs,
context_fit,
)
)
@ -3598,14 +3688,18 @@ class ComplexityRouter(CustomLogger):
specific_deployment=specific_deployment,
conversation_continuing=conversation_continuing,
resolved_messages=resolved_messages,
context_fit=context_fit,
)
response: Final = (
await self._gate_response_health(
await self._gate_response_modality(routed_response, messages, resolved_messages, request_kwargs),
await self._gate_response_modality(
routed_response, messages, resolved_messages, request_kwargs, context_fit
),
messages,
input,
resolved_messages,
request_kwargs,
context_fit,
)
if routed_response is not None
else None
@ -3640,6 +3734,7 @@ class ComplexityRouter(CustomLogger):
specific_deployment: bool | None = False,
conversation_continuing: bool = True,
resolved_messages: Sequence[Mapping[str, object]] | None = None,
context_fit: _RequestContextFit | None = None,
) -> PreRoutingHookResponse | None:
"""
Classifies the request by complexity and returns the appropriate model.
@ -3811,7 +3906,9 @@ class ComplexityRouter(CustomLogger):
plan_floored: Final = tier != pre_floor_tier
if plan_floored:
signals = (*signals, "plan_mode_floor")
context_placement: Final = await self._context_window_placement(tier, resolved_messages, request_kwargs)
context_placement: Final = await self._context_window_placement(
tier, resolved_messages, request_kwargs, context_fit=context_fit
)
tier, signals, context_original_tier = _apply_context_placement(tier, signals, context_placement)
score_repr: Final = f"{score:.3f}" if score is not None else "n/a"
fallback_model: Final = self.config.default_model if not self.config.plugins else None

View file

@ -349,7 +349,9 @@ class DeploymentAffinityCheck(CustomLogger):
first write instead of the last. Re-claiming with the stored value refreshes its
TTL, the same keepalive the complexity router's model pin documents: an active
session must not lose its pin mid-conversation just because it outlives the
original write, so `session_affinity_ttl_seconds` bounds idle time, not total
original write, so the affinity TTL (the Router's
`deployment_affinity_ttl_seconds`, or a pre-routing hook's per-request
`session_affinity_ttl_seconds` override) bounds idle time, not total
session length. On Redis one Lua script does the get-or-set-or-refresh
atomically (same registration seam the rate limiters use) and the in-memory
tier is synchronized to the winner; without Redis, and whenever Redis is

View file

@ -47,6 +47,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager):
aws_web_identity_token: str | None = None,
aws_sts_endpoint: str | None = None,
replica_regions: list[str] | None = None,
kms_key_id: str | None = None,
**kwargs,
):
BaseSecretManager.__init__(self, **kwargs)
@ -61,6 +62,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager):
self.aws_web_identity_token = aws_web_identity_token
self.aws_sts_endpoint = aws_sts_endpoint
self.replica_regions: list[str] = replica_regions or []
self.kms_key_id = kms_key_id
@classmethod
def validate_environment(cls):
@ -106,7 +108,8 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager):
# Remove None values
aws_kwargs = {k: v for k, v in aws_kwargs.items() if v is not None}
litellm.secret_manager_client = cls(**aws_kwargs)
kms_key_id: Final = key_management_settings.kms_key_id if key_management_settings is not None else None
litellm.secret_manager_client = cls(kms_key_id=kms_key_id, **aws_kwargs)
litellm._key_management_system = KeyManagementSystem.AWS_SECRET_MANAGER
except Exception as e:
@ -275,6 +278,9 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager):
if description:
data["Description"] = description
if self.kms_key_id:
data["KmsKeyId"] = self.kms_key_id
# ✅ Normalize tags to AWS format
if tags:
if isinstance(tags, dict):

View file

@ -45,6 +45,9 @@ class KeyManagementSettings(LiteLLMPydanticObjectBase):
tags: dict[str, str] | None = None
"""Optional tags to attach when creating secrets (e.g. {"Environment": "Prod", "Owner": "AI-Platform"})."""
kms_key_id: str | None = None
"""Optional customer-managed KMS key (ID, alias or ARN) used to encrypt secrets created in AWS Secrets Manager."""
custom_secret_manager: str | None = None
"""
Path to custom secret manager class (e.g. "my_secret_manager.InMemorySecretManager")

View file

@ -2920,6 +2920,7 @@ RoutingDecisionCause = Literal[
# same tier served instead. The displaced group rides in signals. Reported even on a kept
# session pin, since the pinned model did not serve the request.
"health_failover",
"health_default_fallback",
"session_affinity_pin",
"session_affinity_escalation",
# classification_mode 'user_turn': the request is an agent loop's continuation turn (no new

View file

@ -5,7 +5,7 @@ from enum import Enum
from typing import Any, Literal
from pydantic import BaseModel
from typing_extensions import TypedDict
from typing_extensions import ReadOnly, TypedDict
class SupportedVectorStoreIntegrations(str, Enum):
@ -96,6 +96,17 @@ class VectorStoreSearchResponse(TypedDict, total=False):
data: list[VectorStoreSearchResult] | None
VectorStoreSearchFailureMode = Literal["annotate", "error"]
class VectorStoreSearchFailure(TypedDict):
"""A configured vector store whose search failed, as reported back to the API caller"""
vector_store_id: ReadOnly[str]
custom_llm_provider: ReadOnly[str | None]
error: ReadOnly[str]
class VectorStoreSearchOptionalRequestParams(TypedDict, total=False):
"""TypedDict for Optional parameters supported by the vector store search API."""

View file

@ -1895,3 +1895,15 @@ async def test_optional_discovery_preserves_cancellation(method: str) -> None:
task.cancel()
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(task, timeout=3)
def test_client_import_before_proxy_credentials_succeeds_in_fresh_process():
import subprocess
result = subprocess.run(
[sys.executable, "-c", "import litellm.experimental_mcp_client.client; from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager; print(MCPServerManager.__name__)"],
capture_output=True, text=True, timeout=60, check=False,
)
assert result.returncode == 0, result.stderr
assert result.stdout.strip() == "MCPServerManager"

View file

@ -11,7 +11,16 @@ from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook i
ProxyServerRuntime,
VectorStorePreCallHook,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.openai import AllMessageValues, ResponsesAPIResponse
from litellm.types.utils import (
CallTypes,
Choices,
Delta,
Message,
ModelResponse,
ModelResponseStream,
StreamingChoices,
)
from litellm.types.vector_stores import (
VectorStoreResultContent,
VectorStoreSearchResponse,
@ -36,6 +45,18 @@ def _search_response(text: str) -> VectorStoreSearchResponse:
)
def _first_message(response: ModelResponse) -> Message:
choice = response.choices[0]
assert isinstance(choice, Choices)
return choice.message
@dataclass(frozen=True)
class ExplodingRegistry:
async def pop_vector_stores_to_run_with_db_fallback(self, **kwargs: object) -> list[LiteLLM_ManagedVectorStore]:
raise RuntimeError("the registry blew up")
@dataclass
class RecordingRouter:
failing_vector_store_ids: frozenset[str] = frozenset()
@ -285,3 +306,264 @@ def test_the_default_runtime_follows_the_proxy_globals(monkeypatch: pytest.Monke
assert runtime.llm_router() is router
assert runtime.prisma_client() is prisma
@pytest.mark.asyncio
async def test_a_failing_vector_store_is_reported_back_to_the_caller(
registry_with: RegisterStores,
) -> None:
"""Regression (LIT-6809): a silently dropped store left the caller with an un-augmented answer and no signal."""
registry_with("vs-broken", "vs-healthy")
logging_obj = FakeLoggingObj({})
await _run_hook(
VectorStorePreCallHook(
proxy_runtime=FakeProxyRuntime(router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"})))
),
["vs-broken", "vs-healthy"],
logging_obj,
)
response = ModelResponse(choices=[Choices(message=Message(content="an answer"))])
await VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)).async_post_call_success_deployment_hook(
request_data={"litellm_logging_obj": logging_obj},
response=response,
call_type=CallTypes.acompletion,
)
provider_specific_fields = _first_message(response).provider_specific_fields or {}
assert provider_specific_fields["vector_store_search_failures"] == (
{
"vector_store_id": "vs-broken",
"custom_llm_provider": "bedrock",
"error": "litellm.BadRequestError: no healthy deployments for vs-broken",
},
)
assert len(provider_specific_fields["search_results"]) == 1
@pytest.mark.asyncio
async def test_a_healthy_vector_store_alone_reports_no_failures(registry_with: RegisterStores) -> None:
registry_with("vs-healthy")
logging_obj = FakeLoggingObj({})
await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=RecordingRouter())),
["vs-healthy"],
logging_obj,
)
response = ModelResponse(choices=[Choices(message=Message(content="an answer"))])
await VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)).async_post_call_success_deployment_hook(
request_data={"litellm_logging_obj": logging_obj},
response=response,
call_type=CallTypes.acompletion,
)
assert "vector_store_search_failures" not in (_first_message(response).provider_specific_fields or {})
@pytest.mark.asyncio
async def test_a_failing_vector_store_is_reported_on_the_responses_api_response(
registry_with: RegisterStores,
) -> None:
"""Regression (LIT-6809): /v1/responses answered 200 with no sign the knowledge base was missing."""
registry_with("vs-broken")
logging_obj = FakeLoggingObj({})
await _run_hook(
VectorStorePreCallHook(
proxy_runtime=FakeProxyRuntime(router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"})))
),
["vs-broken"],
logging_obj,
)
response = ResponsesAPIResponse(id="resp-lit6809", created_at=0, output=[])
await VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)).async_post_call_success_deployment_hook(
request_data={"litellm_logging_obj": logging_obj},
response=response,
call_type=CallTypes.aresponses,
)
assert response.model_dump()["vector_store_search_failures"] == [
{
"vector_store_id": "vs-broken",
"custom_llm_provider": "bedrock",
"error": "litellm.BadRequestError: no healthy deployments for vs-broken",
}
]
@pytest.mark.asyncio
async def test_a_healthy_vector_store_leaves_the_responses_api_response_alone(registry_with: RegisterStores) -> None:
registry_with("vs-healthy")
logging_obj = FakeLoggingObj({})
await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=RecordingRouter())),
["vs-healthy"],
logging_obj,
)
response = ResponsesAPIResponse(id="resp-lit6809", created_at=0, output=[])
await VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)).async_post_call_success_deployment_hook(
request_data={"litellm_logging_obj": logging_obj},
response=response,
call_type=CallTypes.aresponses,
)
assert "vector_store_search_failures" not in response.model_dump()
@pytest.mark.asyncio
async def test_a_failing_vector_store_is_reported_on_the_streaming_chunk(registry_with: RegisterStores) -> None:
registry_with("vs-broken")
logging_obj = FakeLoggingObj({})
await _run_hook(
VectorStorePreCallHook(
proxy_runtime=FakeProxyRuntime(router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"})))
),
["vs-broken"],
logging_obj,
)
chunk = ModelResponseStream(choices=[StreamingChoices(delta=Delta(content="an answer"))])
await VectorStorePreCallHook(
proxy_runtime=FakeProxyRuntime(router=None)
).async_post_call_streaming_deployment_hook(
request_data=logging_obj.model_call_details,
response_chunk=chunk,
call_type=CallTypes.acompletion,
)
assert (chunk.choices[0].delta.provider_specific_fields or {})["vector_store_search_failures"] == (
{
"vector_store_id": "vs-broken",
"custom_llm_provider": "bedrock",
"error": "litellm.BadRequestError: no healthy deployments for vs-broken",
},
)
@pytest.mark.asyncio
async def test_error_mode_fails_the_request_instead_of_answering_without_the_knowledge_base(
registry_with: RegisterStores,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Regression (LIT-6809): opting in must turn an ungrounded answer into a 400 the caller can act on."""
registry_with("vs-broken", "vs-healthy")
monkeypatch.setattr(litellm, "vector_store_search_failure_mode", "error")
with pytest.raises(litellm.VectorStoreSearchError) as raised:
await _run_hook(
VectorStorePreCallHook(
proxy_runtime=FakeProxyRuntime(
router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"}))
)
),
["vs-broken", "vs-healthy"],
FakeLoggingObj({}),
)
assert raised.value.status_code == 400
assert raised.value.failures == (
{
"vector_store_id": "vs-broken",
"custom_llm_provider": "bedrock",
"error": "litellm.BadRequestError: no healthy deployments for vs-broken",
},
)
assert "vs-broken: litellm.BadRequestError: no healthy deployments for vs-broken" in raised.value.message
@pytest.mark.asyncio
async def test_a_misspelled_failure_mode_annotates_instead_of_erroring_the_request(
registry_with: RegisterStores,
monkeypatch: pytest.MonkeyPatch,
warnings: list[logging.LogRecord],
) -> None:
"""Regression (LIT-6809): litellm_settings takes any value, so a typo must not become a 500."""
registry_with("vs-broken")
monkeypatch.setattr(litellm, "vector_store_search_failure_mode", "erorr")
logging_obj = FakeLoggingObj({})
_, messages, _ = await _run_hook(
VectorStorePreCallHook(
proxy_runtime=FakeProxyRuntime(router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"})))
),
["vs-broken"],
logging_obj,
)
assert messages[0]["content"] == "what is litellm?"
assert logging_obj.model_call_details["vector_store_search_failures"] == (
{
"vector_store_id": "vs-broken",
"custom_llm_provider": "bedrock",
"error": "litellm.BadRequestError: no healthy deployments for vs-broken",
},
)
assert any("erorr" in record.getMessage() for record in warnings)
@pytest.mark.asyncio
async def test_error_mode_leaves_a_fully_healthy_request_alone(
registry_with: RegisterStores,
monkeypatch: pytest.MonkeyPatch,
) -> None:
registry_with("vs-healthy")
monkeypatch.setattr(litellm, "vector_store_search_failure_mode", "error")
_, messages, _ = await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=RecordingRouter())),
["vs-healthy"],
FakeLoggingObj({}),
)
assert messages[0]["content"] == "Context:\n\ncontext from vs-healthy\n\n"
@pytest.mark.asyncio
async def test_error_mode_does_not_swallow_the_raise_in_the_hooks_own_catch_all(
registry_with: RegisterStores,
monkeypatch: pytest.MonkeyPatch,
warnings: list[logging.LogRecord],
) -> None:
"""Regression (LIT-6809): the catch-all around the hook must not turn the opted-in failure back into a 200."""
registry_with("vs-broken")
monkeypatch.setattr(litellm, "vector_store_search_failure_mode", "error")
with pytest.raises(litellm.VectorStoreSearchError):
await _run_hook(
VectorStorePreCallHook(
proxy_runtime=FakeProxyRuntime(
router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"}))
)
),
["vs-broken"],
FakeLoggingObj({}),
)
assert [record.levelname for record in warnings] == ["WARNING"]
@pytest.mark.asyncio
async def test_a_crash_outside_the_search_names_the_requested_vector_stores(
monkeypatch: pytest.MonkeyPatch,
warnings: list[logging.LogRecord],
) -> None:
"""Regression (LIT-6809): the catch-all logged no store id, so an operator could not tell which store broke."""
monkeypatch.setattr(litellm, "vector_store_registry", ExplodingRegistry())
_, messages, _ = await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)),
["vs-one", "vs-two"],
FakeLoggingObj({}),
)
assert messages == [{"role": "user", "content": "what is litellm?"}]
assert [record.getMessage() for record in warnings] == [
"Error in VectorStorePreCallHook for vector_store_ids=('vs-one', 'vs-two'): the registry blew up"
]

View file

@ -608,9 +608,11 @@ async def test_streaming_responses_relay_flush_reaches_the_success_callbacks_wit
)
stream = "event: response.completed\ndata: " + json.dumps(RESPONSES_COMPLETED_EVENT) + "\n\n"
await logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=[stream.encode()], provider_config=AzureAIPassthroughConfig()
collector = AzureAIPassthroughConfig().create_stream_collector(
model="gpt-5.4-mini", custom_llm_provider="azure_ai", endpoint="gpt/openai/responses"
)
collector.add(stream.encode())
await logging_obj.async_flush_passthrough_collected_chunks(collector=collector)
info = litellm.get_model_info("azure_ai/gpt-5.4-mini")
assert probe.logged_call_type == "allm_passthrough_route"

View file

@ -1,7 +1,19 @@
import base64
import json
import struct
import tracemalloc
from binascii import crc32
from datetime import datetime
from unittest.mock import patch
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector
from litellm.llms.bedrock.passthrough.transformation import BedrockPassthroughConfig
from litellm.types.utils import ModelResponse
CONVERSE_MODEL = "anthropic.claude-sonnet-4-5-20250929-v1:0"
CONVERSE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/converse-stream"
INVOKE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/invoke-with-response-stream"
def test_bedrock_passthrough_get_complete_url_default_endpoint():
@ -500,3 +512,186 @@ def test_bedrock_passthrough_model_id_without_arn():
f"https://bedrock-runtime.us-east-1.amazonaws.com/model/{model_id}/converse"
)
assert url_str == expected_url
def _event_frame(event_type: str, payload: dict) -> bytes:
def header(name: str, value: str) -> bytes:
name_b, value_b = name.encode(), value.encode()
return struct.pack("!B", len(name_b)) + name_b + struct.pack("!B", 7) + struct.pack("!H", len(value_b)) + value_b
payload_b = json.dumps(payload, separators=(",", ":")).encode()
headers_b = (
header(":event-type", event_type)
+ header(":content-type", "application/json")
+ header(":message-type", "event")
)
prelude = struct.pack("!II", 12 + len(headers_b) + len(payload_b) + 4, len(headers_b))
prelude_crc = crc32(prelude) & 0xFFFFFFFF
message = struct.pack("!I", prelude_crc) + headers_b + payload_b
return prelude + message + struct.pack("!I", crc32(message, prelude_crc) & 0xFFFFFFFF)
def _text_block(index: int, texts: list[str]) -> bytes:
return (
_event_frame("contentBlockStart", {"contentBlockIndex": index, "start": {}})
+ b"".join(
_event_frame("contentBlockDelta", {"contentBlockIndex": index, "delta": {"text": text}}) for text in texts
)
+ _event_frame("contentBlockStop", {"contentBlockIndex": index})
)
def _stream_tail(stop_reason: str, output_tokens: int) -> bytes:
return _event_frame("messageStop", {"stopReason": stop_reason}) + _event_frame(
"metadata",
{
"metrics": {"latencyMs": 1234},
"usage": {"inputTokens": 25, "outputTokens": output_tokens, "totalTokens": 25 + output_tokens},
},
)
def _invoke_chunk(payload: dict) -> bytes:
return _event_frame("chunk", {"bytes": base64.b64encode(json.dumps(payload).encode()).decode()})
def _stream_logging_obj(endpoint: str) -> Logging:
logging_obj = Logging(
model=CONVERSE_MODEL,
messages=[],
stream=True,
call_type="pass_through_endpoint",
start_time=datetime.now(),
litellm_call_id="call-1",
function_id="fn-1",
)
logging_obj.model_call_details["custom_llm_provider"] = "bedrock"
logging_obj.model_call_details["endpoint"] = endpoint
return logging_obj
def _converse_stream_logging_obj() -> Logging:
return _stream_logging_obj(CONVERSE_STREAM_ENDPOINT)
def _stream_collector(endpoint: str) -> PassthroughStreamCollector:
return BedrockPassthroughConfig().create_stream_collector(
model=CONVERSE_MODEL, custom_llm_provider="bedrock", endpoint=endpoint
)
def _converse_stream_collector() -> PassthroughStreamCollector:
return _stream_collector(CONVERSE_STREAM_ENDPOINT)
def _feed(collector: PassthroughStreamCollector, stream: bytes, chunk_size: int = 16384) -> None:
for offset in range(0, len(stream), chunk_size):
collector.add(stream[offset : offset + chunk_size])
def test_converse_stream_collector_keeps_usage_without_retaining_the_stream():
texts = [f"tok{i} " for i in range(4000)]
stream = _event_frame("messageStart", {"role": "assistant"}) + _text_block(0, texts) + _stream_tail("end_turn", 4000)
_feed(_converse_stream_collector(), stream)
tracemalloc.start()
try:
base = tracemalloc.get_traced_memory()[0]
collector = _converse_stream_collector()
_feed(collector, stream)
retained = tracemalloc.get_traced_memory()[0] - base
finally:
tracemalloc.stop()
assert retained < len(stream) // 4
response = collector.build_logged_response(_converse_stream_logging_obj())
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content == "".join(texts)
assert response.choices[0].finish_reason == "stop"
assert (response.usage.prompt_tokens, response.usage.completion_tokens) == (25, 4000)
def test_converse_stream_collector_keeps_tool_calls_between_text_runs():
stream = (
_event_frame("messageStart", {"role": "assistant"})
+ _text_block(0, ["Let me ", "check."])
+ _event_frame(
"contentBlockStart",
{"contentBlockIndex": 1, "start": {"toolUse": {"toolUseId": "tool-1", "name": "get_weather"}}},
)
+ _event_frame("contentBlockDelta", {"contentBlockIndex": 1, "delta": {"toolUse": {"input": '{"city": '}}})
+ _event_frame("contentBlockDelta", {"contentBlockIndex": 1, "delta": {"toolUse": {"input": '"Paris"}'}}})
+ _event_frame("contentBlockStop", {"contentBlockIndex": 1})
+ _text_block(2, ["Done", "."])
+ _stream_tail("tool_use", 12)
)
collector = _converse_stream_collector()
_feed(collector, stream, chunk_size=7)
response = collector.build_logged_response(_converse_stream_logging_obj())
assert isinstance(response, ModelResponse)
message = response.choices[0].message
assert message.content == "Let me check.Done."
assert [(call.function.name, call.function.arguments) for call in message.tool_calls] == [
("get_weather", '{"city": "Paris"}')
]
assert response.choices[0].finish_reason == "tool_calls"
assert (response.usage.prompt_tokens, response.usage.completion_tokens) == (25, 12)
def test_invoke_stream_collector_keeps_usage_without_retaining_the_stream():
texts = [f"tok{i} " for i in range(4000)]
stream = (
_invoke_chunk(
{
"type": "message_start",
"message": {
"id": "msg-1",
"type": "message",
"role": "assistant",
"model": CONVERSE_MODEL,
"content": [],
"stop_reason": None,
"usage": {"input_tokens": 25, "output_tokens": 1},
},
}
)
+ _invoke_chunk({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}})
+ b"".join(
_invoke_chunk({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}})
for text in texts
)
+ _invoke_chunk({"type": "content_block_stop", "index": 0})
+ _invoke_chunk(
{"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 4000}}
)
+ _invoke_chunk({"type": "message_stop"})
)
_feed(_stream_collector(INVOKE_STREAM_ENDPOINT), stream)
tracemalloc.start()
try:
base = tracemalloc.get_traced_memory()[0]
collector = _stream_collector(INVOKE_STREAM_ENDPOINT)
_feed(collector, stream)
retained = tracemalloc.get_traced_memory()[0] - base
finally:
tracemalloc.stop()
assert retained < len(stream) // 4
response = collector.build_logged_response(_stream_logging_obj(INVOKE_STREAM_ENDPOINT))
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content == "".join(texts)
assert response.choices[0].finish_reason == "stop"
assert (response.usage.prompt_tokens, response.usage.completion_tokens) == (25, 4000)
def test_stream_collector_logs_nothing_for_an_unrecognized_endpoint():
collector = BedrockPassthroughConfig().create_stream_collector(
model=CONVERSE_MODEL, custom_llm_provider="bedrock", endpoint=f"/model/{CONVERSE_MODEL}/rerank"
)
collector.add(_event_frame("messageStart", {"role": "assistant"}))
assert collector.build_logged_response(_converse_stream_logging_obj()) is None

View file

@ -1900,3 +1900,198 @@ class TestOCIImageUrlTransformation:
adapt_messages_to_generic_oci_standard(messages)
assert "image_url" in str(exc_info.value)
import itertools
from unittest.mock import patch
from litellm.llms.oci.chat.transformation import OCIStreamWrapper, _iter_sse_events
_STREAM_GENERIC_MODEL = "xai.grok-4"
_STREAM_COHERE_MODEL = "cohere.command-latest"
_GENERIC_TEXT_EVENT = (
'data: {{"index":0,"message":{{"role":"ASSISTANT","content":[{{"type":"TEXT","text":"{text}"}}]}},"pad":"aaa"}}'
)
_GENERIC_TERMINAL_EVENT = (
'data: {"message":{"role":"ASSISTANT","content":[{"type":"TEXT","text":""}]},"finishReason":"stop","pad":"a"}'
)
_COHERE_TEXT_EVENT = 'data: {{"apiFormat":"COHERE","text":"{text}","pad":"aaaaaa"}}'
_COHERE_TERMINAL_EVENT = (
'data: {"apiFormat":"COHERE","text":"123","finishReason":"COMPLETE",'
'"chatHistory":[{"role":"USER","message":"count"},{"role":"CHATBOT","message":"123"}]}'
)
def _make_stream_wrapper(model: str) -> OCIStreamWrapper:
logging_obj = MagicMock()
logging_obj.model_call_details = {"custom_llm_provider": "oci", "litellm_params": {}}
return OCIStreamWrapper(
completion_stream=iter([]),
model=model,
custom_llm_provider="oci",
logging_obj=logging_obj,
)
def _ticking_clock():
"""A ``time.time`` stand-in that advances a full second on every call.
Without it the whole test runs inside one wall-clock second, so a per-chunk
``created`` would coincidentally match and the drift would go unnoticed.
"""
return itertools.count(1_700_000_000.0)
class TestOCIStreamWrapperIdentityPinning:
"""One OCI streaming completion must present one id, one created and the
wrapper's model on every chunk, the way every other provider does."""
def test_generic_stream_shares_one_id_created_and_model(self):
wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL)
events = [
_GENERIC_TEXT_EVENT.format(text="1"),
_GENERIC_TEXT_EVENT.format(text="2"),
_GENERIC_TEXT_EVENT.format(text="3"),
_GENERIC_TERMINAL_EVENT,
]
with patch("time.time", side_effect=_ticking_clock()):
chunks = [wrapper.chunk_creator(event) for event in events]
assert len(chunks) == 4
assert len({chunk.id for chunk in chunks}) == 1
assert chunks[0].id.startswith("chatcmpl-")
assert len({chunk.created for chunk in chunks}) == 1
assert {chunk.model for chunk in chunks} == {_STREAM_GENERIC_MODEL}
assert [chunk.choices[0].delta.content for chunk in chunks[:3]] == ["1", "2", "3"]
assert chunks[-1].choices[0].finish_reason == "stop"
assert all(chunk._hidden_params["custom_llm_provider"] == "oci" for chunk in chunks)
def test_cohere_stream_shares_one_id_created_and_model(self):
"""Rebuilding each chunk through the shared creator must not disturb the
Cohere bookkeeping that suppresses the terminal event's repeated text."""
wrapper = _make_stream_wrapper(_STREAM_COHERE_MODEL)
events = [
_COHERE_TEXT_EVENT.format(text="1"),
_COHERE_TEXT_EVENT.format(text="2"),
_COHERE_TEXT_EVENT.format(text="3"),
_COHERE_TERMINAL_EVENT,
]
with patch("time.time", side_effect=_ticking_clock()):
chunks = [wrapper.chunk_creator(event) for event in events]
assert len(chunks) == 4
assert len({chunk.id for chunk in chunks}) == 1
assert len({chunk.created for chunk in chunks}) == 1
assert {chunk.model for chunk in chunks} == {_STREAM_COHERE_MODEL}
assert [chunk.choices[0].delta.content for chunk in chunks[:3]] == ["1", "2", "3"]
assert chunks[-1].choices[0].finish_reason == "stop"
assert chunks[-1].choices[0].delta.content is None
assert wrapper._cohere_text_emitted is True
def test_id_is_pinned_to_the_wrapper_response_id(self):
wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL)
first = wrapper.chunk_creator(_GENERIC_TEXT_EVENT.format(text="1"))
assert wrapper.response_id == first.id
assert wrapper.created == first.created
class TestOCIStreamWrapperDoneSentinel:
"""OCI's GENERIC apiFormat closes the stream with a literal `[DONE]` line;
parsing it as JSON turned every streaming completion into a 500."""
@pytest.mark.parametrize("done_event", ["data: [DONE]", "data:[DONE]", "data: [DONE] "])
def test_done_sentinel_returns_none(self, done_event):
wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL)
assert wrapper.chunk_creator(done_event) is None
def test_done_sentinel_off_the_sse_splitter_is_skipped(self):
wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL)
wire = (
f"{_GENERIC_TEXT_EVENT.format(text='1')}\n\n"
f"{_GENERIC_TEXT_EVENT.format(text='2')}\n\n"
f"{_GENERIC_TERMINAL_EVENT}\n\n"
"data: [DONE]\n\n"
)
events = list(_iter_sse_events(iter([wire])))
assert events[-1] == "data: [DONE]"
chunks = [wrapper.chunk_creator(event) for event in events]
assert chunks[-1] is None
emitted = [chunk for chunk in chunks if chunk is not None]
assert len(emitted) == 3
assert len({chunk.id for chunk in emitted}) == 1
def test_unparseable_payload_still_raises_oci_error(self):
from litellm.llms.oci.common_utils import OCIError
wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL)
with pytest.raises(OCIError, match="Chunk cannot be parsed as JSON"):
wrapper.chunk_creator("data: not-json-at-all")
def test_done_lookalike_payload_still_raises_oci_error(self):
from litellm.llms.oci.common_utils import OCIError
wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL)
with pytest.raises(OCIError, match="Chunk cannot be parsed as JSON"):
wrapper.chunk_creator("data: [DONE] trailing garbage")
_GENERIC_TOOL_CALL_EVENT = (
'data: {"index":0,"message":{"role":"ASSISTANT","content":[],'
'"toolCalls":[{"type":"FUNCTION","id":"call_1","name":"get_weather","arguments":"{}"}]}}'
)
_GENERIC_TOOL_TERMINAL_EVENT = (
'data: {"index":0,"message":{"role":"ASSISTANT","content":[]},"finishReason":"TOOL_CALLS"}'
)
def _drain_stream(model: str, events: list[str]) -> list:
logging_obj = MagicMock()
logging_obj.model_call_details = {"custom_llm_provider": "oci", "litellm_params": {}}
wrapper = OCIStreamWrapper(
completion_stream=iter(events),
model=model,
custom_llm_provider="oci",
logging_obj=logging_obj,
)
return list(wrapper)
class TestOCIStreamWrapperTerminalChunk:
"""OCI's ``chunk_creator`` override bypasses the shared handler's
finish-reason bookkeeping, so the shared end-of-stream finalizer used to
append a synthetic ``stop`` chunk after OCI's own terminal chunk, silently
downgrading a ``tool_calls`` completion for any client that reads the
finish reason off the last chunk."""
def test_generic_tool_call_stream_ends_on_tool_calls(self):
chunks = _drain_stream(
_STREAM_GENERIC_MODEL,
[_GENERIC_TOOL_CALL_EVENT, _GENERIC_TOOL_TERMINAL_EVENT, "data: [DONE]"],
)
assert [chunk.choices[0].finish_reason for chunk in chunks] == [None, "tool_calls"]
assert len({chunk.id for chunk in chunks}) == 1
def test_generic_text_stream_emits_exactly_one_finish_reason(self):
chunks = _drain_stream(
_STREAM_GENERIC_MODEL,
[_GENERIC_TEXT_EVENT.format(text="1"), _GENERIC_TERMINAL_EVENT, "data: [DONE]"],
)
assert [chunk.choices[0].finish_reason for chunk in chunks] == [None, "stop"]
def test_cohere_stream_emits_exactly_one_finish_reason(self):
chunks = _drain_stream(
_STREAM_COHERE_MODEL,
[_COHERE_TEXT_EVENT.format(text="123"), _COHERE_TERMINAL_EVENT],
)
assert [chunk.choices[0].finish_reason for chunk in chunks] == [None, "stop"]

View file

@ -34,6 +34,34 @@ class _ImmediateExecutor:
fn(*args, **kwargs)
class _RecordingCollector:
def __init__(self) -> None:
self.chunks: List[bytes] = []
def add(self, chunk: bytes) -> None:
self.chunks.append(chunk)
def build_logged_response(self, litellm_logging_obj: MagicMock) -> bytes:
return b"".join(self.chunks)
class _FailingCollector(_RecordingCollector):
def add(self, chunk: bytes) -> None:
raise ValueError("bad frame")
def _provider_config(collector: _RecordingCollector) -> MagicMock:
provider_config = MagicMock()
provider_config.create_stream_collector.return_value = collector
return provider_config
def _spend_payload(flush_mock: MagicMock) -> bytes:
flush_mock.assert_called_once()
collector = flush_mock.call_args.kwargs["collector"]
return collector.build_logged_response(litellm_logging_obj=MagicMock())
@pytest.mark.asyncio
async def test_asyncpassthroughstreamingresponse_flushes_on_normal_completion():
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
@ -48,13 +76,12 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_normal_completion():
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
received = []
received_response = AsyncPassthroughStreamingResponse(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
provider_config=_provider_config(_RecordingCollector()),
)
async for chunk in received_response:
@ -67,12 +94,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_normal_completion():
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == chunks
assert call_kwargs["provider_config"] is provider_config
assert _spend_payload(mock_logging_obj.async_flush_passthrough_collected_chunks) == b"".join(chunks)
@pytest.mark.asyncio
@ -93,12 +115,11 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_client_disconnect():
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
gen = AsyncPassthroughStreamingResponse(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
provider_config=_provider_config(_RecordingCollector()),
)
received = [await gen.__anext__()]
@ -108,11 +129,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_client_disconnect():
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == [chunks[0]]
assert _spend_payload(mock_logging_obj.async_flush_passthrough_collected_chunks) == chunks[0]
@pytest.mark.asyncio
@ -178,14 +195,13 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_upstream_exception_w
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
received = []
async def _drain():
async for chunk in AsyncPassthroughStreamingResponse(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
provider_config=_provider_config(_RecordingCollector()),
):
received.append(chunk)
@ -196,11 +212,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_upstream_exception_w
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == partial_chunks
assert _spend_payload(mock_logging_obj.async_flush_passthrough_collected_chunks) == b"".join(partial_chunks)
def test_passthroughstreamingresponse_flushes_on_normal_completion():
@ -221,12 +233,11 @@ def test_passthroughstreamingresponse_flushes_on_normal_completion():
mock_logging_obj = MagicMock()
mock_logging_obj.flush_passthrough_collected_chunks = MagicMock()
provider_config = MagicMock()
received_responce = PassthroughStreamingResponse(
response=mock_response,
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
provider_config=_provider_config(_RecordingCollector()),
)
with patch("litellm.utils.executor", _ImmediateExecutor()):
@ -237,7 +248,7 @@ def test_passthroughstreamingresponse_flushes_on_normal_completion():
assert received_responce.headers["content-type"] == "application/octet-stream"
assert received_responce.headers["x-request-id"] == "req-123"
mock_logging_obj.flush_passthrough_collected_chunks.assert_called_once()
assert _spend_payload(mock_logging_obj.flush_passthrough_collected_chunks) == b"".join(chunks)
def test_passthroughstreamingresponse_flushes_on_early_close():
@ -258,19 +269,66 @@ def test_passthroughstreamingresponse_flushes_on_early_close():
mock_logging_obj = MagicMock()
mock_logging_obj.flush_passthrough_collected_chunks = MagicMock()
provider_config = MagicMock()
with patch("litellm.utils.executor", _ImmediateExecutor()):
gen = PassthroughStreamingResponse(
response=mock_response,
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
provider_config=_provider_config(_RecordingCollector()),
)
first = next(gen)
gen.close()
assert first == chunks[0]
mock_logging_obj.flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = mock_logging_obj.flush_passthrough_collected_chunks.call_args.kwargs
assert call_kwargs["raw_bytes"] == [chunks[0]]
assert _spend_payload(mock_logging_obj.flush_passthrough_collected_chunks) == chunks[0]
@pytest.mark.asyncio
async def test_asyncpassthroughstreamingresponse_relays_the_stream_when_spend_parsing_fails():
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
chunks = [b"chunk-1", b"chunk-2", b"chunk-3"]
mock_response = _make_streaming_response(chunks)
async def response_coro():
return mock_response
mock_logging_obj = _make_logging_obj()
received = [
chunk
async for chunk in AsyncPassthroughStreamingResponse(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=_provider_config(_FailingCollector()),
)
]
await asyncio.sleep(0)
assert received == chunks
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_not_called()
def test_passthroughstreamingresponse_relays_the_stream_when_spend_parsing_fails():
from litellm.passthrough.main import PassthroughStreamingResponse
chunks = [b"a", b"b", b"c"]
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = httpx.Headers({"content-type": "application/octet-stream"})
mock_response.iter_bytes = lambda: iter(chunks)
mock_logging_obj = MagicMock()
mock_logging_obj.flush_passthrough_collected_chunks = MagicMock()
received = list(
PassthroughStreamingResponse(
response=mock_response,
litellm_logging_obj=mock_logging_obj,
provider_config=_provider_config(_FailingCollector()),
)
)
assert received == chunks
mock_logging_obj.flush_passthrough_collected_chunks.assert_not_called()

View file

@ -1,6 +1,9 @@
"""Classification matrix for upstream OAuth/DCR rejections: who is blamed depends only on the §5.2
code and whose credentials the gateway presented, never on the upstream's HTTP status."""
from typing import Final
import pytest
import httpx
from litellm.proxy._experimental.mcp_server.faults.classify import (
@ -12,6 +15,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import (
GatewayRejected,
UpstreamProtocolFault,
UpstreamReportedFault,
UpstreamRegistrationRefused,
)
@ -145,3 +149,28 @@ def test_dcr_server_error_code_is_not_blamed_on_caller():
log_context="srv",
)
assert isinstance(fault, UpstreamReportedFault)
@pytest.mark.parametrize("status_code", [401, 403])
@pytest.mark.parametrize("body", ["Forbidden", '<html>private upstream details</html>', '{"error": ""}', '{"error": 12}'])
def test_dcr_access_refusal_without_oauth_error(status_code: int, body: str) -> None:
fault: Final = classify_upstream_dcr_rejection(_response(status_code, text_body=body), log_context="srv")
assert isinstance(fault, UpstreamRegistrationRefused)
assert fault.status_code == status_code
@pytest.mark.parametrize("status_code", [401, 403])
def test_dcr_access_refusal_preserves_oauth_error(status_code: int) -> None:
fault: Final = classify_upstream_dcr_rejection(
_response(status_code, json_body={"error": "invalid_redirect_uri", "error_description": "not allowed"}),
log_context="srv",
)
assert fault == CallerRejected(code="invalid_redirect_uri", description="not allowed")
@pytest.mark.parametrize("status_code", [401, 403])
def test_token_access_refusal_remains_protocol_fault(status_code: int) -> None:
fault: Final = classify_upstream_token_rejection(
_response(status_code, text_body="Forbidden"), credential_source="gateway_stored", log_context="srv"
)
assert isinstance(fault, UpstreamProtocolFault)

View file

@ -2,6 +2,9 @@
code can never ship on a server-fault status and gateway-side faults never carry provider prose."""
import json
from typing import Final, Literal
import pytest
from litellm.proxy._experimental.mcp_server.faults.render_oauth import (
dcr_fault_detail,
@ -12,6 +15,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import (
GatewayRejected,
UpstreamProtocolFault,
UpstreamReportedFault,
UpstreamRegistrationRefused,
)
@ -94,3 +98,17 @@ def test_dcr_upstream_reported_fault_maps_to_5xx():
status_code, detail = dcr_fault_detail(UpstreamReportedFault(code="server_error"))
assert status_code == 502
assert "internal error" in detail
@pytest.mark.parametrize("upstream_status", [401, 403])
def test_registration_refusal_gives_configuration_guidance(upstream_status: Literal[401, 403]) -> None:
fault: Final = UpstreamRegistrationRefused(status_code=upstream_status)
status, detail = dcr_fault_detail(fault)
assert status == 403
assert f"HTTP {upstream_status}" in detail
assert "may require a pre-registered OAuth client" in detail
assert "client_id" in detail and "client_secret" in detail
response: Final = render_token_fault(fault)
assert response.status_code == 400
assert json.loads(response.body) == {"error": "unauthorized_client", "error_description": detail}
assert response.headers["cache-control"] == "no-store"

View file

@ -478,3 +478,47 @@ async def test_bearer_auth_advertises_the_header_it_will_occupy():
assert ClientCredentialsBearerAuth("t", refetch, ClientCredentialsConfig()).header_name == "Authorization"
default_carrier = ClientCredentialsConfig(header_name="esb-oauth")
assert ClientCredentialsBearerAuth("t", refetch, default_carrier).header_name == "esb-oauth"
@pytest.mark.asyncio
@pytest.mark.parametrize("mode", ["denied", "invalid", "missing", "success", "timeout", "connect", "cancel"])
async def test_token_exchange_failure_diagnostics(mode, monkeypatch, caplog):
import asyncio
import logging
from litellm.llms.custom_httpx import http_handler
from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import post_client_credentials_grant
class Poster:
async def post(self, url, headers, data):
request = httpx.Request("POST", url, headers=headers, data=data)
if mode == "timeout":
raise httpx.ReadTimeout("private-transport-message", request=request)
if mode == "connect":
raise httpx.ConnectError("private-transport-message", request=request)
if mode == "cancel":
raise asyncio.CancelledError
response = httpx.Response(401 if mode == "denied" else 200, request=request,
content=b"not-json-private" if mode == "invalid" else None,
json=None if mode == "invalid" else {"error": "invalid_client", "client_secret":"first second", **({"access_token":"private-token"} if mode == "success" else {})})
response.raise_for_status()
return response
monkeypatch.setattr(http_handler, "get_async_httpx_client", lambda **kwargs: Poster())
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
if mode == "cancel":
with pytest.raises(asyncio.CancelledError):
await post_client_credentials_grant("https://idp/token", {}, {})
assert not caplog.text
return
result = await post_client_credentials_grant("https://idp/token?key=query-secret", {"client_secret":"first second"}, {"X-Custom":"header-secret"})
for secret in ("first", "second", "query-secret", "header-secret", "private-token", "not-json-private", "private-transport-message"):
assert secret not in caplog.text
if mode == "success":
assert isinstance(result, TokenEndpointSuccess) and result.body["access_token"] == "private-token"
assert not caplog.text
elif mode in {"timeout", "connect"}:
assert isinstance(result, TokenEndpointUnreachable)
assert "POST https://idp/ failed" in caplog.text
else:
assert "POST https://idp/ -> HTTP" in caplog.text
assert {"denied":"denied", "invalid":"invalid response", "missing":"no access token"}[mode] in caplog.text

View file

@ -7150,7 +7150,7 @@ async def test_extract_user_id_rehydrates_cross_replica_dict_cache(proxy_globals
key = "sk-alice-key"
cache = UserApiKeyCache()
cache.in_memory_cache.set_cache(hash_token(key), {"token": hash_token(key), "user_id": "alice"})
cache.set_cache(hash_token(key), {"token": hash_token(key), "user_id": "alice"})
proxy_globals.user_api_key_cache = cache
proxy_globals.prisma_client = object()
@ -11141,3 +11141,48 @@ async def test_enforced_login_warms_verified_token_readable_without_database_loo
assert token.identity_binding_proof == proof
assert token.refresh_token is None
read.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("upstream_status", [401, 403])
@pytest.mark.parametrize("auth_type", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate])
@pytest.mark.parametrize("dcr_bridge", [False, True])
@pytest.mark.parametrize("flow", ["register", "mint"])
async def test_dcr_refusal_is_actionable_without_upstream_body(
upstream_status: int, auth_type: MCPAuth, dcr_bridge: bool, flow: str, monkeypatch: pytest.MonkeyPatch
) -> None:
import httpx
from typing import Final
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
mint_ephemeral_dcr_client,
register_client_with_server,
)
server: Final = _bridge_server(
auth_type=auth_type, dcr_bridge=dcr_bridge, server_id=f"refused-{auth_type}-{dcr_bridge}-{flow}-{upstream_status}",
client_id=None,
)
import respx
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
with respx.mock as upstream:
registration: Final = upstream.post(server.registration_url).mock(
return_value=httpx.Response(upstream_status, text="Forbidden private upstream details")
)
operation: Final = (
mint_ephemeral_dcr_client(_bridge_mock_request(), server)
if flow == "mint"
else register_client_with_server(
request=_bridge_mock_request(), mcp_server=server, client_name="Test client",
grant_types=None, response_types=None, token_endpoint_auth_method=None,
client_redirect_uris=["http://localhost:9999/callback"],
)
)
with pytest.raises(HTTPException) as exc:
await operation
assert registration.call_count == 1
assert exc.value.status_code == 403
assert f"HTTP {upstream_status}" in str(exc.value.detail)
assert "pre-registered OAuth client" in str(exc.value.detail)
assert "private upstream details" not in str(exc.value.detail)

View file

@ -10,9 +10,13 @@ from starlette.types import Message
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import AuthResolution
import httpx
from litellm.proxy._experimental.mcp_server.mcp_debug import (
MCP_DEBUG_REQUEST_HEADER,
MCPDebug,
describe_upstream_http_failure,
MCPAuthDiagnostics,
)
@ -206,6 +210,206 @@ class TestWrapSendWithDebugHeaders:
assert captured[0] == body_msg
class TestDescribeUpstreamHttpFailure:
@staticmethod
def _status_error(*, body: bytes, response_body: bytes | None = None) -> httpx.HTTPStatusError:
request = httpx.Request(
"POST",
"https://upstream.example/apis/mcp",
headers={"Authorization": "Bearer secret-token-abcdef0123456789", "Content-Type": "application/json" if body.startswith(b"{") else "application/x-www-form-urlencoded"},
content=body,
)
response = (
httpx.Response(500, request=request, content=response_body)
if response_body is not None
else httpx.Response(500, request=request, stream=httpx.ByteStream(b'{"error":"boom"}'))
)
return httpx.HTTPStatusError("500", request=request, response=response)
def test_includes_method_url_status_and_request_body(self):
exc = self._status_error(
body=b'{"method":"initialize","jsonrpc":"2.0","id":0}',
response_body=b'{"error":"boom"}',
)
described = describe_upstream_http_failure(exc)
assert described is not None
assert "POST https://upstream.example/ -> HTTP 500" in described
assert '{"method":"initialize"' in described
assert 'response body: {"error":"boom"}' in described
def test_masks_authorization_header_and_secret_body_fields(self):
exc = self._status_error(
body=b"grant_type=client_credentials&client_id=abc&client_secret=super-secret-value-1234",
response_body=b"{}",
)
described = describe_upstream_http_failure(exc)
assert described is not None
assert "secret-token-abcdef0123456789" not in described
assert "super-secret-value-1234" not in described
assert "client_id=abc" in described
assert "client_secret=" in described
def test_reports_unread_streamed_response_body(self):
described = describe_upstream_http_failure(self._status_error(body=b"{}"))
assert described is not None
assert "response body: (not read)" in described
def test_finds_response_behind_cause_chain(self):
wrapper = RuntimeError("token minting failed")
wrapper.__cause__ = self._status_error(body=b"{}", response_body=b'{"error":"invalid_client"}')
described = describe_upstream_http_failure(wrapper)
assert described is not None
assert "invalid_client" in described
def test_returns_none_without_http_response(self):
assert describe_upstream_http_failure(ConnectionError("refused")) is None
@pytest.mark.parametrize("body", [
b'{"password":"first second","token":"demo-secret"}',
b'{"nested":[{"access_token":"first,second"}]}',
b'client%5Fsecret=first+second&token=demo-secret',
])
def test_failure_log_fully_redacts_structured_secrets(body):
request = httpx.Request("POST", "https://upstream/mcp?credential=query-secret",
headers={"X-Custom-Credential": "custom-secret"}, content=body)
response = httpx.Response(500, request=request, content=body)
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
assert detail is not None
for secret in ("first", "second", "demo-secret", "custom-secret", "query-secret"):
assert secret not in detail
def test_failure_log_omits_unstructured_body():
request = httpx.Request("POST", "https://upstream/mcp", content=b"arbitrary-secret")
response = httpx.Response(500, request=request, content=b"<html>arbitrary-secret</html>")
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
assert detail is not None
assert "arbitrary-secret" not in detail
assert "omitted" in detail
@pytest.mark.asyncio
@pytest.mark.parametrize("mode", ["error", "empty", "large", "timeout", "read_failure", "closed", "success", "cancel"])
async def test_error_capture_is_bounded_and_preserves_success_and_cancellation(mode):
from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response
class Stream(httpx.AsyncByteStream):
def __init__(self):
self.reads = 0
async def __aiter__(self):
self.reads += 1
if mode == "timeout":
await asyncio.sleep(10)
if mode == "closed":
raise httpx.StreamClosed()
if mode == "read_failure":
raise httpx.ReadError("private-read-error")
if mode == "cancel":
raise asyncio.CancelledError
yield b"" if mode == "empty" else b'{"error":"missing_scope","password":"first second"}' if mode != "large" else b"x" * 20000
stream = Stream()
request = httpx.Request("POST", "https://upstream/mcp")
response = httpx.Response(200 if mode == "success" else 500, request=request, stream=stream)
if mode == "cancel":
with pytest.raises(asyncio.CancelledError):
await capture_upstream_error_response(response)
return
await capture_upstream_error_response(response)
if mode == "success":
assert stream.reads == 0
assert await response.aread() == b'{"error":"missing_scope","password":"first second"}'
return
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
assert detail is not None
assert "first" not in detail and "second" not in detail and "private-read-error" not in detail
expected = {"empty": "(empty)", "error": "missing_scope", "large": "capture limit", "timeout": "read failed", "read_failure": "read failed", "closed":"read failed"}
assert expected[mode] in detail
if mode == "error":
assert await response.aread() == b'{"error":"missing_scope","password":"first second"}'
@pytest.mark.parametrize("body", [b"", b'"scalar"', b'{"hint":"line1\\nline2"}', b'{"hint":"' + b'x' * 600 + b'"}'])
def test_failure_preview_handles_empty_scalar_control_and_long_bodies(body):
request = httpx.Request("POST", "https://user:secret@upstream/mcp?key=private#private", content=body)
response = httpx.Response(500, request=request, content=body)
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
assert detail is not None
assert "private" not in detail and "user:secret" not in detail and "\n" not in detail
if not body:
assert "(empty)" in detail
elif body.startswith(b'"'):
assert "omitted" in detail
elif len(body) > 512:
assert "truncated" in detail and len(detail) < 1300
else:
assert "line1\\nline2" in detail
@pytest.mark.asyncio
@pytest.mark.parametrize("slow_error", [False, True])
async def test_error_capture_preserves_httpx_auth_retry(slow_error):
from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response
class RetryAuth(httpx.Auth):
def auth_flow(self, request):
response = yield request
if response.status_code == 401:
request.headers["Authorization"] = "Bearer refreshed"
yield request
class SlowStream(httpx.AsyncByteStream):
async def __aiter__(self):
await asyncio.sleep(10)
yield b'{"error":"expired_token"}'
def upstream(request):
if request.headers.get("Authorization"):
return httpx.Response(200, json={"ok": True})
return httpx.Response(401, stream=SlowStream()) if slow_error else httpx.Response(401, json={"error":"expired_token"})
async with httpx.AsyncClient(transport=httpx.MockTransport(upstream), auth=RetryAuth(),
event_hooks={"response":[capture_upstream_error_response]}) as client:
response = await client.get("https://upstream/mcp")
assert response.status_code == 200 and response.json() == {"ok":True}
if slow_error:
assert response.history[0].content == b""
else:
assert response.history[0].json() == {"error":"expired_token"}
def test_failure_diagnostics_without_request_and_with_streamed_request():
response = httpx.Response(503)
exc = httpx.HTTPStatusError("failed", request=httpx.Request("GET", "https://upstream"), response=response)
assert describe_upstream_http_failure(exc) == "HTTP 503 | request unavailable"
request = httpx.Request("POST", "https://upstream", content=iter((b"private-body",)))
response = httpx.Response(503, request=request)
described = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response))
assert described is not None and "streamed, not captured" in described and "private-body" not in described
def test_deep_error_body_is_bounded_without_exposing_nested_values():
body = b'{"nested":' * 18 + b'{"password":"hidden-value"}' + b'}' * 18
request = httpx.Request("POST", "https://upstream/mcp", content=body)
response = httpx.Response(500, request=request, content=body)
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
assert detail is not None and "hidden-value" not in detail
assert "nested" in detail and "REDACTED" in detail
@pytest.mark.parametrize("body", [b'client%5Fsecret=first+second&client_id=visible', b'client_secret=first%26second&client_id=visible'])
def test_encoded_form_credentials_are_decoded_before_redaction(body):
request = httpx.Request("POST", "https://upstream/token", content=body,
headers={"Content-Type":"application/x-www-form-urlencoded"})
response = httpx.Response(400, request=request, content=body,
headers={"Content-Type":"application/x-www-form-urlencoded"})
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
assert detail is not None and "client_id=visible" in detail
assert "first" not in detail and "second" not in detail
@pytest.mark.asyncio
@pytest.mark.parametrize("source", tuple(AuthResolution))
@pytest.mark.parametrize("method", ("GET", "DELETE", "POST"))
@ -291,3 +495,99 @@ async def test_concurrent_mcp_messages_record_on_their_own_http_scope() -> None:
await asyncio.gather(record(first, AuthResolution.stored_user_token), record(second, AuthResolution.per_request_header))
assert first.resolution() == "stored-user-token"
assert second.resolution() == "per-request-header"
@pytest.mark.parametrize("source", ["header", "bearer", "basic", "cookie", "query", "form", "json"])
def test_reflected_credentials_are_removed_from_normal_response_fields(source):
import base64
secret = "generic-credential-123"
headers = {"X-Custom":secret} if source == "header" else {"Authorization":"Bearer " + secret} if source == "bearer" else {"Authorization":"Basic " + base64.b64encode(("client:" + secret).encode()).decode()} if source == "basic" else {"Cookie":"session=" + secret} if source == "cookie" else {}
request = httpx.Request("POST", "https://upstream/token" + ("?credential=" + secret if source == "query" else ""),
headers=headers, data={"client_secret":secret} if source == "form" else None,
json={"nested":{"client_secret":secret}} if source == "json" else None)
response = httpx.Response(401, request=request, json={"error":"invalid_client", "error_description":"Rejected " + secret})
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response))
assert detail is not None and "invalid_client" in detail
assert secret not in detail and "REDACTED" in detail
@pytest.mark.parametrize("secret", ['value"with\ncharacters€', "R"])
def test_reflected_values_are_redacted_before_truncation_without_expanding_replacements(secret):
request = httpx.Request("POST", "https://upstream/token", json={"client_secret":secret})
response = httpx.Response(401, request=request, json={"error":"invalid_client", "detail":"x" * 460 + secret})
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response))
assert detail is not None and "invalid_client" in detail
assert "value" not in detail and "characters" not in detail and len(detail) < 1400
@pytest.mark.parametrize("headers", [{"Authorization":"Basic !!!"}, {"Cookie":"bad@key=opaque"}])
def test_malformed_auth_headers_do_not_break_failure_diagnostics(headers):
request = httpx.Request("POST", "https://upstream/token", headers=headers)
response = httpx.Response(401, request=request, json={"error":"invalid_client"})
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response))
assert detail is not None and "invalid_client" in detail
assert "!!!" not in detail and "opaque" not in detail
def test_oversized_request_omits_potentially_reflected_response_credentials():
request = httpx.Request("POST", "https://upstream/token", content=b"x" * 17000)
response = httpx.Response(401, request=request, json={"error_description":"unknown-secret"})
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response))
assert detail is not None and "capture limit" in detail and "credentials unavailable" in detail
assert "unknown-secret" not in detail
@pytest.mark.asyncio
async def test_streamed_error_redacts_reflected_credentials_before_capture():
import json
from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response
secret = "generic-credential-123"
request = httpx.Request("POST", "https://upstream/token", data={"client_secret":secret})
raw = json.dumps({"error":"invalid_client", "error_description":"Rejected " + secret}).encode()
response = httpx.Response(401, request=request, stream=httpx.ByteStream(raw))
await capture_upstream_error_response(response)
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response))
assert detail is not None and "invalid_client" in detail and "Rejected" in detail
assert secret not in detail and "REDACTED" in detail
assert await response.aread() == raw
@pytest.mark.parametrize("path", ["/credential-path-value/mcp", "/oauth/credential-path-value/token"])
def test_failure_diagnostics_omit_credential_bearing_url_paths(path):
request = httpx.Request("POST", "https://upstream.example" + path)
response = httpx.Response(401, request=request, json={"error": "access_denied"})
error = httpx.HTTPStatusError("denied", request=request, response=response)
diagnostic = describe_upstream_http_failure(error)
assert diagnostic is not None
assert "credential-path-value" not in diagnostic
assert "POST https://upstream.example/ -> HTTP 401" in diagnostic
assert "access_denied" in diagnostic
def test_deep_request_omits_response_when_credentials_cannot_be_inspected():
from litellm.proxy._experimental.mcp_server.utils import MAX_STRUCTURED_CONTENT_SCAN_DEPTH
raw = "[" * (MAX_STRUCTURED_CONTENT_SCAN_DEPTH + 1) + '{"client_secret":"nested-credential"}' + "]" * (MAX_STRUCTURED_CONTENT_SCAN_DEPTH + 1)
request = httpx.Request("POST", "https://upstream/token", content=raw, headers={"Content-Type": "application/json"})
response = httpx.Response(401, request=request, json={"error_description": "Rejected nested-credential"})
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response))
assert detail is not None and "HTTP 401" in detail
assert "response body: (omitted: request credentials unavailable)" in detail
assert "nested-credential" not in detail
@pytest.mark.parametrize("field", ["accessToken", "refreshToken", "clientSecret", "apikey", "CLIENTASSERTION", "cost_token"])
@pytest.mark.parametrize("encoding", ["json", "form"])
def test_compact_credential_fields_and_reflected_values_are_redacted(field, encoding):
secret = "generic-private-value"
fields = {field: secret}
request = httpx.Request("POST", "https://upstream/token", json=fields if encoding == "json" else None,
data=fields if encoding == "form" else None)
response = httpx.Response(401, request=request, json={field: secret, "error": "invalid_client", "detail": "Rejected " + secret})
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response))
assert detail is not None and "invalid_client" in detail
assert "REDACTED" in detail and secret not in detail

View file

@ -1,5 +1,6 @@
"""Unit tests for MCP OAuth passthrough tool-fetch behavior."""
import logging
import sys
from unittest.mock import AsyncMock, MagicMock
@ -11,7 +12,7 @@ if sys.version_info < (3, 11):
from exceptiongroup import ExceptionGroup
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.exceptions import MCPServerListError, MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
_extract_upstream_auth_failure,
@ -434,3 +435,43 @@ async def test_aggregate_with_single_accessible_server_still_absorbs():
assert listing.tools == []
assert listing.outcomes["delegate_docs"].tag == "auth_required"
@pytest.mark.asyncio
async def test_fetch_tools_logs_upstream_request_details_on_500(caplog):
manager = MCPServerManager()
request = httpx.Request(
"POST",
"https://upstream/apis/mcp",
headers={"Authorization": "Bearer upstream-token-0123456789"},
content=b'{"method":"initialize","jsonrpc":"2.0","id":0}',
)
response = httpx.Response(500, request=request)
mock_client = MagicMock()
mock_client.list_tools = AsyncMock(
side_effect=httpx.HTTPStatusError("500", request=request, response=response)
)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
with pytest.raises(MCPServerListError):
await manager._fetch_tools_with_timeout(mock_client, "sample_docs")
assert "POST https://upstream/ -> HTTP 500" in caplog.text
assert '"method":"initialize"' in caplog.text
assert "upstream-token-0123456789" not in caplog.text
@pytest.mark.asyncio
async def test_client_creation_failure_logs_sanitized_exchange(monkeypatch, caplog):
manager = MCPServerManager()
server = MCPServer(server_id="sample", name="sample", url="https://upstream/mcp", transport=MCPTransport.http, auth_type=MCPAuth.none)
request = httpx.Request("POST", "https://upstream/mcp?credential=query-secret")
response = httpx.Response(500, request=request, json={"error":"missing_scope"})
error = httpx.HTTPStatusError("query-secret", request=request, response=response)
monkeypatch.setattr(manager, "_create_mcp_client", AsyncMock(side_effect=error))
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
with pytest.raises(MCPServerListError):
await manager._get_tools_from_server(server)
assert "POST https://upstream/ -> HTTP 500" in caplog.text
assert "missing_scope" in caplog.text and "query-secret" not in caplog.text

View file

@ -12787,6 +12787,125 @@ async def test_debug_reports_legacy_signing_and_non_http_transport(transport: Li
finally:
request_ctx.reset(token)
@pytest.mark.asyncio
async def test_temporary_server_discovery_reuses_resolved_metadata_without_publishing() -> None:
manager: Final = MCPServerManager()
server: Final = MCPServer(
server_id="temporary-oauth-discovery", name="temporary", url="https://idp.example.com/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough,
)
manager._set_oauth_discovery_deferred(server.server_id, True)
metadata: Final = MCPOAuthMetadata(
authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token",
registration_url="https://idp.example.com/register",
)
with patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery:
resolved: Final = await manager.ensure_oauth_metadata_discovered(server)
repeated: Final = await manager.ensure_oauth_metadata_discovered(server)
assert resolved.authorization_url == metadata.authorization_url
assert resolved.token_url == metadata.token_url
assert resolved.registration_url == metadata.registration_url
assert repeated is resolved
assert server.server_id not in manager.registry
assert server.server_id not in manager.config_mcp_servers
discovery.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type", [MCPAuth.oauth2, MCPAuth.true_passthrough])
async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> None:
manager: Final = MCPServerManager()
server: Final = MCPServer(
server_id="repeated-stale", name="stale", url="https://idp.example.com/mcp",
transport=MCPTransport.http, auth_type=auth_type, oauth2_flow="authorization_code",
)
manager.registry[server.server_id] = server
manager._set_oauth_discovery_deferred(server.server_id, True)
metadata: Final = MCPOAuthMetadata(
authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token",
)
with (
patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery,
patch.object(manager, "_publish_resolved_oauth_server", return_value=None),
):
if auth_type == MCPAuth.true_passthrough:
assert await manager.ensure_oauth_metadata_discovered(server) is server
else:
with pytest.raises(HTTPException) as exc:
await manager.ensure_oauth_metadata_discovered(server)
assert exc.value.status_code == 503
assert "changed repeatedly" in str(exc.value.detail)
assert discovery.await_count == 2
@pytest.mark.asyncio
async def test_stale_discovery_falls_back_to_resolved_registered_server() -> None:
manager: Final = MCPServerManager()
original: Final = MCPServer(
server_id="resolved-replacement", name="replacement", url="https://old.example.com/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code",
)
replacement: Final = original.model_copy(update={
"url": "https://new.example.com/mcp", "authorization_url": "https://new.example.com/authorize",
"token_url": "https://new.example.com/token",
})
manager.registry[original.server_id] = replacement
assert await manager._rejoin_oauth_metadata_discovery(original, retry_stale=False) is replacement
def test_stale_discovery_cannot_overwrite_new_registered_server() -> None:
manager: Final = MCPServerManager()
original: Final = MCPServer(
server_id="stale-publication", name="publication", url="https://old.example.com/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.oauth2,
)
manager._set_oauth_discovery_deferred(original.server_id, True)
original_slot: Final = manager._oauth_discovery_slot(original.server_id)
assert original_slot is not None
replacement: Final = original.model_copy(update={"url": "https://new.example.com/mcp"})
manager.registry[original.server_id] = replacement
manager._set_oauth_discovery_deferred(original.server_id, True)
assert manager._publish_resolved_oauth_server(original, original_slot.generation) is None
assert manager.registry[original.server_id] is replacement
@pytest.mark.asyncio
async def test_temporary_oauth_discovery_expires_without_more_requests() -> None:
manager: Final = MCPServerManager()
server: Final = MCPServer(
server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough,
authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token",
)
manager._set_oauth_discovery_deferred(server.server_id, True)
resolved: Final = await manager.ensure_oauth_metadata_discovered(server)
assert manager._oauth_discovery_slot(server.server_id) is not None
loop: Final = asyncio.get_running_loop()
expired: Final = loop.create_future()
with patch.object(loop, "time", return_value=loop.time() + 301):
loop.call_later(0, expired.set_result, None)
await expired
assert resolved.authorization_url == server.authorization_url
assert manager._oauth_discovery_slot(server.server_id) is None
def test_old_temporary_discovery_expiry_preserves_replacement() -> None:
manager: Final = MCPServerManager()
manager._set_oauth_discovery_deferred("reused-session", True)
old_slot: Final = manager._oauth_discovery_slot("reused-session")
assert old_slot is not None
manager._set_oauth_discovery_deferred("reused-session", True)
replacement: Final = manager._oauth_discovery_slot("reused-session")
manager._expire_temporary_oauth_discovery("reused-session", old_slot.generation)
assert manager._oauth_discovery_slot("reused-session") is replacement
assert replacement is not None
manager._expire_temporary_oauth_discovery("reused-session", replacement.generation)
assert manager._oauth_discovery_slot("reused-session") is None
manager._expire_temporary_oauth_discovery("reused-session", replacement.generation)
assert manager._oauth_discovery_slot("reused-session") is None
@pytest.mark.asyncio
async def test_openapi_health_coalesces_concurrent_checks_and_reuses_results(respx_mock, monkeypatch):
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")

View file

@ -1,4 +1,5 @@
import asyncio
import hashlib
import json
from typing import Iterable, List, Optional, Tuple
from unittest.mock import patch
@ -7,6 +8,7 @@ import pytest
from redis.asyncio import Redis
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
AUTH_CACHE_INVALIDATION_CHANNEL,
AuthCacheInvalidationSubscriber,
@ -145,6 +147,35 @@ async def test_subscriber_deletes_local_cache_entry_on_message() -> None:
assert pubsub.subscribed_channels == [AUTH_CACHE_INVALIDATION_CHANNEL]
@pytest.mark.asyncio
async def test_subscriber_deletes_key_object_partition_entry_on_message() -> None:
"""
LIT-7563 moved user-key objects into their own in-memory partition; a key
invalidation broadcast must still evict the hashed-token entry there, or a
deleted key keeps authenticating on other workers until its TTL expires.
"""
hashed_token = hashlib.sha256(b"sk-lit7563-hot-key").hexdigest()
cache = UserApiKeyCache()
cache.set_cache(hashed_token, UserAPIKeyAuth(token=hashed_token), model_type=UserAPIKeyAuth)
assert cache.get_cache(hashed_token, model_type=UserAPIKeyAuth) is not None
pubsub = _QueuePubSub(initial_messages=[_invalidation_message(hashed_token)])
subscriber = AuthCacheInvalidationSubscriber(
redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[pubsub])),
user_api_key_cache=cache,
)
subscriber.start()
try:
for _ in range(200):
if cache.get_cache(hashed_token, model_type=UserAPIKeyAuth) is None:
break
await asyncio.sleep(0.01)
finally:
await subscriber.stop()
assert cache.get_cache(hashed_token, model_type=UserAPIKeyAuth) is None
@pytest.mark.asyncio
async def test_subscriber_deletes_additional_in_memory_cache_entry_on_message() -> None:
"""

View file

@ -1,3 +1,4 @@
import hashlib
import json
from typing import Any
@ -10,10 +11,14 @@ from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
end_user_cache_key,
get_management_object_ttl,
is_user_key_cache_key,
)
from litellm.proxy.proxy_server import UserAPIKeyCacheTTLEnum
HASHED_TOKEN = hashlib.sha256(b"sk-lit7563-hot-key").hexdigest()
class CapturingInMemoryCache(InMemoryCache):
"""Records ``ttl`` passed into ``set_cache`` (what DualCache injects)."""
@ -204,9 +209,7 @@ class TestUserApiKeyCache:
# Bypass UserApiKeyCache.serialize: CacheCodec rejects non-dict cached values
# for dict-based models (deserialize returns None).
await cache.in_memory_cache.async_set_cache(
key="k", value="invalid-payload-not-a-dict"
)
await cache.in_memory_cache.async_set_cache(key="k", value="invalid-payload-not-a-dict")
value = await cache.async_get_cache("k", model_type=UserAPIKeyAuth)
assert value is None
@ -224,6 +227,141 @@ class TestUserApiKeyCache:
fake.set_cache("k2", {"ok": NotSerializable()})
class TestUserKeyObjectPartition:
"""
Regression for LIT-7563: user-key objects share one 200-entry ``InMemoryCache`` with
every other management object, so end-user / team / tag churn evicts hot keys and
forces a ``LiteLLM_VerificationToken`` lookup on the next request.
"""
@pytest.mark.parametrize(
("key", "expected"),
[
(HASHED_TOKEN, True),
(HASHED_TOKEN.upper(), False),
(f"team_id:{HASHED_TOKEN}", False),
(end_user_cache_key("u1"), False),
("sk-lit7563-hot-key", False),
],
)
def test_is_user_key_cache_key(self, key: str, expected: bool):
assert is_user_key_cache_key(key) is expected
@pytest.mark.asyncio
async def test_management_object_churn_does_not_evict_key_object(self):
cache = UserApiKeyCache(in_memory_cache=InMemoryCache(max_size_in_memory=2))
await cache.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth, ttl=100)
for i in range(2):
await cache.async_set_cache(end_user_cache_key(f"u{i}"), {"user_id": f"u{i}"}, ttl=200)
key_obj = await cache.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth)
assert key_obj is not None
assert key_obj.token == HASHED_TOKEN
assert cache.get_cache(end_user_cache_key("u1")) == {"user_id": "u1"}
assert HASHED_TOKEN not in cache.in_memory_cache.cache_dict
def test_sync_write_and_read_route_to_key_object_partition(self):
cache = UserApiKeyCache(in_memory_cache=InMemoryCache(max_size_in_memory=2))
cache.set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth, ttl=100)
for i in range(2):
cache.set_cache(end_user_cache_key(f"u{i}"), {"user_id": f"u{i}"}, ttl=200)
key_obj = cache.get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth)
assert key_obj is not None
assert key_obj.token == HASHED_TOKEN
@pytest.mark.asyncio
async def test_redis_hit_backfills_key_object_partition_with_configured_ttl(self):
redis = FakeRedisCache()
writer = UserApiKeyCache(redis_cache=redis, default_in_memory_ttl=30)
await writer.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth)
key_partition = CapturingInMemoryCache()
reader = UserApiKeyCache(redis_cache=redis, default_in_memory_ttl=30, key_object_in_memory_cache=key_partition)
key_obj = await reader.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth)
assert key_obj is not None
assert key_obj.token == HASHED_TOKEN
assert key_partition.last_ttl == 30
assert HASHED_TOKEN not in reader.in_memory_cache.cache_dict
@pytest.mark.asyncio
async def test_update_cache_ttl_applies_to_key_object_partition(self):
key_partition = CapturingInMemoryCache()
cache = UserApiKeyCache(default_in_memory_ttl=60, key_object_in_memory_cache=key_partition)
cache.update_cache_ttl(default_in_memory_ttl=7, default_redis_ttl=7)
await cache.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth)
assert key_partition.last_ttl == 7
@pytest.mark.asyncio
async def test_attach_redis_cache_applies_to_key_object_partition(self):
redis = FakeRedisCache()
cache = UserApiKeyCache()
cache.attach_redis_cache(redis)
await cache.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth)
other_worker = UserApiKeyCache(redis_cache=redis)
key_obj = await other_worker.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth)
assert key_obj is not None
assert key_obj.token == HASHED_TOKEN
@pytest.mark.asyncio
async def test_delete_removes_key_object_from_partition_and_redis(self):
redis = FakeRedisCache()
cache = UserApiKeyCache(redis_cache=redis)
await cache.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth)
assert await cache.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) is not None
cache.delete_cache(HASHED_TOKEN)
assert await cache.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) is None
assert await redis.async_get_cache(HASHED_TOKEN) is None
@pytest.mark.asyncio
async def test_async_delete_removes_key_object_from_partition_and_redis(self):
redis = FakeRedisCache()
cache = UserApiKeyCache(redis_cache=redis)
await cache.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth)
await cache.async_delete_cache(HASHED_TOKEN)
assert await cache.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) is None
assert await redis.async_get_cache(HASHED_TOKEN) is None
@pytest.mark.asyncio
async def test_pipeline_write_routes_each_entry_to_its_partition(self):
cache = UserApiKeyCache(in_memory_cache=InMemoryCache(max_size_in_memory=2))
await cache.async_set_cache_pipeline(
[(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN))]
+ [(end_user_cache_key(f"u{i}"), {"user_id": f"u{i}"}) for i in range(2)],
ttl=100,
)
key_obj = await cache.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth)
assert key_obj is not None
assert key_obj.token == HASHED_TOKEN
assert HASHED_TOKEN not in cache.in_memory_cache.cache_dict
assert cache.get_cache(end_user_cache_key("u1")) == {"user_id": "u1"}
def test_flush_clears_key_object_partition(self):
cache = UserApiKeyCache()
cache.set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth)
cache.set_cache(end_user_cache_key("u1"), {"user_id": "u1"})
cache.flush_cache()
assert cache.get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) is None
assert cache.get_cache(end_user_cache_key("u1")) is None
def test_in_memory_cache_for_routes_by_key(self):
cache = UserApiKeyCache()
assert cache.in_memory_cache_for(HASHED_TOKEN) is cache.key_object_cache.in_memory_cache
assert cache.in_memory_cache_for(end_user_cache_key("u1")) is cache.in_memory_cache
class TestManagementObjectTTL:
"""
Regression for LIT-3338: ``general_settings.user_api_key_cache_ttl`` (which the
@ -238,19 +376,13 @@ class TestManagementObjectTTL:
def test_falls_back_to_constant_when_no_default_configured(self):
cache = UserApiKeyCache()
assert cache.default_in_memory_ttl is None
assert (
get_management_object_ttl(cache)
== DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
)
assert get_management_object_ttl(cache) == DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
def test_resolves_on_a_plain_dual_cache(self):
# Many call sites are typed UserApiKeyCache but exercised in tests with a
# bare DualCache; the resolver must work on the base type, not just the subclass.
assert get_management_object_ttl(DualCache(default_in_memory_ttl=300)) == 300
assert (
get_management_object_ttl(DualCache())
== DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
)
assert get_management_object_ttl(DualCache()) == DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
@pytest.mark.asyncio
async def test_management_write_uses_configured_ttl_over_constant(self):
@ -260,9 +392,7 @@ class TestManagementObjectTTL:
redis_cache=FakeRedisCache(),
default_in_memory_ttl=300,
)
assert get_management_object_ttl(cache) != (
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
)
assert get_management_object_ttl(cache) != (DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL)
await cache.async_set_cache(
"team_id:abc",

View file

@ -4,7 +4,7 @@ import contextlib
import json
import os
import traceback
from collections.abc import Mapping
from collections.abc import Iterator, Mapping
from types import MappingProxyType, SimpleNamespace
from typing import Final
from unittest import mock
@ -12,7 +12,9 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch
import httpx
import pytest
import respx
from fastapi import HTTPException, Request, Response
from fastapi.routing import APIRoute
from fastapi.responses import StreamingResponse
from fastapi.testclient import TestClient
from starlette.datastructures import FormData
@ -45,6 +47,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
)
from litellm.proxy._types import LitellmUserRoles, SpecialHeaders, UserAPIKeyAuth
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
@ -3339,8 +3342,11 @@ class TestOpenAIPassthroughRoute:
def _resolve_route_name(method: str, path: str) -> str | None:
from starlette.routing import Match
from litellm.proxy._lazy_features import LAZY_FEATURES, _force_load
from litellm.proxy.proxy_server import app
asyncio.run(_force_load(app, next(f for f in LAZY_FEATURES if f.name == "llm_passthrough")))
scope: Final = {
"type": "http",
"method": method,
@ -3350,8 +3356,8 @@ def _resolve_route_name(method: str, path: str) -> str | None:
"root_path": "",
}
for route in app.router.routes:
if route.matches(scope)[0] == Match.FULL:
return getattr(route, "name", None)
if isinstance(route, APIRoute) and route.matches(scope)[0] == Match.FULL:
return route.name
return None
@ -3376,7 +3382,7 @@ def test_openai_passthrough_prefix_wins_over_native_provider_routes(method, path
/{provider}/v1/files and /{provider}/v1/batches routes must never capture it
with provider="openai_passthrough" (which 500s on the LlmProviders lookup).
"""
assert _resolve_route_name(method, path) == "openai_proxy_route"
assert _resolve_route_name(method, path) == "openai_passthrough_route"
@pytest.mark.parametrize(
@ -3393,6 +3399,41 @@ def test_native_provider_routes_are_unchanged(method, path, expected_name):
assert _resolve_route_name(method, path) == expected_name
@pytest.fixture
def openai_passthrough_client(monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]:
from litellm.proxy.proxy_server import app
monkeypatch.setenv("OPENAI_API_KEY", "sk-upstream")
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, lambda: UserAPIKeyAuth(api_key="sk-virtual"))
yield TestClient(app)
@pytest.mark.parametrize(
"method, path, body",
[
("POST", "/v1/responses", {"model": "gpt-5.1", "input": "hi"}),
("GET", "/v1/files", None),
("POST", "/v1/batches", {"input_file_id": "file-abc123", "endpoint": "/v1/responses"}),
],
)
def test_openai_passthrough_forwards_verbatim_to_openai(
openai_passthrough_client: TestClient, method: str, path: str, body: dict[str, str] | None
) -> None:
"""Every /openai_passthrough request, including the /v1/files and /v1/batches
paths that native provider routes also claim, must reach OpenAI unchanged."""
with respx.mock(assert_all_called=True) as upstream:
route = upstream.request(method, f"https://api.openai.com{path}").mock(
return_value=httpx.Response(200, json={"id": "upstream_123"})
)
response = openai_passthrough_client.request(method, f"/openai_passthrough{path}", json=body)
assert (response.status_code, response.json()) == (200, {"id": "upstream_123"})
assert route.calls.last.request.headers["authorization"] == "Bearer sk-upstream"
class TestCursorProxyRoute:
"""Tests for the Cursor Cloud Agents pass-through route."""

View file

@ -7313,6 +7313,88 @@ async def test_update_general_settings_apply_user_budget_to_team_keys_yaml_wins(
assert ps.general_settings["apply_user_budget_to_team_keys"] is True
def _fill_user_api_key_cache(cache: DualCache, count: int) -> None:
for index in range(count):
cache.set_cache(key=f"key-{index}", value={"token": f"key-{index}"}, local_only=True)
@pytest.mark.asyncio
async def test_update_general_settings_user_api_key_cache_max_size_resizes_the_running_cache(monkeypatch):
"""The Admin UI writes the capacity to the DB config, so the running cache has
to pick it up on reload; otherwise the knob only works after a restart."""
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.proxy_server import ProxyConfig
cache = UserApiKeyCache()
monkeypatch.setattr(proxy_server_module, "general_settings", {})
monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache)
await ProxyConfig()._update_general_settings(db_general_settings={"user_api_key_cache_max_size": 300})
assert proxy_server_module.general_settings["user_api_key_cache_max_size"] == 300
_fill_user_api_key_cache(cache, 250)
assert cache.get_cache(key="key-0", local_only=True) == {"token": "key-0"}
@pytest.mark.asyncio
async def test_update_general_settings_clearing_user_api_key_cache_max_size_restores_the_default(monkeypatch):
"""Blanking the field in the dashboard deletes the key, so the cache must fall
back to the default capacity rather than keep the last configured size."""
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.proxy_server import ProxyConfig
cache = UserApiKeyCache()
cache.update_in_memory_max_size(5000)
monkeypatch.setattr(proxy_server_module, "general_settings", {"user_api_key_cache_max_size": 5000})
monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache)
await ProxyConfig()._update_general_settings(db_general_settings={"store_model_in_db": True})
assert "user_api_key_cache_max_size" not in proxy_server_module.general_settings
_fill_user_api_key_cache(cache, 201)
assert cache.get_cache(key="key-0", local_only=True) is None
@pytest.mark.asyncio
@pytest.mark.parametrize("db_value", [0, -5, "lots"])
async def test_update_general_settings_ignores_an_invalid_user_api_key_cache_max_size(db_value, monkeypatch):
"""A non-positive capacity would make the eviction loop pop an empty heap on the
next write, so a bad DB value must leave the running cache untouched."""
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.proxy_server import ProxyConfig
cache = UserApiKeyCache()
cache.update_in_memory_max_size(300)
monkeypatch.setattr(proxy_server_module, "general_settings", {})
monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache)
await ProxyConfig()._update_general_settings(db_general_settings={"user_api_key_cache_max_size": db_value})
assert "user_api_key_cache_max_size" not in proxy_server_module.general_settings
_fill_user_api_key_cache(cache, 250)
assert cache.get_cache(key="key-0", local_only=True) == {"token": "key-0"}
@pytest.mark.asyncio
async def test_update_general_settings_user_api_key_cache_max_size_yaml_wins(monkeypatch):
"""A DB value must not silently override an explicit YAML capacity on reload."""
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.proxy_server import ProxyConfig
proxy_config = ProxyConfig()
proxy_config._yaml_general_settings_keys = {"user_api_key_cache_max_size"}
cache = UserApiKeyCache()
cache.update_in_memory_max_size(300)
monkeypatch.setattr(proxy_server_module, "general_settings", {"user_api_key_cache_max_size": 300})
monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache)
await proxy_config._update_general_settings(db_general_settings={"user_api_key_cache_max_size": 10})
assert proxy_server_module.general_settings["user_api_key_cache_max_size"] == 300
_fill_user_api_key_cache(cache, 250)
assert cache.get_cache(key="key-0", local_only=True) == {"token": "key-0"}
@pytest.mark.asyncio
@pytest.mark.parametrize(
"db_value,expected",
@ -9183,6 +9265,126 @@ class TestLazyFeaturesNotImportedAtStartup:
class TestLazyFeatureMiddleware:
"""Behavior of the middleware itself, exercised in isolation."""
@pytest.mark.asyncio
async def test_llm_passthrough_loads_on_first_provider_request(self, monkeypatch):
"""An app that never registered the provider passthrough routes 404s a
provider request; behind the middleware the same request registers the
routes and is forwarded to the provider with the configured key."""
import respx
from fastapi import FastAPI
from litellm.proxy._lazy_features import LAZY_FEATURES, LazyFeatureMiddleware
monkeypatch.setenv("MISTRAL_API_KEY", "sk-upstream")
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
feat = next(f for f in LAZY_FEATURES if f.name == "llm_passthrough")
target_app = FastAPI()
target_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(api_key="sk-virtual")
mw = LazyFeatureMiddleware(target_app, fastapi_app=target_app, features=(feat,))
with respx.mock() as upstream:
route = upstream.get("https://api.mistral.ai/v1/models").mock(
return_value=httpx.Response(200, json={"object": "list", "data": []})
)
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=target_app), base_url="http://t") as bare:
assert (await bare.get("/mistral/v1/models")).status_code == 404
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=mw), base_url="http://t") as lazy:
response = await lazy.get("/mistral/v1/models")
assert (response.status_code, response.json()) == (200, {"object": "list", "data": []})
assert route.calls.last.request.headers["authorization"] == "Bearer sk-upstream"
def test_llm_passthrough_prefixes_cover_every_route_the_module_registers(self):
"""A route the module registers under a prefix the feature does not claim
would 404 until an unrelated provider request happens to load the module."""
from litellm.proxy._lazy_features import LAZY_FEATURES
feat = next(f for f in LAZY_FEATURES if f.name == "llm_passthrough")
paths = [r.path for r in importlib.import_module(feat.module_path).router.routes]
assert {"/mistral/{endpoint:path}", "/openai/{endpoint:path}"} <= set(paths)
unreachable = [p for p in paths if not feat.matches(p.replace("{endpoint:path}", "x"))]
assert unreachable == [], f"routes the middleware would never load: {unreachable}"
@pytest.mark.asyncio
@pytest.mark.parametrize("first_hit", ["/v1/realtime/calls", "/openai/v1/models"])
async def test_lazy_routes_land_in_registry_order_not_first_hit_order(self, first_hit):
"""Two lazy features with overlapping paths must answer a request with the
same handler no matter which one a deployment happens to hit first."""
from fastapi import APIRouter, FastAPI
from litellm.proxy._lazy_features import LazyFeature, LazyFeatureMiddleware
def make_register(path, handler):
def register(app, module):
router = APIRouter()
router.add_api_route(path, lambda: {"handler": handler}, methods=["POST"])
app.include_router(router)
return register
catch_all = LazyFeature(
name="catch_all",
module_path="json",
path_prefixes=("/openai/",),
register_fn=make_register("/openai/{endpoint:path}", "catch_all"),
)
specific = LazyFeature(
name="specific",
module_path="base64",
path_prefixes=("/openai/v1/realtime", "/v1/realtime"),
register_fn=make_register("/openai/v1/realtime/calls", "specific"),
)
target_app = FastAPI()
mw = LazyFeatureMiddleware(target_app, fastapi_app=target_app, features=(catch_all, specific))
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=mw), base_url="http://t") as client:
await client.post(first_hit)
await client.post("/openai/v1/models")
response = await client.post("/openai/v1/realtime/calls")
assert response.json() == {"handler": "catch_all"}
@pytest.mark.asyncio
@pytest.mark.parametrize("root_path", ["", "/api"])
async def test_reserved_slot_keeps_lazy_catch_all_ahead_of_later_eager_routes(self, root_path):
"""/{mcp_server_name}/mcp is registered after the provider passthrough router
at startup, so /mistral/mcp must keep reaching the provider catch-all once
that router loads lazily instead of being swallowed by the MCP route. The
native /mistral/v1/files route sits ahead of it, so that path neither loads
the feature nor changes owner, with or without a SERVER_ROOT_PATH prefix."""
from fastapi import APIRouter, FastAPI
from litellm.proxy._lazy_features import LazyFeature, LazyFeatureMiddleware, reserve_lazy_slot
def register(app, module):
router = APIRouter()
router.add_api_route("/mistral/{endpoint:path}", lambda: {"handler": "passthrough"}, methods=["POST"])
app.include_router(router)
passthrough = LazyFeature(
name="llm_passthrough", module_path="json", path_prefixes=("/mistral/",), register_fn=register
)
target_app = FastAPI(root_path=root_path)
target_app.add_api_route("/mistral/v1/files", lambda: {"handler": "files"}, methods=["POST"])
reserve_lazy_slot(target_app, "llm_passthrough", features=(passthrough,))
target_app.add_api_route("/{mcp_server_name}/mcp", lambda: {"handler": "mcp"}, methods=["POST"])
target_app.add_middleware(LazyFeatureMiddleware, fastapi_app=target_app, features=(passthrough,))
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=target_app), base_url="http://t") as client:
files_first = (await client.post(f"{root_path}/mistral/v1/files")).json()["handler"]
loaded_after_files = frozenset(target_app.state.lazy_loaded)
handlers = [
(await client.post(f"{root_path}{path}")).json()["handler"]
for path in ("/mistral/mcp", "/mistral/v1/files")
]
assert (files_first, loaded_after_files) == ("files", frozenset())
assert handlers == ["passthrough", "files"]
@pytest.mark.asyncio
async def test_first_request_triggers_load_subsequent_does_not(self):
from fastapi import FastAPI
@ -10224,6 +10426,27 @@ def test_get_config_list_includes_apply_user_budget_to_team_keys(monkeypatch):
app.dependency_overrides.clear()
def test_get_config_list_includes_user_api_key_cache_max_size(monkeypatch):
"""The Admin UI General Settings table renders whatever /config/list returns,
so the cache capacity has to be exposed there as an Integer to be editable."""
mock_prisma = MagicMock()
mock_config_table = MagicMock()
mock_config_table.find_first = AsyncMock(return_value=None)
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
monkeypatch.setattr(proxy_server_module, "prisma_client", mock_prisma)
app.dependency_overrides[proxy_server_module.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
)
try:
client = TestClient(app)
resp = client.get("/config/list", params={"config_type": "general_settings"})
assert resp.status_code == 200, resp.text
fields = {item["field_name"]: item for item in resp.json()}
assert fields["user_api_key_cache_max_size"]["field_type"] == "Integer"
finally:
app.dependency_overrides.clear()
def test_get_config_list_includes_budget_exceeded_throttle_percentage(monkeypatch):
"""The throttle fraction is a litellm_settings scalar surfaced on the General
Settings table as a Float field so it sits with the other global limits; it
@ -12826,6 +13049,48 @@ async def test_load_config_router_authorizes_fallback_targets_against_the_callin
assert router.fallback_access_check is router_fallback_access_check
@pytest.mark.asyncio
async def test_load_config_user_api_key_cache_max_size_keeps_more_than_200_entries(tmp_path, monkeypatch):
"""The auth cache used to be pinned at InMemoryCache's 200 entry default, so a
deployment with more keys than that evicted constantly and every request
fell through to the DB. The YAML knob has to raise the cap on the live cache."""
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.proxy_server import ProxyConfig
config_file = tmp_path / "config.yaml"
config_file.write_text(yaml.dump({"general_settings": {"user_api_key_cache_max_size": "1000"}}))
cache = UserApiKeyCache()
monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache)
await ProxyConfig().load_config(router=None, config_file_path=str(config_file))
_fill_user_api_key_cache(cache, 999)
assert cache.get_cache(key="key-0", local_only=True) == {"token": "key-0"}
@pytest.mark.asyncio
@pytest.mark.parametrize("bad_value", [0, -1, "unbounded"])
async def test_load_config_rejects_a_non_positive_user_api_key_cache_max_size(tmp_path, bad_value, monkeypatch):
"""InMemoryCache treats 0 as 'cache nothing' and a negative cap makes eviction
pop an empty heap, so the proxy must refuse to boot with such a value instead
of silently disabling auth caching."""
from pydantic import ValidationError
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.proxy_server import ProxyConfig
config_file = tmp_path / "config.yaml"
config_file.write_text(yaml.dump({"general_settings": {"user_api_key_cache_max_size": bad_value}}))
cache = UserApiKeyCache()
monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache)
with pytest.raises(ValidationError):
await ProxyConfig().load_config(router=None, config_file_path=str(config_file))
_fill_user_api_key_cache(cache, 150)
assert cache.get_cache(key="key-0", local_only=True) == {"token": "key-0"}
def test_docs_redoc_openapi_are_reachable_by_default():
"""
LIT-6745: the interactive/machine-readable docs surfaces are on by

View file

@ -0,0 +1,171 @@
import sys
from types import ModuleType
import httpx
import pytest
from fastapi import APIRouter, FastAPI
from fastapi.testclient import TestClient
from starlette.routing import Match
from litellm.proxy.route_priority import HOT_ROUTE_PATHS, hot_routes_first
FILLER_COUNT = 300
def _routes_scanned_before_dispatch(app: FastAPI, method: str, path: str) -> int:
"""Number of route.matches() calls Starlette's Router.app makes before it finds a full match."""
scope = {"type": "http", "method": method, "path": path, "root_path": "", "headers": [], "query_string": b""}
for i, route in enumerate(app.router.routes):
match, _ = route.matches(dict(scope))
if match == Match.FULL:
return i + 1
raise AssertionError(f"{method} {path} has no route")
def _hot_router() -> APIRouter:
router = APIRouter()
@router.get("/health/liveliness")
@router.get("/health/liveness")
async def liveliness():
return "I'm alive!"
@router.post("/v1/chat/completions")
@router.post("/chat/completions")
async def chat():
return {"object": "chat.completion"}
return router
def _app_with_filler_then_hot_routes() -> FastAPI:
app = FastAPI()
for i in range(FILLER_COUNT):
@app.get(f"/filler/{i}")
async def filler(i: int = i):
return {"filler": i}
app.include_router(_hot_router())
return app
def test_hot_routes_first_puts_hot_routes_ahead_of_everything_else():
app = _app_with_filler_then_hot_routes()
assert _routes_scanned_before_dispatch(app, "GET", "/health/liveliness") > FILLER_COUNT
app.router.routes = hot_routes_first(app.router.routes)
hot_count = sum(1 for r in app.router.routes if getattr(r, "path", None) in HOT_ROUTE_PATHS)
assert _routes_scanned_before_dispatch(app, "GET", "/health/liveliness") <= hot_count
assert _routes_scanned_before_dispatch(app, "GET", "/health/liveness") <= hot_count
assert _routes_scanned_before_dispatch(app, "POST", "/v1/chat/completions") <= hot_count
assert _routes_scanned_before_dispatch(app, "POST", "/chat/completions") <= hot_count
def test_hot_routes_first_keeps_the_other_routes_in_order_and_dispatching():
app = _app_with_filler_then_hot_routes()
before = [r.path for r in app.router.routes if getattr(r, "path", "").startswith("/filler/")]
app.router.routes = hot_routes_first(app.router.routes)
after = [r.path for r in app.router.routes if getattr(r, "path", "").startswith("/filler/")]
assert after == before
client = TestClient(app)
assert client.get("/health/liveliness").json() == "I'm alive!"
assert client.get("/filler/7").json() == {"filler": 7}
assert client.post("/v1/chat/completions").json() == {"object": "chat.completion"}
assert client.get("/v1/chat/completions").status_code == 405
assert client.get("/does/not/exist").status_code == 404
def test_hot_routes_first_is_idempotent():
app = _app_with_filler_then_hot_routes()
once = hot_routes_first(app.router.routes)
assert hot_routes_first(once) == once
@pytest.mark.asyncio
async def test_lazy_loaded_hot_route_moves_to_the_front(monkeypatch):
from litellm.proxy._lazy_features import LazyFeature, LazyFeatureMiddleware
messages_router = APIRouter()
@messages_router.post("/v1/messages")
async def messages():
return {"type": "message"}
fake_module = ModuleType("fake_anthropic_endpoints")
fake_module.router = messages_router
monkeypatch.setitem(sys.modules, fake_module.__name__, fake_module)
target_app = _app_with_filler_then_hot_routes()
target_app.router.routes = hot_routes_first(target_app.router.routes)
async def downstream(scope, receive, send):
await send({"type": "http.response.start", "status": 200, "headers": []})
await send({"type": "http.response.body", "body": b""})
feat = LazyFeature(name="anthropic", module_path=fake_module.__name__, path_prefixes=("/v1/messages",))
mw = LazyFeatureMiddleware(downstream, fastapi_app=target_app, features=(feat,))
async def receive():
return {"type": "http.request", "body": b"", "more_body": False}
async def send(message):
pass
await mw({"type": "http", "path": "/v1/messages", "method": "POST", "headers": []}, receive, send)
hot_count = sum(1 for r in target_app.router.routes if getattr(r, "path", None) in HOT_ROUTE_PATHS)
assert _routes_scanned_before_dispatch(target_app, "POST", "/v1/messages") <= hot_count
assert TestClient(target_app).post("/v1/messages").json() == {"type": "message"}
@pytest.mark.asyncio
async def test_hot_routes_first_keeps_reserved_lazy_slot_ahead_of_later_eager_routes():
"""Liveness is registered after the provider passthrough slot, so pulling it to the
front must not shift where the lazily loaded catch-all is spliced back in."""
from litellm.proxy._lazy_features import LazyFeature, LazyFeatureMiddleware, reserve_lazy_slot
def register(app, module):
router = APIRouter()
router.add_api_route("/mistral/{endpoint:path}", lambda: {"handler": "passthrough"}, methods=["POST"])
app.include_router(router)
passthrough = LazyFeature(
name="llm_passthrough", module_path="json", path_prefixes=("/mistral/",), register_fn=register
)
target_app = FastAPI()
target_app.add_api_route("/mistral/v1/files", lambda: {"handler": "files"}, methods=["POST"])
target_app.add_api_route("/mistral/v1/batches", lambda: {"handler": "batches"}, methods=["POST"])
reserve_lazy_slot(target_app, "llm_passthrough", features=(passthrough,))
target_app.include_router(_hot_router())
target_app.add_api_route("/{mcp_server_name}/mcp", lambda: {"handler": "mcp"}, methods=["POST"])
target_app.router.routes = hot_routes_first(target_app.router.routes)
target_app.add_middleware(LazyFeatureMiddleware, fastapi_app=target_app, features=(passthrough,))
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=target_app), base_url="http://t") as client:
batches_first = (await client.post("/mistral/v1/batches")).json()["handler"]
loaded_after_batches = frozenset(target_app.state.lazy_loaded)
handlers = [
(await client.post(path)).json()["handler"]
for path in ("/mistral/mcp", "/mistral/v1/files", "/mistral/v1/batches")
]
assert (batches_first, loaded_after_batches) == ("batches", frozenset())
assert handlers == ["passthrough", "files", "batches"]
hot_count = sum(1 for r in target_app.router.routes if getattr(r, "path", None) in HOT_ROUTE_PATHS)
assert _routes_scanned_before_dispatch(target_app, "GET", "/health/liveliness") <= hot_count
def test_proxy_app_dispatches_liveness_and_chat_completions_before_the_rest():
from litellm.proxy.proxy_server import app
hot_count = sum(1 for r in app.router.routes if getattr(r, "path", None) in HOT_ROUTE_PATHS)
assert hot_count >= 4
assert len(app.router.routes) > 100
assert _routes_scanned_before_dispatch(app, "GET", "/health/liveliness") <= hot_count
assert _routes_scanned_before_dispatch(app, "GET", "/health/liveness") <= hot_count
assert _routes_scanned_before_dispatch(app, "POST", "/v1/chat/completions") <= hot_count
assert _routes_scanned_before_dispatch(app, "POST", "/chat/completions") <= hot_count

View file

@ -236,6 +236,38 @@ async def test_record_turn_attributes_satisfaction_to_previous_response_model():
assert smart_after.alpha == pytest.approx(smart_before.alpha)
@pytest.mark.asyncio
async def test_external_default_keeps_feedback_history_without_entering_bandit_pool():
r = _make_router()
before = r._cells[(RequestType.GENERAL, "fast")]
await r.record_turn(
session_id="fallback",
model_name="fast",
request_type=RequestType.GENERAL,
turn=Turn(user_content="fix this retry bug", assistant_content="clear the cache"),
)
await r.record_turn(
session_id="fallback",
model_name="external-default",
request_type=RequestType.GENERAL,
turn=Turn(user_content="the fix is still broken", assistant_content="keep cache entries"),
)
assert r._cells[(RequestType.GENERAL, "fast")].beta > before.beta
await r.record_turn(
session_id="fallback",
model_name="smart",
request_type=RequestType.GENERAL,
turn=Turn(
user_content="the fix is still broken",
assistant_content="use the corrected entry",
tool_results=[{"is_error": True, "content": "failure"}],
),
)
assert r._feedback_contexts["fallback"].model_name == "smart"
assert all(model != "external-default" for _, model in r._cells)
assert r.config.available_models == ["fast", "smart"]
@pytest.mark.asyncio
async def test_record_turn_bounds_feedback_contexts_and_evicts_least_recent_session():
r = _make_router()

View file

@ -5,18 +5,19 @@ Tests the rule-based complexity scoring and tier assignment logic.
"""
import asyncio
from collections.abc import AsyncIterator
import json
from copy import deepcopy
from functools import partial
import logging
import sys
import time
from typing import Dict, Final, List
from collections.abc import AsyncIterator, Mapping
from copy import deepcopy
from functools import partial
from typing import Dict, Final, List, Literal
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import httpx
import pytest
import respx
from pydantic import ValidationError
import litellm
@ -69,6 +70,7 @@ from litellm.router_strategy.complexity_router.tier_predictor import (
from litellm.types.router import (
Deployment,
LiteLLM_Params,
RouterErrors,
TaggedPreRoutingStrategy,
)
from litellm.types.llms.openai import ResponsesAPIResponse
@ -2671,7 +2673,8 @@ class TestEncryptedTaskClassifier:
assert call["metadata"]["user_api_key_hash"] == "caller-key-hash"
assert call["proxy_server_request"]["body"]["input"] == call["input"]
assert call["proxy_server_request"]["originating_request_masked"] == {
"input": [task], "metadata": {"authorization": "REDACTED"},
"input": [task],
"metadata": {"authorization": "REDACTED"},
}
assert "source-secret" not in json.dumps(call)
assert "originating_request_masked" not in call["proxy_server_request"]["body"]
@ -3315,19 +3318,23 @@ class TestLLMClassifier:
]
@pytest.mark.asyncio
@pytest.mark.parametrize("source_body", [
{"model": "router", "messages": [{"role": "user", "content": "source-only"}]},
{"model": "router", "system": "source-only", "messages": [{"role": "user", "content": "ask"}]},
{"model": "router", "instructions": "source-only", "input": "ask"},
])
@pytest.mark.parametrize(
"source_body",
[
{"model": "router", "messages": [{"role": "user", "content": "source-only"}]},
{"model": "router", "system": "source-only", "messages": [{"role": "user", "content": "ask"}]},
{"model": "router", "instructions": "source-only", "input": "ask"},
],
)
async def test_classifier_source_is_masked_and_separate_from_provider_input(
self, llm_complexity_router, mock_router_instance, source_body
):
mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}'))
outcome = await llm_complexity_router.aclassify(
"classify-this-ask", request_kwargs={"proxy_server_request": {
"body": {**source_body, "metadata": {"authorization": "source-secret"}}
}}
"classify-this-ask",
request_kwargs={
"proxy_server_request": {"body": {**source_body, "metadata": {"authorization": "source-secret"}}}
},
)
assert outcome.cause == "llm_classifier"
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
@ -8138,7 +8145,9 @@ class TestContextAwareClassifier:
),
),
)
def test_only_text_reminder_tails_are_ignored_for_new_asks(self, tail: list[dict[str, object]], expected: bool) -> None:
def test_only_text_reminder_tails_are_ignored_for_new_asks(
self, tail: list[dict[str, object]], expected: bool
) -> None:
from litellm.router_strategy.complexity_router.complexity_router import (
_CODEX_REMINDER_MARKERS,
_newest_turn_is_human_ask,
@ -13222,6 +13231,528 @@ class TestModalityRouting:
assert cache.async_set_cache.await_args.kwargs["value"] == {"model": "text-cheap", "tier": "SIMPLE"}
@pytest.mark.usefixtures("local_model_cost_map")
class TestHealthFallbackDispatch:
@pytest.fixture(autouse=True)
def httpx_transport(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
@staticmethod
def _router(
surface: str = "chat",
*,
peer: bool = False,
session: bool = False,
tagged: bool = False,
budgeted: bool = False,
config: Mapping[str, object] | None = None,
) -> Router:
provider: Final = "anthropic/claude-sonnet-5" if surface == "messages" else "openai/gpt-5.6"
base_suffix: Final = "" if surface == "messages" else "/v1"
return Router(
model_list=[
{
"model_name": "health-router",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_default_model": (config or {}).get("default_model", "fallback"),
"complexity_router_config": {
"tiers": {"SIMPLE": ["primary", "peer"] if peer else "primary", "MEDIUM": "primary"},
"session_affinity": session,
"deployment_affinity": False,
"max_tokens_from_tier_model": False,
**(config or {}),
},
},
},
*[
{
"model_name": name,
"litellm_params": {
"model": provider,
"api_key": "test-only",
"api_base": f"https://{name}.test{base_suffix}",
**({"tags": [name]} if tagged else {}),
**(
{"max_budget": 1.0, "budget_duration": "1d"}
if budgeted and name == "primary"
else {}
),
},
"model_info": {"id": f"{name}-id"},
}
for name in ("primary", "peer", "fallback")
],
],
num_retries=0,
enable_health_check_routing=True,
enable_tag_filtering=tagged,
)
@staticmethod
def _unavailable(router: Router, model_id: str, source: Literal["health", "cooldown"]) -> None:
if source == "health":
router.health_state_cache.set_deployment_health_states(
{model_id: {"is_healthy": False, "timestamp": time.time()}}
)
else:
router.cooldown_cache.add_deployment_to_cooldown(
model_id=model_id,
original_exception=RuntimeError("unavailable"),
exception_status=503,
cooldown_time=60,
)
@staticmethod
def _http_response(request: httpx.Request) -> httpx.Response:
body: Final = json.loads(request.content)
text: Final = request.url.host.split(".")[0]
payload: Final[Mapping[str, object]]
events: Final[tuple[Mapping[str, object], ...]]
if request.url.path.endswith("/responses"):
from litellm.responses.main import mock_responses_api_response
payload = mock_responses_api_response(text).model_dump()
events = (
{"type": "response.created", "response": {**payload, "status": "in_progress"}, "sequence_number": 0},
{
"type": "response.output_text.delta",
"delta": text,
"item_id": "msg_test",
"output_index": 0,
"content_index": 0,
"sequence_number": 1,
},
{"type": "response.completed", "response": payload, "sequence_number": 2},
)
elif request.url.path.endswith("/messages"):
payload = {
"id": "msg_test",
"type": "message",
"role": "assistant",
"model": body["model"],
"content": [{"type": "text", "text": text}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 10, "output_tokens": 1},
}
events = (
{"type": "message_start", "message": {**payload, "content": [], "stop_reason": None}},
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}},
{"type": "content_block_stop", "index": 0},
{"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 1}},
{"type": "message_stop"},
)
else:
payload = {
"id": "chatcmpl-test",
"object": "chat.completion",
"created": 1,
"model": body["model"],
"choices": [{"index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 1, "total_tokens": 11},
}
events = (
{
**payload,
"object": "chat.completion.chunk",
"choices": [{"index": 0, "delta": {"content": text}, "finish_reason": None}],
},
{
**payload,
"object": "chat.completion.chunk",
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
},
)
if not body.get("stream"):
return httpx.Response(200, json=payload)
wire: Final = "".join(
(f"event: {event['type']}\n" if "type" in event else "") + f"data: {json.dumps(event)}\n\n"
for event in events
)
return httpx.Response(
200,
text=wire + ("data: [DONE]\n\n" if "type" not in events[0] else ""),
headers={"content-type": "text/event-stream"},
)
@staticmethod
async def _request(router: Router, surface: str, stream: bool, metadata: dict[str, object]) -> str:
if surface == "responses":
result = await router.aresponses(
model="health-router", input="Hello!", stream=stream, litellm_metadata=metadata
)
elif surface == "messages":
result = await router.aanthropic_messages(
model="health-router",
messages=[{"role": "user", "content": "Hello!"}],
max_tokens=32,
stream=stream,
litellm_metadata=metadata,
)
else:
result = await router.acompletion(
model="health-router",
messages=[{"role": "user", "content": "Hello!"}],
stream=stream,
metadata=metadata,
)
if not stream:
payload = result if isinstance(result, dict) else result.model_dump()
if surface == "responses":
return payload["output"][0]["content"][0]["text"]
if surface == "messages":
return payload["content"][0]["text"]
return payload["choices"][0]["message"]["content"]
if surface == "messages":
wire: Final = b"".join([chunk async for chunk in result]).decode()
events = tuple(json.loads(line[6:]) for line in wire.splitlines() if line.startswith("data: "))
assert events[-1]["type"] == "message_stop"
return "".join(c["delta"]["text"] for c in events if c["type"] == "content_block_delta")
chunks: Final = [chunk.model_dump() async for chunk in result]
if surface == "responses":
assert chunks[-1]["type"] == "response.completed"
return "".join(c["delta"] for c in chunks if c["type"] == "response.output_text.delta")
assert chunks[-1]["choices"][0]["finish_reason"] == "stop"
return "".join(c["choices"][0]["delta"].get("content") or "" for c in chunks if c["choices"])
@pytest.mark.asyncio
@pytest.mark.parametrize("surface", ["chat", "responses", "messages"])
@pytest.mark.parametrize("stream", [False, True])
@pytest.mark.parametrize("source", ["health", "cooldown"])
async def test_public_call_falls_back_and_recovers(
self, surface: str, stream: bool, source: Literal["health", "cooldown"]
) -> None:
router: Final = self._router(surface, session=True)
self._unavailable(router, "primary-id", source)
metadata: Final[dict[str, object]] = {"session_id": "outage"}
with respx.mock(assert_all_mocked=True) as upstream:
upstream.post(host__regex=r"^(primary|peer|fallback)\.test$").mock(side_effect=self._http_response)
assert await self._request(router, surface, stream, metadata) == "fallback"
assert metadata["routing_decision"]["cause"] == "health_default_fallback"
assert "tier" not in metadata["routing_decision"]
assert "health_displaced:primary" in metadata["routing_decision"]["signals"]
assert [c.request.url.host for c in upstream.calls] == ["fallback.test"]
strategy: Final = router.complexity_routers["health-router"][0].strategy
key: Final = strategy._get_session_affinity_cache_key("outage", {})
assert await router.cache.async_get_cache(key=key) is None
if source == "health":
router.health_state_cache.set_deployment_health_states(
{"primary-id": {"is_healthy": True, "timestamp": time.time()}}
)
else:
router.cooldown_cache.cooldown_store.delete_cache(
router.cooldown_cache.get_cooldown_cache_key("primary-id")
)
recovered: Final[dict[str, object]] = {"session_id": "outage"}
assert await self._request(router, surface, stream, recovered) == "primary"
assert recovered["routing_decision"]["routed_model"] == "primary"
assert [c.request.url.host for c in upstream.calls] == ["fallback.test", "primary.test"]
@pytest.mark.asyncio
@pytest.mark.parametrize("source", ["health", "cooldown"])
async def test_partial_group_then_peer_then_default(self, source: Literal["health", "cooldown"]) -> None:
router: Final = self._router(peer=True, session=True)
router.add_deployment(
Deployment(
model_name="primary",
litellm_params=LiteLLM_Params(
model="openai/gpt-5.6", api_key="test-only", api_base="https://primary.test/v1"
),
model_info={"id": "primary-sibling-id"},
)
)
strategy: Final = router.complexity_routers["health-router"][0].strategy
key: Final = strategy._get_session_affinity_cache_key("precedence", {})
await router.cache.async_set_cache(key=key, value={"model": "primary", "tier": "SIMPLE"}, ttl=600)
with respx.mock(assert_all_mocked=True) as upstream:
upstream.post(host__regex=r"^(primary|peer|fallback)\.test$").mock(side_effect=self._http_response)
for model_id, expected, cause in (
("primary-id", "primary", "session_affinity_pin"),
("primary-sibling-id", "peer", "health_failover"),
("peer-id", "fallback", "health_default_fallback"),
):
self._unavailable(router, model_id, source)
metadata: Final[dict[str, object]] = {"session_id": "precedence"}
assert await self._request(router, "chat", False, metadata) == expected
assert metadata["routing_decision"]["cause"] == cause
assert await router.cache.async_get_cache(key=key) == {"model": "primary", "tier": "SIMPLE"}
assert [c.request.url.host for c in upstream.calls] == ["primary.test", "peer.test", "fallback.test"]
@pytest.mark.asyncio
async def test_spent_deployment_budget_falls_back_to_the_default(self, monkeypatch: pytest.MonkeyPatch) -> None:
"""A spent budget leaves the tier with nothing that may serve the request, and the budget
filter reports that as a bare ValueError instead of a typed router error. Reading it as
capacity skips the recovery and fails the request the recovery exists for."""
async def _no_sync(*args: object, **kwargs: object) -> None:
return None
monkeypatch.setattr(
"litellm.router_strategy.budget_limiter.RouterBudgetLimiting.periodic_sync_in_memory_spend_with_redis",
_no_sync,
)
monkeypatch.setattr(litellm, "callbacks", [])
router: Final = self._router(budgeted=True)
limiter: Final = router.router_budget_logger
assert limiter is not None, "a deployment max_budget must install the budget limiter"
await router.cache.async_set_cache(key="deployment_spend:primary-id:1d", value=2.0)
with respx.mock(assert_all_mocked=True) as upstream:
upstream.post(host__regex=r"^(primary|fallback)\.test$").mock(side_effect=self._http_response)
metadata: Final[dict[str, object]] = {}
assert await self._request(router, "chat", False, metadata) == "fallback"
assert metadata["routing_decision"]["cause"] == "health_default_fallback"
assert [c.request.url.host for c in upstream.calls] == ["fallback.test"]
@pytest.mark.asyncio
async def test_concurrent_tag_scopes_keep_fallbacks_request_local(self) -> None:
router: Final = self._router(tagged=True)
router.add_deployment(
Deployment(
model_name="fallback",
litellm_params=LiteLLM_Params(
model="openai/gpt-5.6", api_key="test-only", api_base="https://peer.test/v1", tags=["peer"]
),
model_info={"id": "fallback-peer-id"},
)
)
self._unavailable(router, "primary-id", "cooldown")
with respx.mock(assert_all_mocked=True) as upstream:
upstream.post(host__regex=r"^(peer|fallback)\.test$").mock(side_effect=self._http_response)
scopes: Final = tuple({"tags": [name], "session_id": name} for name in ("peer", "fallback"))
results: Final = await asyncio.gather(
*(self._request(router, "chat", False, metadata) for metadata in scopes)
)
assert results == ["peer", "fallback"]
assert [m["tags"] for m in scopes] == [["peer"], ["fallback"]]
assert [m["routing_decision"]["routed_model"] for m in scopes] == ["fallback", "fallback"]
assert sorted(c.request.url.host for c in upstream.calls) == ["fallback.test", "peer.test"]
@pytest.mark.asyncio
async def test_probe_preserves_consumed_request_exclusions(self) -> None:
router: Final = self._router()
self._unavailable(router, "primary-id", "cooldown")
kwargs: Final = {"_excluded_deployment_ids": ["fallback-id"], "_target_order": 1}
strategy: Final = router.complexity_routers["health-router"][0].strategy
response: Final = await strategy.async_pre_routing_hook(
model="health-router", messages=[{"role": "user", "content": "Hello!"}], request_kwargs=kwargs
)
assert response.model == "primary"
assert kwargs == {"_excluded_deployment_ids": ["fallback-id"], "_target_order": 1}
@pytest.mark.asyncio
@pytest.mark.parametrize("default_state", ["cooldown", "unconfigured", "same-model"])
async def test_unavailable_default_preserves_no_deployment_error(self, default_state: str) -> None:
from litellm.types.router import RouterRateLimitError
router: Final = self._router(config={"default_model": "primary"} if default_state == "same-model" else None)
self._unavailable(router, "primary-id", "cooldown")
if default_state == "unconfigured":
router.delete_deployment(id="fallback-id")
elif default_state == "cooldown":
self._unavailable(router, "fallback-id", "cooldown")
with respx.mock(assert_all_mocked=True) as upstream:
with pytest.raises(RouterRateLimitError, match="No deployments available"):
await self._request(router, "chat", False, {})
assert not upstream.calls
@pytest.mark.asyncio
@pytest.mark.parametrize("plan_active", [False, True])
async def test_plan_floor_outage_cannot_use_untiered_default(self, plan_active: bool) -> None:
from litellm.types.router import RouterRateLimitError
router: Final = self._router(
config={"tiers": {"SIMPLE": "primary", "MEDIUM": "peer"}, "plan_mode_min_tier": "MEDIUM"}
)
self._unavailable(router, "primary-id", "cooldown")
self._unavailable(router, "peer-id", "cooldown")
metadata: Final = {}
with respx.mock(assert_all_mocked=True, assert_all_called=False) as upstream:
upstream.post(host="fallback.test").mock(side_effect=self._http_response)
if plan_active:
with pytest.raises(RouterRateLimitError, match="No deployments available"):
await router.acompletion(
model="health-router",
messages=[
{"role": "system", "content": "Plan mode is active"},
{"role": "user", "content": "Hello!"},
],
metadata=metadata,
)
assert not upstream.calls
assert metadata["routing_decision"]["routed_model"] == "peer"
assert metadata["routing_decision"]["tier"] == "MEDIUM"
else:
assert await self._request(router, "chat", False, metadata) == "fallback"
@pytest.mark.asyncio
async def test_default_dispatch_drops_displaced_tier_params(self) -> None:
router: Final = self._router(
config={"tiers": {"SIMPLE": {"model_name": "primary", "litellm_params": {"max_tokens": 9}}}}
)
with respx.mock(assert_all_mocked=True) as upstream:
upstream.post(host__regex=r"^(primary|fallback)\.test$").mock(side_effect=self._http_response)
await router.acompletion(
model="health-router", messages=[{"role": "user", "content": "Hello!"}], max_tokens=32
)
assert json.loads(upstream.calls[-1].request.content)["max_completion_tokens"] == 9
self._unavailable(router, "primary-id", "cooldown")
await router.acompletion(
model="health-router", messages=[{"role": "user", "content": "Hello!"}], max_tokens=32
)
assert json.loads(upstream.calls[-1].request.content)["max_completion_tokens"] == 32
assert upstream.calls[-1].request.url.host == "fallback.test"
@pytest.mark.asyncio
@pytest.mark.parametrize("source", ["health", "cooldown"])
async def test_pinned_session_returns_to_primary_after_outage(self, source: Literal["health", "cooldown"]) -> None:
router: Final = self._router(session=True)
with respx.mock(assert_all_mocked=True) as upstream:
upstream.post(host__regex=r"^(primary|fallback)\.test$").mock(side_effect=self._http_response)
assert await self._request(router, "chat", False, {"session_id": "pinned"}) == "primary"
self._unavailable(router, "primary-id", source)
outage: Final[dict[str, object]] = {"session_id": "pinned"}
assert await self._request(router, "chat", False, outage) == "fallback"
assert outage["routing_decision"]["cause"] == "health_default_fallback"
if source == "health":
router.health_state_cache.set_deployment_health_states(
{"primary-id": {"is_healthy": True, "timestamp": time.time()}}
)
else:
router.cooldown_cache.cooldown_store.delete_cache(
router.cooldown_cache.get_cooldown_cache_key("primary-id")
)
recovered: Final[dict[str, object]] = {"session_id": "pinned"}
assert await self._request(router, "chat", False, recovered) == "primary"
assert recovered["routing_decision"]["cause"] == "session_affinity_pin"
assert [c.request.url.host for c in upstream.calls] == ["primary.test", "fallback.test", "primary.test"]
@pytest.mark.asyncio
async def test_policy_plugin_does_not_escape_to_live_default(self) -> None:
from litellm.types.router import RouterRateLimitError, RoutingContext
class PrimaryOnly:
async def run(self, context: RoutingContext) -> RoutingContext:
context.candidate_models = [name for name in context.candidate_models if name == "primary"]
return context
router: Final = self._router(peer=True, config={"plugins": [PrimaryOnly()]})
self._unavailable(router, "primary-id", "cooldown")
with respx.mock(assert_all_mocked=True) as upstream:
with pytest.raises(RouterRateLimitError, match="No deployments available"):
await self._request(router, "chat", False, {})
assert not upstream.calls
@pytest.mark.asyncio
@pytest.mark.parametrize("live_tier", [True, False])
@pytest.mark.parametrize("default_fits", [True, False])
async def test_context_recovery_precedes_default_with_prechecks_off(
self, live_tier: bool, default_fits: bool
) -> None:
from litellm.types.router import RouterRateLimitError
router: Final = self._router(config={"tiers": {"SIMPLE": "primary", "MEDIUM": "peer", "COMPLEX": "large"}})
router.add_deployment(
Deployment(
model_name="large",
litellm_params=LiteLLM_Params(
model="openai/gpt-5.6", api_key="test-only", api_base="https://large.test/v1"
),
model_info={"id": "large-id", "max_input_tokens": 10000},
)
)
for deployment in router.model_list:
deployment["model_info"]["max_input_tokens"] = (
10
if deployment["model_name"] == "primary"
or (deployment["model_name"] == "fallback" and not default_fits)
else 10000
)
self._unavailable(router, "peer-id", "cooldown")
if not live_tier:
self._unavailable(router, "large-id", "cooldown")
assert router.enable_pre_call_checks is False
metadata: Final = {}
messages: Final = [{"role": "user", "content": "hello " * 100}]
with respx.mock(assert_all_mocked=True, assert_all_called=False) as upstream:
upstream.post(host__regex=r"^(large|fallback)\.test$").mock(side_effect=self._http_response)
if not live_tier and not default_fits:
with pytest.raises(RouterRateLimitError, match="No deployments available"):
await router.acompletion(model="health-router", messages=messages, metadata=metadata)
assert not upstream.calls
else:
result: Final = await router.acompletion(model="health-router", messages=messages, metadata=metadata)
expected: Final = "large" if live_tier else "fallback"
assert result.choices[0].message.content == expected
assert upstream.calls[-1].request.url.host == f"{expected}.test"
assert metadata["routing_decision"].get("tier") == ("COMPLEX" if live_tier else None)
@pytest.mark.asyncio
@pytest.mark.parametrize("live_tier", [True, False])
async def test_modality_recovery_precedes_default(self, live_tier: bool) -> None:
router: Final = self._router(
config={"modality_routing": True, "tiers": {"SIMPLE": "primary", "MEDIUM": "peer", "COMPLEX": "vision"}}
)
router.add_deployment(
Deployment(
model_name="vision",
litellm_params=LiteLLM_Params(
model="openai/gpt-5.6", api_key="test-only", api_base="https://vision.test/v1"
),
model_info={"id": "vision-id", "supports_vision": True},
)
)
for deployment in router.model_list:
deployment["model_info"]["supports_vision"] = deployment["model_name"] != "primary"
self._unavailable(router, "peer-id", "cooldown")
if not live_tier:
self._unavailable(router, "vision-id", "cooldown")
with respx.mock(assert_all_mocked=True) as upstream:
upstream.post(host__regex=r"^(vision|fallback)\.test$").mock(side_effect=self._http_response)
result: Final = await router.acompletion(
model="health-router",
messages=[
{
"role": "user",
"content": [
{"type": "text", "text": "Hello!"},
{"type": "image_url", "image_url": {"url": "data:image/png;base64,aGk="}},
],
}
],
)
expected: Final = "vision" if live_tier else "fallback"
assert result.choices[0].message.content == expected
assert upstream.calls[-1].request.url.host == f"{expected}.test"
@pytest.mark.asyncio
@pytest.mark.parametrize("default_fits", [True, False])
async def test_modality_default_must_also_fit_context(self, default_fits: bool) -> None:
router: Final = self._router(config={"modality_routing": True, "tiers": {"SIMPLE": "primary"}})
for deployment in router.model_list:
deployment["model_info"]["supports_vision"] = deployment["model_name"] == "fallback"
deployment["model_info"]["max_input_tokens"] = 10000 if default_fits else 10
with respx.mock(assert_all_mocked=True, assert_all_called=False) as upstream:
upstream.post(host="fallback.test").mock(side_effect=self._http_response)
messages: Final = [
{
"role": "user",
"content": [
{"type": "text", "text": "hello " * 100},
{"type": "image_url", "image_url": {"url": "data:image/png;base64,aGk="}},
],
}
]
if default_fits:
result: Final = await router.acompletion(model="health-router", messages=messages)
assert result.choices[0].message.content == "fallback"
else:
with pytest.raises(litellm.BadRequestError, match="modality_routing is enabled"):
await router.acompletion(model="health-router", messages=messages)
assert not upstream.calls
class TestTierHealthFailover:
"""A tier whose decided model group is entirely in cooldown falls back to a live peer."""
@ -13256,7 +13787,7 @@ class TestTierHealthFailover:
probed_prompts = []
async def get_healthy_deployments(
model, request_kwargs, messages=None, input=None, parent_otel_span=None, **kwargs
model, request_kwargs, messages=None, input=None, parent_otel_span=None, health_check_probe=False
):
probed_kwargs.append(request_kwargs)
probed_prompts.append((messages, input))
@ -13758,6 +14289,51 @@ class TestTierHealthFailover:
for _, probed_input in router.litellm_router_instance.probed_prompts
), "the eligibility probe must forward `input` to the owner"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"raised, expected",
[
(ValueError(f"{RouterErrors.no_deployments_with_tag_routing.value}. Passed model=b"), {"live-c"}),
(
ValueError(f"{RouterErrors.no_deployments_with_provider_budget_routing.value}: b over budget"),
{"live-c"},
),
(ValueError("cannot unpack non-sequence"), {"exhausted-b", "live-c"}),
],
)
async def test_a_marked_exhaustion_value_error_is_a_verdict_and_an_unmarked_one_is_not(
self, mock_router_instance, raised, expected
):
"""Budget and tag filters exhaust a group without a typed error, signalling it only by a
RouterErrors marker on a bare ValueError. Those are verdicts; any other ValueError is a
fault, and a fault must still read as capacity rather than silently rerouting."""
router = self._router(
mock_router_instance,
{
"tiers": {
"SIMPLE": ["dead-a", "exhausted-b", "live-c"],
"MEDIUM": "mid",
"COMPLEX": "big",
"REASONING": "top",
},
"session_affinity": True,
},
{"dead-a": ["id-a1"], "exhausted-b": ["id-b1"], "live-c": ["id-c1"]},
cooling=("id-a1",),
raises_for={"exhausted-b": raised},
)
key = router._get_session_affinity_cache_key("sess-exhausted", {})
await router.litellm_router_instance.cache.async_set_cache(
key=key, value={"model": "dead-a", "tier": "SIMPLE"}, ttl=600
)
results = [
await router.async_pre_routing_hook(
model="m", request_kwargs={"metadata": {"session_id": "sess-exhausted"}}, messages=self.SIMPLE_MESSAGE
)
for _ in range(20)
]
assert {r.model for r in results} == expected
@pytest.mark.asyncio
async def test_a_group_the_router_has_no_deployment_for_is_not_a_failover_target(self, mock_router_instance):
"""The owner answers an unconfigured group with BadRequestError. Reading that as live

View file

@ -5,11 +5,68 @@ Tests the write/read/delete cycle for JSON and simple string secrets.
"""
import json
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import respx
import litellm
from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2
from litellm.types.secret_managers.main import KeyManagementSettings
_STATIC_CREDENTIALS = {"aws_access_key_id": "test-key", "aws_secret_access_key": "test-secret"}
_CMK_ARN = "arn:aws:kms:us-east-1:123456789012:key/11111111-2222-3333-4444-555555555555"
async def _create_secret_body_for_settings(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter, settings: KeyManagementSettings
) -> dict[str, object]:
"""Boot the manager from settings the way the proxy does and return the CreateSecret body it posts to AWS."""
monkeypatch.setattr(litellm, "secret_manager_client", None)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", None)
monkeypatch.setenv("AWS_REGION_NAME", "us-east-1")
AWSSecretsManagerV2.load_aws_secret_manager(use_aws_secret_manager=True, key_management_settings=settings)
manager = litellm.secret_manager_client
assert isinstance(manager, AWSSecretsManagerV2)
route = respx_mock.post("https://secretsmanager.us-east-1.amazonaws.com/").respond(
json={"ARN": "arn", "Name": "litellm/test-key"}
)
await manager.async_write_secret(
secret_name="litellm/test-key",
secret_value="sk-test-value",
optional_params=dict(_STATIC_CREDENTIALS),
)
assert route.call_count == 1
request = route.calls.last.request
assert request.headers["X-Amz-Target"] == "secretsmanager.CreateSecret"
return json.loads(request.content)
@pytest.mark.asyncio
async def test_create_secret_uses_customer_managed_kms_key_from_settings(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
body = await _create_secret_body_for_settings(
monkeypatch,
respx_mock,
KeyManagementSettings(store_virtual_keys=True, aws_region_name="us-east-1", kms_key_id=_CMK_ARN),
)
assert body["KmsKeyId"] == _CMK_ARN
assert body["Name"] == "litellm/test-key"
assert body["SecretString"] == "sk-test-value"
@pytest.mark.asyncio
async def test_create_secret_omits_kms_key_id_when_not_configured(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
body = await _create_secret_body_for_settings(
monkeypatch, respx_mock, KeyManagementSettings(store_virtual_keys=True, aws_region_name="us-east-1")
)
assert "KmsKeyId" not in body
assert body["Name"] == "litellm/test-key"
@pytest.mark.asyncio

View file

@ -855,3 +855,242 @@ class TestAsyncPostCallSuccessHook:
)
assert result == mock_response
_FABRICATED_PROVIDER_RESPONSE_ID = "resp_fabricatedprovideridaaaaaaaaaaaaaaaa"
_FABRICATED_UNMANAGED_ID = "resp_fabricatedunmanagedidbbbbbbbbbbbbbbbb"
_UNIT_TEST_SALT_KEY = "lit6837-unit-test-salt-key"
_ADDRESSED_ID_FIELD_BY_CALL_TYPE = {
"aresponses": "previous_response_id",
"aget_responses": "response_id",
"adelete_responses": "response_id",
"acancel_responses": "response_id",
"alist_input_items": "response_id",
}
@pytest.fixture
def salt_key_env(monkeypatch):
"""Give the encrypt/decrypt helpers a real salt key so ids round-trip for real."""
monkeypatch.setenv("LITELLM_SALT_KEY", _UNIT_TEST_SALT_KEY)
return _UNIT_TEST_SALT_KEY
def _hook(general_settings=None, signing_key=_UNIT_TEST_SALT_KEY):
settings = general_settings if general_settings is not None else {}
return ResponsesIDSecurity(
general_settings_reader=lambda: settings,
signing_key_reader=lambda: signing_key,
)
def _auth(user_id="owner-user", team_id="owner-team", user_role=None):
from litellm.proxy._types import UserAPIKeyAuth
return UserAPIKeyAuth(user_id=user_id, team_id=team_id, user_role=user_role)
def _issue_managed_id(hook, owner, provider_response_id=_FABRICATED_PROVIDER_RESPONSE_ID):
"""Mint an id exactly the way the proxy hands one to a client on create."""
issued = hook._encrypt_response_id(
ResponsesAPIResponse(
id=provider_response_id, created_at=1234567890, output=[], status="completed"
),
owner,
)
return issued.id
class TestUnrecognizedResponseIdIsRejected:
"""An id this proxy never issued carries no owner, so it must not reach the provider."""
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type", sorted(_ADDRESSED_ID_FIELD_BY_CALL_TYPE))
async def test_unmanaged_id_is_rejected_and_not_forwarded(self, mock_cache, salt_key_env, call_type):
field = _ADDRESSED_ID_FIELD_BY_CALL_TYPE[call_type]
data = {field: _FABRICATED_UNMANAGED_ID}
with pytest.raises(HTTPException) as exc_info:
await _hook().async_pre_call_hook(
user_api_key_dict=_auth(),
cache=mock_cache,
data=data,
call_type=call_type,
)
assert exc_info.value.status_code == 403
assert "allow_unmanaged_response_ids" in exc_info.value.detail
assert data[field] == _FABRICATED_UNMANAGED_ID
@pytest.mark.asyncio
async def test_owner_can_still_address_the_id_the_proxy_issued_it(self, mock_cache, salt_key_env):
hook = _hook()
owner = _auth()
data = {"response_id": _issue_managed_id(hook, owner)}
result = await hook.async_pre_call_hook(
user_api_key_dict=owner,
cache=mock_cache,
data=data,
call_type="aget_responses",
)
assert result["response_id"] == _FABRICATED_PROVIDER_RESPONSE_ID
@pytest.mark.asyncio
async def test_stranger_cannot_address_an_id_issued_to_someone_else(self, mock_cache, salt_key_env):
hook = _hook()
issued_id = _issue_managed_id(hook, _auth())
data = {"response_id": issued_id}
with pytest.raises(HTTPException) as exc_info:
await hook.async_pre_call_hook(
user_api_key_dict=_auth(user_id="stranger-user", team_id="stranger-team"),
cache=mock_cache,
data=data,
call_type="aget_responses",
)
assert exc_info.value.status_code == 403
assert data["response_id"] == issued_id
@pytest.mark.asyncio
async def test_unmanaged_previous_response_id_cannot_seed_a_new_response(self, mock_cache, salt_key_env):
data = {"model": "gpt-fake", "previous_response_id": _FABRICATED_UNMANAGED_ID}
with pytest.raises(HTTPException) as exc_info:
await _hook().async_pre_call_hook(
user_api_key_dict=_auth(),
cache=mock_cache,
data=data,
call_type="aresponses",
)
assert exc_info.value.status_code == 403
assert data["previous_response_id"] == _FABRICATED_UNMANAGED_ID
@pytest.mark.asyncio
async def test_re_entering_the_hook_on_the_same_request_does_not_reject(self, mock_cache, salt_key_env):
"""The rate-limit fallback retry runs pre-call twice over one already-rewritten dict."""
hook = _hook()
owner = _auth()
data = {"model": "gpt-fake", "previous_response_id": _issue_managed_id(hook, owner)}
first = await hook.async_pre_call_hook(
user_api_key_dict=owner, cache=mock_cache, data=data, call_type="aresponses"
)
second = await hook.async_pre_call_hook(
user_api_key_dict=owner, cache=mock_cache, data=first, call_type="aresponses"
)
assert second["previous_response_id"] == _FABRICATED_PROVIDER_RESPONSE_ID
class TestUnmanagedResponseIdEscapeHatches:
"""Deployments that pass provider ids through on purpose must keep working."""
@pytest.mark.asyncio
@pytest.mark.parametrize(
"general_settings",
[{"allow_unmanaged_response_ids": True}, {"disable_responses_id_security": True}],
)
async def test_opted_in_settings_forward_the_id_untouched(self, mock_cache, salt_key_env, general_settings):
data = {"response_id": _FABRICATED_UNMANAGED_ID}
result = await _hook(general_settings=general_settings).async_pre_call_hook(
user_api_key_dict=_auth(),
cache=mock_cache,
data=data,
call_type="aget_responses",
)
assert result["response_id"] == _FABRICATED_UNMANAGED_ID
@pytest.mark.asyncio
async def test_proxy_without_a_signing_key_forwards_the_id_untouched(self, mock_cache, monkeypatch):
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
data = {"response_id": _FABRICATED_UNMANAGED_ID}
result = await _hook(signing_key=None).async_pre_call_hook(
user_api_key_dict=_auth(),
cache=mock_cache,
data=data,
call_type="aget_responses",
)
assert result["response_id"] == _FABRICATED_UNMANAGED_ID
@pytest.mark.asyncio
async def test_proxy_admin_may_address_an_unmanaged_id(self, mock_cache, salt_key_env):
from litellm.proxy._types import LitellmUserRoles
data = {"response_id": _FABRICATED_UNMANAGED_ID}
result = await _hook().async_pre_call_hook(
user_api_key_dict=_auth(user_role=LitellmUserRoles.PROXY_ADMIN),
cache=mock_cache,
data=data,
call_type="aget_responses",
)
assert result["response_id"] == _FABRICATED_UNMANAGED_ID
class TestClientSuppliedRetainedIdCannotBypassAuthorization:
"""The retained-id key travels in the request body, so it is re-authorized, never trusted."""
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type", sorted(_ADDRESSED_ID_FIELD_BY_CALL_TYPE))
async def test_forged_retained_id_is_still_authorized(self, mock_cache, salt_key_env, call_type):
field = _ADDRESSED_ID_FIELD_BY_CALL_TYPE[call_type]
data = {
field: _FABRICATED_UNMANAGED_ID,
"_litellm_addressed_response_id": _FABRICATED_UNMANAGED_ID,
}
with pytest.raises(HTTPException) as exc_info:
await _hook().async_pre_call_hook(
user_api_key_dict=_auth(),
cache=mock_cache,
data=data,
call_type=call_type,
)
assert exc_info.value.status_code == 403
assert data[field] == _FABRICATED_UNMANAGED_ID
@pytest.mark.asyncio
@pytest.mark.parametrize("forged", [{"nested": "value"}, ["list"], 42, "", None])
async def test_non_string_retained_id_falls_back_to_the_addressed_field(self, mock_cache, salt_key_env, forged):
data = {"response_id": _FABRICATED_UNMANAGED_ID, "_litellm_addressed_response_id": forged}
with pytest.raises(HTTPException) as exc_info:
await _hook().async_pre_call_hook(
user_api_key_dict=_auth(),
cache=mock_cache,
data=data,
call_type="aget_responses",
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_stranger_forging_their_own_id_never_reaches_someone_elses_response(
self, mock_cache, salt_key_env
):
hook = _hook()
stranger = _auth(user_id="stranger-user", team_id="stranger-team")
stranger_id = _issue_managed_id(hook, stranger, provider_response_id="resp_strangerownprovideridcccccccc")
victim_provider_id = "resp_victimprovideriddddddddddddddddddddd"
data = {"response_id": victim_provider_id, "_litellm_addressed_response_id": stranger_id}
result = await hook.async_pre_call_hook(
user_api_key_dict=stranger,
cache=mock_cache,
data=data,
call_type="aget_responses",
)
assert result["response_id"] == "resp_strangerownprovideridcccccccc"
assert result["response_id"] != victim_provider_id

View file

@ -7381,6 +7381,63 @@ async def test_async_get_fully_unhealthy_model_names_marks_name_when_all_unhealt
assert await router.async_get_fully_unhealthy_model_names() == {"gpt-4o"}
@pytest.mark.asyncio
@pytest.mark.parametrize("health_check_probe", [False, True])
@pytest.mark.parametrize(
"state, health_routing, fails_policy, scoped, strict_ids",
[
("absent", True, False, False, ("dep-0", "dep-1")),
("partial", True, False, False, ("dep-1",)),
("all", True, False, False, ()),
("stale", True, False, False, ("dep-0", "dep-1")),
("all", False, False, False, ("dep-0", "dep-1")),
("all", True, True, False, ("dep-0", "dep-1")),
("all", True, True, True, ()),
],
)
async def test_health_probe_preserves_normal_caller_policy(
health_check_probe: bool,
state: str,
health_routing: bool,
fails_policy: bool,
scoped: bool,
strict_ids: tuple[str, ...],
) -> None:
import time
from litellm.types.router import AllowedFailsPolicy, RouterRateLimitError
router: Final = Router(
model_list=[
{
"model_name": "health-group",
"litellm_params": {"model": "openai/gpt-5.6", "api_key": "test-only"},
"model_info": {"id": model_id},
}
for model_id in ("dep-0", "dep-1")
],
enable_health_check_routing=health_routing,
allowed_fails_policy=AllowedFailsPolicy(ServiceUnavailableErrorAllowedFails=2) if fails_policy else None,
background_health_check_model_groups=["health-group"] if scoped else None,
)
if state != "absent":
_seed_unhealthy_states(
router,
("dep-0",) if state == "partial" else ("dep-0", "dep-1"),
time.time() - router.health_state_cache.staleness_threshold - 10 if state == "stale" else None,
)
expected: Final = strict_ids if strict_ids or health_check_probe else ("dep-0", "dep-1")
if not expected:
with pytest.raises(RouterRateLimitError, match="No deployments available"):
await router.async_get_healthy_deployments(model="health-group", request_kwargs={}, health_check_probe=True)
else:
deployments: Final = await router.async_get_healthy_deployments(
model="health-group", request_kwargs={}, health_check_probe=health_check_probe
)
assert {d["model_info"]["id"] for d in deployments} == set(expected)
assert await router.cooldown_cache.async_get_active_cooldowns(["dep-0", "dep-1"], parent_otel_span=None) == []
@pytest.mark.asyncio
async def test_async_get_fully_unhealthy_model_names_keeps_name_when_partial():
router = _router_with_two_deployments([False, False])

View file

@ -9547,26 +9547,6 @@ export interface paths {
patch?: never;
trace?: never;
};
"/openai/": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
/**
* WebSocket: openai_websocket_proxy_route
* @description WebSocket connection endpoint
*/
get: operations["websocket_openai_websocket_proxy_route_get"];
put?: never;
post?: never;
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/openai/deployments/{model}/chat/completions": {
parameters: {
query?: never;
@ -10122,26 +10102,6 @@ export interface paths {
patch: operations["openai_proxy_route_openai__endpoint__patch"];
trace?: never;
};
"/openai_passthrough/": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
/**
* WebSocket: openai_websocket_proxy_route
* @description WebSocket connection endpoint
*/
get: operations["websocket_openai_websocket_proxy_route_get_2"];
put?: never;
post?: never;
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/openai_passthrough/{endpoint}": {
parameters: {
query?: never;
@ -10150,132 +10110,72 @@ export interface paths {
cookie?: never;
};
/**
* Openai Proxy Route
* @description Pass-through endpoint for OpenAI API calls.
*
* Available on both routes:
* - /openai/{endpoint:path} - Standard OpenAI passthrough route
* - /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API)
*
* Use /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts
* with LiteLLM's native implementations (e.g., for the Responses API at /v1/responses).
* Openai Passthrough Route
* @description Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native
* implementations (e.g. the Responses API at /v1/responses).
*
* Examples:
* Standard route:
* - /openai/v1/chat/completions
* - /openai/v1/assistants
* - /openai/v1/threads
*
* Dedicated passthrough (for Responses API):
* - /openai_passthrough/v1/responses
* - /openai_passthrough/v1/responses/{response_id}
* - /openai_passthrough/v1/responses/{response_id}/input_items
*
* [Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)
*/
get: operations["openai_proxy_route_openai_passthrough__endpoint__get"];
get: operations["openai_passthrough_route_openai_passthrough__endpoint__get"];
/**
* Openai Proxy Route
* @description Pass-through endpoint for OpenAI API calls.
*
* Available on both routes:
* - /openai/{endpoint:path} - Standard OpenAI passthrough route
* - /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API)
*
* Use /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts
* with LiteLLM's native implementations (e.g., for the Responses API at /v1/responses).
* Openai Passthrough Route
* @description Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native
* implementations (e.g. the Responses API at /v1/responses).
*
* Examples:
* Standard route:
* - /openai/v1/chat/completions
* - /openai/v1/assistants
* - /openai/v1/threads
*
* Dedicated passthrough (for Responses API):
* - /openai_passthrough/v1/responses
* - /openai_passthrough/v1/responses/{response_id}
* - /openai_passthrough/v1/responses/{response_id}/input_items
*
* [Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)
*/
put: operations["openai_proxy_route_openai_passthrough__endpoint__put"];
put: operations["openai_passthrough_route_openai_passthrough__endpoint__put"];
/**
* Openai Proxy Route
* @description Pass-through endpoint for OpenAI API calls.
*
* Available on both routes:
* - /openai/{endpoint:path} - Standard OpenAI passthrough route
* - /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API)
*
* Use /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts
* with LiteLLM's native implementations (e.g., for the Responses API at /v1/responses).
* Openai Passthrough Route
* @description Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native
* implementations (e.g. the Responses API at /v1/responses).
*
* Examples:
* Standard route:
* - /openai/v1/chat/completions
* - /openai/v1/assistants
* - /openai/v1/threads
*
* Dedicated passthrough (for Responses API):
* - /openai_passthrough/v1/responses
* - /openai_passthrough/v1/responses/{response_id}
* - /openai_passthrough/v1/responses/{response_id}/input_items
*
* [Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)
*/
post: operations["openai_proxy_route_openai_passthrough__endpoint__post"];
post: operations["openai_passthrough_route_openai_passthrough__endpoint__post"];
/**
* Openai Proxy Route
* @description Pass-through endpoint for OpenAI API calls.
*
* Available on both routes:
* - /openai/{endpoint:path} - Standard OpenAI passthrough route
* - /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API)
*
* Use /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts
* with LiteLLM's native implementations (e.g., for the Responses API at /v1/responses).
* Openai Passthrough Route
* @description Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native
* implementations (e.g. the Responses API at /v1/responses).
*
* Examples:
* Standard route:
* - /openai/v1/chat/completions
* - /openai/v1/assistants
* - /openai/v1/threads
*
* Dedicated passthrough (for Responses API):
* - /openai_passthrough/v1/responses
* - /openai_passthrough/v1/responses/{response_id}
* - /openai_passthrough/v1/responses/{response_id}/input_items
*
* [Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)
*/
delete: operations["openai_proxy_route_openai_passthrough__endpoint__delete"];
delete: operations["openai_passthrough_route_openai_passthrough__endpoint__delete"];
options?: never;
head?: never;
/**
* Openai Proxy Route
* @description Pass-through endpoint for OpenAI API calls.
*
* Available on both routes:
* - /openai/{endpoint:path} - Standard OpenAI passthrough route
* - /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API)
*
* Use /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts
* with LiteLLM's native implementations (e.g., for the Responses API at /v1/responses).
* Openai Passthrough Route
* @description Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native
* implementations (e.g. the Responses API at /v1/responses).
*
* Examples:
* Standard route:
* - /openai/v1/chat/completions
* - /openai/v1/assistants
* - /openai/v1/threads
*
* Dedicated passthrough (for Responses API):
* - /openai_passthrough/v1/responses
* - /openai_passthrough/v1/responses/{response_id}
* - /openai_passthrough/v1/responses/{response_id}/input_items
*
* [Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)
*/
patch: operations["openai_proxy_route_openai_passthrough__endpoint__patch"];
patch: operations["openai_passthrough_route_openai_passthrough__endpoint__patch"];
trace?: never;
};
"/organization/daily/activity": {
@ -21891,52 +21791,6 @@ export interface paths {
patch?: never;
trace?: never;
};
"/vertex-ai/{endpoint}": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
/**
* Vertex Proxy Route
* @description Call LiteLLM proxy via Vertex AI SDK.
*
* [Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai)
*/
get: operations["vertex_proxy_route_vertex_ai__endpoint__get_2"];
/**
* Vertex Proxy Route
* @description Call LiteLLM proxy via Vertex AI SDK.
*
* [Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai)
*/
put: operations["vertex_proxy_route_vertex_ai__endpoint__put_2"];
/**
* Vertex Proxy Route
* @description Call LiteLLM proxy via Vertex AI SDK.
*
* [Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai)
*/
post: operations["vertex_proxy_route_vertex_ai__endpoint__post_2"];
/**
* Vertex Proxy Route
* @description Call LiteLLM proxy via Vertex AI SDK.
*
* [Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai)
*/
delete: operations["vertex_proxy_route_vertex_ai__endpoint__delete_2"];
options?: never;
head?: never;
/**
* Vertex Proxy Route
* @description Call LiteLLM proxy via Vertex AI SDK.
*
* [Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai)
*/
patch: operations["vertex_proxy_route_vertex_ai__endpoint__patch_2"];
trace?: never;
};
"/vertex_ai/discovery/{endpoint}": {
parameters: {
query?: never;
@ -25775,6 +25629,11 @@ export interface components {
* @description opt-in to RFC 8628 verification_uri_complete for the CLI SSO device flow, pre-filling the user_code in the browser. Off by default; intended for same-host clients where the device that starts the flow and the browser run on the same machine
*/
allow_cli_sso_verification_uri_complete?: boolean | null;
/**
* Allow Unmanaged Response Ids
* @description If True, lets keys address Responses API ids that this proxy did not issue (raw provider ids, or ids issued before response-id encryption was configured). Such an id carries no owner, so no ownership check can run on it; ids this proxy did issue keep full ownership enforcement. Off by default, in which case an unrecognized response id is rejected with 403
*/
allow_unmanaged_response_ids?: boolean | null;
/**
* Allowed Routes
* @description Proxy API Endpoints you want users to be able to access
@ -25894,6 +25753,11 @@ export interface components {
* @description If True and SSO is configured (MICROSOFT_CLIENT_ID, GOOGLE_CLIENT_ID, GENERIC_CLIENT_ID, or SAML_IDP_METADATA_URL/XML), disables username/password login on /login, /v2/login, and /v3/login so SSO is the only way to reach the Admin UI. An admin locked out of the UI can still administer the proxy over the API with the master key; unset this setting and restart the proxy to restore UI username/password login. Default is False.
*/
disable_password_login_when_sso_enabled?: boolean | null;
/**
* Disable Responses Id Security
* @description If True, disables ownership enforcement on Responses API ids. Keys may then retrieve, cancel, delete, and chain from any response id, including ids belonging to another user or team and ids this proxy never issued. WARNING: this removes tenant isolation on /v1/responses
*/
disable_responses_id_security?: boolean | null;
/**
* Enable Openai Websocket Passthrough
* @description Serve the OpenAI pass-through WebSocket route, which relays frames to OpenAI under the proxy's own provider credential without reading them. Off by default.
@ -26153,6 +26017,11 @@ export interface components {
* @description If True and LiteLLM_SpendLogs has been converted to a range-partitioned table (db_scripts/partition_spend_logs.sql), retention cleanup drops expired partitions instead of deleting rows, and pre-creates upcoming partitions. Default is False.
*/
use_spend_logs_partitioning?: boolean | null;
/**
* User Api Key Cache Max Size
* @description max number of entries (virtual keys, teams, users, end users, memberships, ...) each worker keeps in its in-memory auth cache. Defaults to 200. Raise this if you have more active keys than that or auth lookups keep hitting the DB
*/
user_api_key_cache_max_size?: number | null;
/** User Header Mappings */
user_header_mappings?: components["schemas"]["UserHeaderMapping"][] | null;
/**
@ -36460,7 +36329,7 @@ export interface components {
* Cause
* @enum {string}
*/
cause?: "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "heuristic_first_short_circuit" | "hybrid_short_circuit" | "classifier_plugin" | "classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "modality_pin_override" | "health_failover" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit";
cause?: "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "heuristic_first_short_circuit" | "hybrid_short_circuit" | "classifier_plugin" | "classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "modality_pin_override" | "health_failover" | "health_default_fallback" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit";
/** Classifier Cost */
classifier_cost?: number;
/** Classifier Model */
@ -52513,24 +52382,6 @@ export interface operations {
};
};
};
websocket_openai_websocket_proxy_route_get: {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
requestBody?: never;
responses: {
/** @description WebSocket Protocol Switched */
101: {
headers: {
[name: string]: unknown;
};
content?: never;
};
};
};
chat_completion_openai_deployments__model__chat_completions_post: {
parameters: {
query?: never;
@ -53430,25 +53281,7 @@ export interface operations {
};
};
};
websocket_openai_websocket_proxy_route_get_2: {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
requestBody?: never;
responses: {
/** @description WebSocket Protocol Switched */
101: {
headers: {
[name: string]: unknown;
};
content?: never;
};
};
};
openai_proxy_route_openai_passthrough__endpoint__get: {
openai_passthrough_route_openai_passthrough__endpoint__get: {
parameters: {
query?: never;
header?: never;
@ -53479,7 +53312,7 @@ export interface operations {
};
};
};
openai_proxy_route_openai_passthrough__endpoint__put: {
openai_passthrough_route_openai_passthrough__endpoint__put: {
parameters: {
query?: never;
header?: never;
@ -53510,7 +53343,7 @@ export interface operations {
};
};
};
openai_proxy_route_openai_passthrough__endpoint__post: {
openai_passthrough_route_openai_passthrough__endpoint__post: {
parameters: {
query?: never;
header?: never;
@ -53541,7 +53374,7 @@ export interface operations {
};
};
};
openai_proxy_route_openai_passthrough__endpoint__delete: {
openai_passthrough_route_openai_passthrough__endpoint__delete: {
parameters: {
query?: never;
header?: never;
@ -53572,7 +53405,7 @@ export interface operations {
};
};
};
openai_proxy_route_openai_passthrough__endpoint__patch: {
openai_passthrough_route_openai_passthrough__endpoint__patch: {
parameters: {
query?: never;
header?: never;
@ -67979,161 +67812,6 @@ export interface operations {
};
};
};
vertex_proxy_route_vertex_ai__endpoint__get_2: {
parameters: {
query?: never;
header?: never;
path: {
endpoint: string;
};
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": unknown;
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
vertex_proxy_route_vertex_ai__endpoint__put_2: {
parameters: {
query?: never;
header?: never;
path: {
endpoint: string;
};
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": unknown;
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
vertex_proxy_route_vertex_ai__endpoint__post_2: {
parameters: {
query?: never;
header?: never;
path: {
endpoint: string;
};
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": unknown;
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
vertex_proxy_route_vertex_ai__endpoint__delete_2: {
parameters: {
query?: never;
header?: never;
path: {
endpoint: string;
};
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": unknown;
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
vertex_proxy_route_vertex_ai__endpoint__patch_2: {
parameters: {
query?: never;
header?: never;
path: {
endpoint: string;
};
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": unknown;
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
vertex_discovery_proxy_route_vertex_ai_discovery__endpoint__get: {
parameters: {
query?: never;