mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_bedrock_openai_xhigh_flags
This commit is contained in:
commit
7b25c6a29e
71 changed files with 9673 additions and 933 deletions
5
.github/ci-coverage-allowlist.yml
vendored
5
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -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
73
.github/workflows/ai-gateway-image.yml
vendored
Normal 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
|
||||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__ = [
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 _:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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=(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]] = {}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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"),
|
||||
|
|
|
|||
24
litellm/proxy/route_priority.py
Normal file
24
litellm/proxy/route_priority.py
Normal 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))
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
171
tests/test_litellm/proxy/test_route_priority.py
Normal file
171
tests/test_litellm/proxy/test_route_priority.py
Normal 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
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
404
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
404
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue