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

This commit is contained in:
Devin AI 2026-07-18 03:39:29 +00:00
commit cd63b40255
156 changed files with 9789 additions and 1621 deletions

3
.gitignore vendored
View file

@ -15,6 +15,9 @@ litellm/rust_bridge/_native*.so
litellm/rust_bridge/_native*.pyd
litellm-rust/target/
# Python package build output
dist/
bun.lockb
**/.DS_Store
.aider*

Binary file not shown.

View file

@ -28,12 +28,6 @@ version = "1.1.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0"
[[package]]
name = "autocfg"
version = "1.5.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53"
[[package]]
name = "axum"
version = "0.7.9"
@ -392,21 +386,6 @@ dependencies = [
"percent-encoding",
]
[[package]]
name = "futures"
version = "0.3.32"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8b147ee9d1f6d097cef9ce628cd2ee62288d963e16fb287bd9286455b241382d"
dependencies = [
"futures-channel",
"futures-core",
"futures-executor",
"futures-io",
"futures-sink",
"futures-task",
"futures-util",
]
[[package]]
name = "futures-channel"
version = "0.3.32"
@ -423,17 +402,6 @@ version = "0.3.32"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d"
[[package]]
name = "futures-executor"
version = "0.3.32"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "baf29c38818342a3b26b5b923639e7b1f4a61fc5e76102d4b1981c6dc7a7579d"
dependencies = [
"futures-core",
"futures-task",
"futures-util",
]
[[package]]
name = "futures-io"
version = "0.3.32"
@ -469,7 +437,6 @@ version = "0.3.32"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6"
dependencies = [
"futures-channel",
"futures-core",
"futures-io",
"futures-macro",
@ -928,15 +895,6 @@ dependencies = [
"serde_core",
]
[[package]]
name = "indoc"
version = "2.0.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "79cf5c93f93228cf8efb3ba362535fb11199ac548a09ce117c9b1adc3030d706"
dependencies = [
"rustversion",
]
[[package]]
name = "ipnet"
version = "2.12.0"
@ -1089,15 +1047,6 @@ version = "2.8.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "88904434abc2901f197fe8cc55f0445e7ded921dba5911dad2e2b39b48e663c4"
[[package]]
name = "memoffset"
version = "0.9.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "488016bfae457b036d996092f6cb448677611ce4449e970ceaf42695203f218a"
dependencies = [
"autocfg",
]
[[package]]
name = "mime"
version = "0.3.17"
@ -1225,29 +1174,26 @@ dependencies = [
[[package]]
name = "pyo3"
version = "0.23.5"
version = "0.29.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7778bffd85cf38175ac1f545509665d0b9b92a198ca7941f131f85f7a4f9a872"
checksum = "cd274650b21d4bfc26a0a47587962c1edb425f69287324355cd040c3ea66071c"
dependencies = [
"cfg-if",
"indoc",
"libc",
"memoffset",
"once_cell",
"portable-atomic",
"pyo3-build-config",
"pyo3-ffi",
"pyo3-macros",
"unindent",
]
[[package]]
name = "pyo3-async-runtimes"
version = "0.23.0"
version = "0.29.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "977dc837525cfd22919ba6a831413854beb7c99a256c03bf8624ad707e45810e"
checksum = "b3ef68daa7316a3fac65e5e18b2203f010346de1c1c53456811a2624673ab046"
dependencies = [
"futures",
"futures-channel",
"futures-util",
"once_cell",
"pin-project-lite",
"pyo3",
@ -1256,19 +1202,18 @@ dependencies = [
[[package]]
name = "pyo3-build-config"
version = "0.23.5"
version = "0.29.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "94f6cbe86ef3bf18998d9df6e0f3fc1050a8c5efa409bf712e661a4366e010fb"
checksum = "c5e2a7d2f0d013342f295c048ad19237add5154a55b1c5a254c0ec93d4109078"
dependencies = [
"once_cell",
"target-lexicon",
]
[[package]]
name = "pyo3-ffi"
version = "0.23.5"
version = "0.29.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e9f1b4c431c0bb1c8fb0a338709859eed0d030ff6daa34368d3b152a63dfdd8d"
checksum = "ca85c467da1bbc8d866eea5deff9cf29ea5f7785054a17da36e65bda9c05845b"
dependencies = [
"libc",
"pyo3-build-config",
@ -1276,9 +1221,9 @@ dependencies = [
[[package]]
name = "pyo3-macros"
version = "0.23.5"
version = "0.29.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fbc2201328f63c4710f68abdf653c89d8dbc2858b88c5d88b0ff38a75288a9da"
checksum = "9ac53762fd065daa3194dd09337a38bd793a188100fd1a9304c4ab312d901771"
dependencies = [
"proc-macro2",
"pyo3-macros-backend",
@ -1288,13 +1233,12 @@ dependencies = [
[[package]]
name = "pyo3-macros-backend"
version = "0.23.5"
version = "0.29.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fca6726ad0f3da9c9de093d6f116a93c1a38e417ed73bf138472cf4064f72028"
checksum = "4ca3a1557399783172dc5bf39cfca835157732532cba56b71d2292161e53b362"
dependencies = [
"heck",
"proc-macro2",
"pyo3-build-config",
"quote",
"syn",
]
@ -1974,9 +1918,9 @@ dependencies = [
[[package]]
name = "target-lexicon"
version = "0.12.16"
version = "0.13.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "61c41af27dd6d1e27b1b16b489db798443478cef1f06a660c96db617ba5de3b1"
checksum = "adb6935a6f5c20170eeceb1a3835a49e12e19d792f6dd344ccc76a985ca5a6ca"
[[package]]
name = "thiserror"
@ -2250,12 +2194,6 @@ version = "1.0.24"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75"
[[package]]
name = "unindent"
version = "0.2.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7264e107f553ccae879d21fbea1d6724ac785e8c3bfc762137959b5802826ef3"
[[package]]
name = "untrusted"
version = "0.9.0"

View file

@ -16,8 +16,8 @@ rust-version = "1.86"
litellm-core = { path = "crates/core" }
litellm-ai-gateway = { path = "crates/ai-gateway", default-features = false }
axum = "0.7"
pyo3 = "0.23.5"
pyo3-async-runtimes = { version = "0.23.0", features = ["tokio-runtime"] }
pyo3 = "0.29.0"
pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
rand = "0.8"
reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "multipart", "rustls-tls", "http2", "stream"] }
serde = { version = "1.0", features = ["derive"] }

View file

@ -17,7 +17,7 @@ use crate::gil;
/// Load the router's `model_list` from `config_path` via the Python reader.
pub fn load_router_from_config(config_path: &str) -> CoreResult<Router> {
gil::record_acquisition();
Python::with_gil(|py| {
Python::attach(|py| {
let model_list = py
.import("litellm.proxy.read_model_list")
.and_then(|module| module.getattr("read_model_list"))

View file

@ -2,7 +2,7 @@
//!
//! A single chokepoint for releasing the GIL around blocking work. Every
//! blocking call in the bridge goes through [`release_gil`] instead of calling
//! `Python::allow_threads` directly, so the release count stays accurate and we
//! `Python::detach` directly, so the release count stays accurate and we
//! have one place to extend later (timing histograms, per-call labels, etc.).
use std::sync::atomic::{AtomicU64, Ordering};
@ -23,7 +23,7 @@ where
T: Send,
{
GIL_RELEASES.fetch_add(1, Ordering::Relaxed);
py.allow_threads(f)
py.detach(f)
}
/// Total GIL releases performed by the bridge so far.

View file

@ -174,7 +174,7 @@ fn aocr(
.await
.map_err(|err| Python::with_gil(|py| core_error_to_pyerr(py, err)))?;
Python::with_gil(|py| json_to_py(py, value))
Python::attach(|py| json_to_py(py, value))
})
}

View file

@ -315,6 +315,11 @@ disable_token_counter: bool = False
disable_add_transform_inline_image_block: bool = False
disable_add_user_agent_to_request_tags: bool = False
disable_anthropic_gemini_context_caching_transform: bool = False
enable_anthropic_prompt_caching: bool = os.getenv("LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING", "false").lower() == "true"
_anthropic_prompt_caching_ttl_env: Optional[str] = os.getenv("LITELLM_ANTHROPIC_PROMPT_CACHING_TTL")
anthropic_prompt_caching_ttl: Optional[Literal["5m", "1h"]] = (
"1h" if _anthropic_prompt_caching_ttl_env == "1h" else "5m" if _anthropic_prompt_caching_ttl_env == "5m" else None
)
disable_vertex_batch_output_transformation: bool = False
extra_spend_tag_headers: Optional[List[str]] = None
in_memory_llm_clients_cache: "LLMClientCache"

View file

@ -296,18 +296,148 @@ class AnthropicCacheControlHook(CustomPromptManagement):
return processed_messages, processed_system, remaining_points
@staticmethod
def _default_control() -> ChatCompletionCachedContent:
"""Build the cache_control block for auto-injected breakpoints.
Defaults to Anthropic's 5-minute ephemeral cache; honors the optional
``litellm.anthropic_prompt_caching_ttl`` override ("5m" or "1h").
"""
import litellm
ttl = litellm.anthropic_prompt_caching_ttl
if ttl == "5m" or ttl == "1h":
return ChatCompletionCachedContent(type="ephemeral", ttl=ttl)
return ChatCompletionCachedContent(type="ephemeral")
@staticmethod
def _request_has_cache_control(
messages: list[AllMessageValues],
system: str | list | None,
tools: list | None = None,
) -> bool:
"""Return True if the request already carries any client-supplied cache_control.
When the client (e.g. Claude Code) already marks its own breakpoints we
stand down entirely rather than add more, per the auto-caching contract.
Tools count: they are a breakpoint the client can mark, they count toward
the provider's four-block limit, and caching only the tool definitions is
a common pattern, so injecting alongside them can exceed the cap.
"""
if any(AnthropicCacheControlHook._count_cache_control_blocks(msg) for msg in messages):
return True
if isinstance(system, list):
if any(isinstance(block, dict) and block.get("cache_control") is not None for block in system):
return True
if tools is not None:
return any(isinstance(tool, dict) and tool.get("cache_control") is not None for tool in tools)
return False
@staticmethod
def get_default_injection_points(
messages: list[AllMessageValues],
system: str | list | None,
model: str,
custom_llm_provider: str | None,
tools: list | None = None,
) -> list[CacheControlInjectionPoint]:
"""Default breakpoints when ``litellm.enable_anthropic_prompt_caching`` is on.
Caches the system prompt and the trailing turn, so the stable prefix
(system + tools + history) is reused while the breakpoint advances with
the conversation. Returns [] (stand down) when the flag is off, the
provider does not consume cache_control breakpoints (only anthropic /
bedrock do), the model lacks prompt-caching support, or the request
already carries client-supplied cache_control.
"""
import litellm
if litellm.enable_anthropic_prompt_caching is not True:
return []
provider = custom_llm_provider
if provider is None:
from litellm.litellm_core_utils.get_llm_provider_logic import (
get_llm_provider,
)
try:
_, provider, _, _ = get_llm_provider(model=model)
except Exception: # noqa: BLE001 # unroutable model must never block the call, just skip auto-caching
return []
if provider not in ("anthropic", "bedrock"):
return []
from litellm.utils import supports_prompt_caching
if not supports_prompt_caching(model=model, custom_llm_provider=provider):
return []
if AnthropicCacheControlHook._request_has_cache_control(messages, system, tools):
return []
control = AnthropicCacheControlHook._default_control()
points: list[CacheControlInjectionPoint] = [
CacheControlMessageInjectionPoint(location="message", role="system", index=None, control=control),
CacheControlMessageInjectionPoint(location="message", role=None, index=-1, control=control),
]
return points
@staticmethod
def maybe_seed_default_injection_points(
non_default_params: dict[str, Any],
messages: list[AllMessageValues],
model: str,
custom_llm_provider: str | None,
tools: list | None = None,
) -> None:
"""For /chat/completions: add default injection points to the request params.
No-op when injection points are already configured (explicit config wins).
Seeding the param lets the existing prompt-management gate and the
AnthropicCacheControlHook run unchanged.
"""
if non_default_params.get("cache_control_injection_points"):
return
points = AnthropicCacheControlHook.get_default_injection_points(
messages=messages,
system=None,
model=model,
custom_llm_provider=custom_llm_provider,
tools=tools,
)
if points:
non_default_params["cache_control_injection_points"] = points
@staticmethod
def maybe_inject_cache_control(
messages: List[Dict],
system: str | list | None,
kwargs: Dict[str, Any],
model: str | None = None,
custom_llm_provider: str | None = None,
tools: list[dict] | None = None,
) -> Tuple[List[Dict], str | list | None]:
"""Extract cache_control_injection_points from kwargs and apply if present.
When none are configured but ``litellm.enable_anthropic_prompt_caching``
is on, synthesize default breakpoints for the native /v1/messages path.
Pops the key from kwargs; if remaining (non-message) points exist they
are written back so downstream transforms can handle them.
"""
injection_points = kwargs.pop("cache_control_injection_points", None)
configured = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list
list[CacheControlInjectionPoint] | None, kwargs.pop("cache_control_injection_points", None)
)
injection_points: list[CacheControlInjectionPoint] = configured or []
if not injection_points and model is not None:
injection_points = AnthropicCacheControlHook.get_default_injection_points(
messages=cast(list[AllMessageValues], messages), # cast-ok: Anthropic-shaped dicts from v1/messages
system=system,
tools=tools,
model=model,
custom_llm_provider=custom_llm_provider,
)
if not injection_points:
return messages, system

View file

@ -1,40 +1,38 @@
import configparser
import os
import time
import uuid
from typing import Any, Dict, Final, List, Optional, Tuple
CONFIG_FILE_PATH_DEFAULT: Final[str] = "~/.opik.config"
def create_uuid7():
ns = time.time_ns()
last = [0, 0, 0, 0]
def create_uuid7() -> str:
"""Generate an RFC 9562 conformant UUIDv7 string.
# Simple uuid7 implementation
sixteen_secs = 16_000_000_000
t1, rest1 = divmod(ns, sixteen_secs)
t2, rest2 = divmod(rest1 << 16, sixteen_secs)
t3, _ = divmod(rest2 << 12, sixteen_secs)
t3 |= 7 << 12 # Put uuid version in top 4 bits, which are 0 in t3
The top 48 bits encode the Unix timestamp in milliseconds. Opik's backend
validates this embedded timestamp on ingestion (it must fall within a window
around "now"), so the encoding has to be correct or trace/span batches are
rejected with HTTP 400. Implemented with the standard library only, so no
extra dependency is added to litellm. See ``opik.id_helpers`` for the
reference implementation.
"""
unix_ts_ms = int(time.time() * 1000)
# The next two bytes are an int (t4) with two bits for
# the variant 2 and a 14 bit sequence counter which increments
# if the time is unchanged.
if t1 == last[0] and t2 == last[1] and t3 == last[2]:
# Stop the seq counter wrapping past 0x3FFF.
# This won't happen in practice, but if it does,
# uuids after the 16383rd with that same timestamp
# will not longer be correctly ordered but
# are still unique due to the 6 random bytes.
if last[3] < 0x3FFF:
last[3] += 1
else:
last[:] = (t1, t2, t3, 0)
t4 = (2 << 14) | last[3] # Put variant 0b10 in top two bits
# Fill the 16-byte buffer with random data, then overwrite the structured
# parts (timestamp, version, variant) defined by the UUIDv7 layout.
uuid_bytes = bytearray(os.urandom(16))
# Six random bytes for the lower part of the uuid
rand = os.urandom(6)
return f"{t1:>08x}-{t2:>04x}-{t3:>04x}-{t4:>04x}-{rand.hex()}"
# First 48 bits (6 bytes): Unix timestamp in milliseconds.
uuid_bytes[0:6] = unix_ts_ms.to_bytes(6, byteorder="big")
# Version 7 in the top 4 bits of byte 6.
uuid_bytes[6] = 0x70 | (uuid_bytes[6] & 0x0F)
# Variant 0b10 in the top 2 bits of byte 8.
uuid_bytes[8] = 0x80 | (uuid_bytes[8] & 0x3F)
return str(uuid.UUID(bytes=bytes(uuid_bytes)))
def _read_opik_config_file() -> Dict[str, str]:

View file

@ -1453,6 +1453,9 @@ class Logging(LiteLLMLoggingBaseClass):
response_cost = litellm.response_cost_calculator(**response_cost_calculator_kwargs)
verbose_logger.debug(f"response_cost: {response_cost}")
additional_response_cost: object = self.model_call_details.get("additional_response_cost")
if isinstance(additional_response_cost, (int, float)) and additional_response_cost > 0:
return (response_cost or 0.0) + additional_response_cost
return response_cost
except Exception as e: # error calculating cost
debug_info = StandardLoggingModelCostFailureDebugInformation(

View file

@ -906,16 +906,17 @@ def strip_advisor_blocks_from_messages(messages: List[Any], replace_with_text: b
def is_anthropic_invalid_thinking_signature_error(error_text: str) -> bool:
"""
Detect Anthropic 400 when encrypted thinking signatures in history do not match
the current deployment (e.g. user rotated API key or switched model endpoint).
Detect Anthropic 400 errors caused by missing or invalid thinking signatures.
Example API message:
Known error formats:
{"message":"messages.2.content.0.thinking.signature.str: Input should be a valid string"}
messages.N.content.M.thinking.signature.str: Input should be a valid string
messages.N.content.M: Invalid `signature` in `thinking` block
"""
if not error_text:
return False
lower = error_text.lower()
return "invalid" in lower and "signature" in lower and "thinking" in lower and "block" in lower
return "thinking" in lower and "signature" in lower and ("invalid" in lower or "valid string" in lower)
def strip_thinking_blocks_from_anthropic_messages(messages: List[Any]) -> List[Any]:

View file

@ -237,7 +237,9 @@ async def anthropic_messages(
AnthropicCacheControlHook,
)
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(messages, system, kwargs)
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(
messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools
)
original_stream = stream or kwargs.get("_websearch_interception_converted_stream", False)
@ -426,7 +428,9 @@ def anthropic_messages_handler(
AnthropicCacheControlHook,
)
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(messages, system, kwargs)
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(
messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools
)
metadata = validate_anthropic_api_metadata(metadata)

View file

@ -75,8 +75,9 @@ class AnthropicResponsesStreamWrapper:
# ---- message_start ----
if event_type == "response.created":
self._sent_message_start = True
self._chunk_queue.append(self._make_message_start())
if not self._sent_message_start:
self._sent_message_start = True
self._chunk_queue.append(self._make_message_start())
return
# ---- content_block_start for a new output message item ----

View file

@ -75,10 +75,23 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
model_info = get_model_info(model=base_model, custom_llm_provider="fireworks_ai")
## CALCULATE INPUT COST
prompt_tokens_details = usage.prompt_tokens_details
cached_tokens: int = (
prompt_tokens_details.cached_tokens
if prompt_tokens_details is not None and prompt_tokens_details.cached_tokens is not None
else 0
)
input_cost_per_token: float = model_info["input_cost_per_token"] or 0.0
cache_read_input_token_cost = model_info.get("cache_read_input_token_cost")
cache_read_cost_per_token: float = (
cache_read_input_token_cost if cache_read_input_token_cost is not None else input_cost_per_token
)
non_cached_prompt_tokens: int = max(usage.prompt_tokens - cached_tokens, 0)
prompt_cost: float = usage["prompt_tokens"] * model_info["input_cost_per_token"]
prompt_cost: float = non_cached_prompt_tokens * input_cost_per_token + cached_tokens * cache_read_cost_per_token
## CALCULATE OUTPUT COST
completion_cost = usage["completion_tokens"] * model_info["output_cost_per_token"]
output_cost_per_token: float = model_info["output_cost_per_token"] or 0.0
completion_cost: float = usage.completion_tokens * output_cost_per_token
return prompt_cost, completion_cost

View file

@ -510,6 +510,20 @@ async def acompletion(
#########################################################
#########################################################
litellm_logging_obj = kwargs.get("litellm_logging_obj", None)
from litellm.integrations.anthropic_cache_control_hook import (
AnthropicCacheControlHook,
)
from litellm.types.llms.openai import AllMessageValues
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=kwargs,
messages=cast(list[AllMessageValues], messages), # cast-ok: acompletion types messages as a bare List
model=model,
custom_llm_provider=cast(Optional[str], custom_llm_provider), # cast-ok: read from untyped kwargs
tools=tools,
)
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and (
litellm_logging_obj.should_run_prompt_management_hooks(
prompt_id=kwargs.get("prompt_id", None),
@ -5055,6 +5069,19 @@ def completion( # type: ignore
litellm_params = {} # used to prevent unbound var errors
## PROMPT MANAGEMENT HOOKS ##
from litellm.integrations.anthropic_cache_control_hook import (
AnthropicCacheControlHook,
)
from litellm.types.llms.openai import AllMessageValues
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=non_default_params,
messages=cast(list[AllMessageValues], messages), # cast-ok: completion types messages as a bare List
model=model,
custom_llm_provider=cast(Optional[str], kwargs.get("custom_llm_provider")), # cast-ok: untyped kwargs
tools=tools,
)
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and (
litellm_logging_obj.should_run_prompt_management_hooks(
prompt_id=prompt_id, non_default_params=non_default_params

View file

@ -3451,7 +3451,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2.2e-05,
"output_cost_per_token": 2.64e-06,
"supports_audio_input": true,
@ -3470,7 +3470,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.00022,
"output_cost_per_token": 2.2e-05,
"supports_audio_input": true,
@ -3489,7 +3489,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2.2e-05,
"supported_modalities": [
@ -4687,7 +4687,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -4707,7 +4707,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -4739,7 +4739,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -4771,7 +4771,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -4832,7 +4832,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.0002,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -4850,7 +4850,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supported_modalities": [
@ -7922,7 +7922,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2.2e-05,
"output_cost_per_token": 2.64e-06,
"supports_audio_input": true,
@ -7941,7 +7941,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.00022,
"output_cost_per_token": 2.2e-05,
"supports_audio_input": true,
@ -7960,7 +7960,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2.2e-05,
"supported_modalities": [
@ -16272,7 +16272,7 @@
"supports_vision": false
},
"fireworks_ai/accounts/fireworks/models/glm-5p2": {
"cache_read_input_token_cost": 2.6e-07,
"cache_read_input_token_cost": 1.4e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
@ -16686,7 +16686,7 @@
"supports_vision": false
},
"fireworks_ai/glm-5p2": {
"cache_read_input_token_cost": 2.6e-07,
"cache_read_input_token_cost": 1.4e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
@ -22094,7 +22094,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -22113,7 +22113,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -22207,7 +22207,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -22225,7 +22225,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -22243,7 +22243,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -24438,7 +24438,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24470,7 +24470,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24502,7 +24502,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24535,7 +24535,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 2.4e-05,
"regional_processing_uplift_multiplier_eu": 1.1,
@ -24570,7 +24570,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"regional_processing_uplift_multiplier_eu": 1.1,
@ -24603,7 +24603,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -24635,7 +24635,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -43573,7 +43573,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -43606,7 +43606,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [

View file

@ -1220,11 +1220,29 @@ class MCPRequestHandler:
global_mcp_server_manager,
)
key_tools = (
key_direct_tools = (
global_mcp_server_manager.expand_tool_permissions(key_obj_perm.mcp_tool_permissions).get(server_id)
if key_obj_perm
else None
)
# Tools granted through the key's toolsets restrict this server exactly
# as direct tool permissions do; union with any direct grants so the
# tool-level check sees the key's full effective tool scope
key_toolset_ids = (key_obj_perm.mcp_toolsets or []) if key_obj_perm else []
key_toolset_tools = (
(await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=key_toolset_ids)).get(
server_id
)
if key_toolset_ids
else None
)
key_tools = (
list(set(key_direct_tools or []) | set(key_toolset_tools or []))
if key_direct_tools is not None or key_toolset_tools is not None
else None
)
team_tools = (
global_mcp_server_manager.expand_tool_permissions(team_obj_perm.mcp_tool_permissions).get(server_id)
if team_obj_perm
@ -1430,8 +1448,18 @@ class MCPRequestHandler:
global_mcp_server_manager.expand_tool_permissions(key_object_permission.mcp_tool_permissions).keys()
)
# servers referenced by the key's toolset grants are part of the key's
# scope on every path (list, call, REST), subject to the same team/org
# ceilings as any other key-level grant
toolset_ids = key_object_permission.mcp_toolsets or []
toolset_servers = (
list((await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=toolset_ids)).keys())
if toolset_ids
else []
)
# Combine all lists
all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers
all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers + toolset_servers
return list(set(all_servers))
except Exception as e:
verbose_logger.warning(f"Failed to get allowed MCP servers for key: {str(e)}")

View file

@ -88,3 +88,19 @@ class MCPToolResultError(Exception):
into two identities, breaking ``isinstance`` checks against instances
created before the reload.
"""
class MCPServerListError(Exception):
"""Carrier for a classified per-server listing fault (``faults.list_outcomes.ServerListFault``).
Raised where a server fetch used to silently return an empty tool list, so each boundary can
apply its own policy: the aggregate listing absorbs it into that server's outcome, while
single-server routes relay a truthful HTTP status instead of empty-success. The fault value is
typed as ``object`` here only to avoid a circular import with the faults package; construction
sites always pass a ``ServerListFault``.
"""
def __init__(self, fault: object, server_name: str) -> None:
self.fault = fault
self.server_name = server_name
super().__init__(f"Listing tools from MCP server {server_name!r} failed")

View file

@ -0,0 +1,190 @@
"""Per-server outcomes for the aggregate MCP tools/list fan-out.
The aggregate listing deliberately keeps serving the healthy subset when one server fails, but a
failed server must contribute a classified outcome instead of silently shrinking the list: an empty
contribution with no signal makes a broken upstream indistinguishable from a healthy server with no
tools. Outcomes carry only machine fields (category and status code) so nothing from an upstream
body crosses the trust boundary; classification is total, so any exception out of a server fetch
becomes an outcome, never a second failure.
"""
from __future__ import annotations
from collections.abc import Iterator
from typing import Literal, NamedTuple, NoReturn, TypeAlias
import httpx
from mcp.types import Tool as MCPTool
from pydantic import BaseModel, ConfigDict
from typing_extensions import assert_never
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPServerListError,
MCPUpstreamAuthError,
)
ListFaultCategory: TypeAlias = Literal[
"auth_required",
"forbidden",
"timeout",
"unreachable",
"upstream_error",
"internal",
]
class ServerListOk(BaseModel):
model_config = ConfigDict(frozen=True)
tag: Literal["ok"] = "ok"
tool_count: int
class ServerListFault(BaseModel):
"""Why a server contributed nothing to a listing: the caller must authenticate upstream
(``auth_required``/``forbidden``), the upstream did not answer (``timeout``/``unreachable``),
the upstream answered outside its contract (``upstream_error``), or the gateway itself failed
(``internal``). ``status_code`` is the upstream HTTP status when one exists."""
model_config = ConfigDict(frozen=True)
tag: ListFaultCategory
status_code: int | None = None
ServerOutcome: TypeAlias = ServerListOk | ServerListFault
SERVER_OUTCOMES_META_KEY = "litellm.ai/server_outcomes"
"""The tools/list result ``_meta`` key carrying per-server outcomes. Prefixed with the litellm.ai
domain per the MCP spec's ``_meta`` key format so it cannot collide with spec-reserved names."""
class AggregateToolListing(NamedTuple):
tools: list[MCPTool]
outcomes: dict[str, ServerOutcome]
def _iter_upstream_responses(exc: BaseException) -> Iterator[httpx.Response]:
"""Yield every ``httpx.Response`` in the exception tree (``__cause__``/``__context__``/
ExceptionGroup members) in deliberate order, mirroring how upstream failures surface through the
MCP SDK's task groups. Explicit links come first: each node's ``raise ... from`` cause, then
group members in raise order, then the incidental ``__context__`` chain, so a response raised
while handling the real failure can never shadow one on the explicit causal chain. Consumers
apply their own predicate over the stream: selecting the first response and THEN testing it
would miss a causal auth response sitting behind an unrelated earlier one."""
seen: set[int] = set()
stack = [exc]
while stack:
current = stack.pop()
if id(current) in seen:
continue
seen.add(id(current))
response = getattr(current, "response", None)
if isinstance(response, httpx.Response):
yield response
if current.__context__ is not None:
stack.append(current.__context__)
exceptions = getattr(current, "exceptions", None)
if isinstance(exceptions, tuple):
stack.extend(reversed(exceptions))
if current.__cause__ is not None:
stack.append(current.__cause__)
def _find_upstream_response(exc: BaseException) -> httpx.Response | None:
return next(_iter_upstream_responses(exc), None)
def upstream_auth_challenge(exc: BaseException) -> tuple[int, str | None] | None:
"""The first upstream 401/403 in deliberate order and its ``WWW-Authenticate`` challenge, both
read from the SAME response, so the status that picks the carrier channel and the challenge that
rides with it can never come from two different responses in the tree. Non-auth responses do not
end the scan: a causal 401 behind an unrelated 5xx must still be found, or the client never
receives the challenge it needs to re-authenticate."""
for response in _iter_upstream_responses(exc):
if response.status_code in (401, 403):
return response.status_code, response.headers.get("www-authenticate")
return None
def raise_classified_list_failure(
exc: BaseException,
server_name: str,
suppress_challenge: bool = False,
) -> NoReturn:
"""The one place a failed server fetch chooses its carrier: an upstream 401/403 travels as
``MCPUpstreamAuthError`` with the upstream's own challenge preserved (a challenge is only ever
fabricated at the HTTP edge, and only for a 401), everything else as ``MCPServerListError`` with
a classified fault. Every fetch site delegates here so the two channels cannot drift apart per
call site. ``suppress_challenge`` is for dcr_bridge servers, whose upstream challenge points
clients at the wrong protected-resource metadata and must never relay."""
auth = upstream_auth_challenge(exc)
if auth is not None:
status_code, challenge = auth
raise MCPUpstreamAuthError(
status_code=status_code,
www_authenticate=None if suppress_challenge else challenge,
server_name=server_name,
) from exc
raise MCPServerListError(classify_list_exception(exc), server_name) from exc
def classify_list_exception(exc: BaseException) -> ServerListFault:
"""Classify a per-server listing failure into exactly one outcome. Total: an exception this
function cannot recognize is the gateway's own fault (``internal``), never a re-raise."""
if isinstance(exc, MCPServerListError) and isinstance(exc.fault, ServerListFault):
return exc.fault
if isinstance(exc, MCPUpstreamAuthError):
tag = "forbidden" if exc.status_code == 403 else "auth_required"
return ServerListFault(tag=tag, status_code=exc.status_code)
if isinstance(exc, TimeoutError):
return ServerListFault(tag="timeout")
if isinstance(exc, ConnectionError):
return ServerListFault(tag="unreachable")
auth = upstream_auth_challenge(exc)
if auth is not None:
status_code, _ = auth
return ServerListFault(
tag="forbidden" if status_code == 403 else "auth_required",
status_code=status_code,
)
response = _find_upstream_response(exc)
if response is not None:
return ServerListFault(tag="upstream_error", status_code=response.status_code)
if isinstance(exc, (httpx.TimeoutException,)):
return ServerListFault(tag="timeout")
if isinstance(exc, httpx.TransportError):
return ServerListFault(tag="unreachable")
return ServerListFault(tag="internal")
def outcome_wire_value(outcome: ServerOutcome) -> dict[str, object]:
"""The client-visible form of one outcome, for the tools/list result ``_meta`` and the REST
response: category plus status code only, never upstream prose or URLs."""
match outcome.tag:
case "ok":
return {"status": "ok", "tool_count": outcome.tool_count}
case "auth_required" | "forbidden" | "timeout" | "unreachable" | "upstream_error" | "internal":
return {
"status": outcome.tag,
**({"http_status": outcome.status_code} if outcome.status_code is not None else {}),
}
case _:
assert_never(outcome.tag)
def list_fault_http_status(fault: ServerListFault) -> int:
"""The truthful HTTP status for a single-upstream listing fault per RFC 9110: the upstream's own
401/403 for auth, 504 for a timeout, 502 for an unreachable or misbehaving upstream, and 500 only
for the gateway's own failure."""
match fault.tag:
case "auth_required":
return fault.status_code or 401
case "forbidden":
return 403
case "timeout":
return 504
case "unreachable" | "upstream_error":
return 502
case "internal":
return 500
case _:
assert_never(fault.tag)

View file

@ -50,7 +50,15 @@ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
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.faults.list_outcomes import (
ServerListFault,
raise_classified_list_failure,
upstream_auth_challenge,
)
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
MCP_ELICITATION_AVAILABLE,
)
@ -633,49 +641,14 @@ def _caller_authorization_fans_out(
def _extract_upstream_auth_failure(
exc: BaseException,
) -> Optional[tuple[int, Optional[str]]]:
"""Walk the exception tree looking for an HTTP 401/403 response from the
upstream MCP server.
"""The upstream 401/403 and its ``WWW-Authenticate`` header from the exception tree, or ``None``.
The MCP SDK wraps transport errors in anyio ``ExceptionGroup`` objects and
may chain through ``__cause__`` / ``__context__``. We inspect all of those
layers for an ``httpx.Response``-bearing exception (typically
``httpx.HTTPStatusError``) and extract the status code and any upstream
``WWW-Authenticate`` header.
Returns ``(status_code, www_authenticate)`` on match, else ``None``.
"""
seen: set[int] = set()
stack: list[BaseException] = [exc]
while stack:
current = stack.pop()
if id(current) in seen:
continue
seen.add(id(current))
response = getattr(current, "response", None)
if response is not None:
status_code = getattr(response, "status_code", None)
if isinstance(status_code, int) and status_code in (401, 403):
www_authenticate: Optional[str] = None
headers = getattr(response, "headers", None)
if headers is not None:
try:
www_authenticate = headers.get("www-authenticate")
except Exception:
www_authenticate = None
return status_code, www_authenticate
# anyio / PEP 654 ExceptionGroup
sub_exceptions = getattr(current, "exceptions", None)
if sub_exceptions:
stack.extend(sub_exceptions)
if current.__cause__ is not None:
stack.append(current.__cause__)
if current.__context__ is not None and current.__context__ is not current.__cause__:
stack.append(current.__context__)
return None
Delegates to the shared traversal in ``faults.list_outcomes`` so every consumer (tool listing,
tool calls, the connect-time probe) selects the same response with the same deliberate order:
explicit ``raise ... from`` causes first, ExceptionGroup members in raise order, the incidental
``__context__`` chain last. A response raised while handling the real failure can therefore never
shadow the causal one."""
return upstream_auth_challenge(exc)
def _warn_on_server_name_fields(
@ -3047,10 +3020,12 @@ class MCPServerManager:
server_name=server.name,
) from e
verbose_logger.warning(f"Failed to get tools from server {server.name}: {str(e)}")
return []
raise MCPServerListError(ServerListFault(tag="internal", status_code=e.status_code), server.name) from e
except MCPServerListError:
raise
except Exception as e:
verbose_logger.warning(f"Failed to get tools from server {server.name}: {str(e)}")
return []
raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge)
async def get_prompts_from_server(
self,
@ -3682,16 +3657,17 @@ class MCPServerManager:
Uses anyio.fail_after() instead of asyncio.wait_for() to avoid conflicts
with the MCP SDK's anyio TaskGroup. See GitHub issue #20715 for details.
An upstream HTTP 401 is converted into :class:`MCPUpstreamAuthError`
instead of being swallowed to an empty tool list, regardless of the
server's auth_type. Callers route it by surface: the single-server HTTP
routes turn it into a 401 + ``WWW-Authenticate`` challenge so standards-
compliant MCP clients trigger the upstream OAuth flow, while the
multi-server ``/mcp`` aggregator absorbs it to an empty list so one
unauthenticated server doesn't fail the whole listing. Only a 401
(missing/invalid credential) drives the re-auth challenge; a 403
(authenticated but forbidden, e.g. insufficient scope) is not a re-auth
signal and, like other non-auth errors, returns an empty list.
Failures never return an empty tool list. An upstream 401 or 403 raises
:class:`MCPUpstreamAuthError` carrying the upstream's own
``WWW-Authenticate`` challenge when one was sent (a challenge is only
ever fabricated at the HTTP edge, and only for a 401: a 403 means the
caller is authenticated but not allowed, so prompting re-auth would be
wrong, while an upstream-sent 403 challenge is the RFC 6750
insufficient_scope step-up and relays verbatim). Every other failure
raises :class:`MCPServerListError` with a classified fault. Each
boundary then applies its own policy: single-server routes relay the
truthful status, the multi-server aggregator absorbs the failure into
that server's listing outcome.
Args:
client: MCP client instance
@ -3705,27 +3681,18 @@ class MCPServerManager:
tools = await client.list_tools(raise_on_error=True)
verbose_logger.debug(f"Tools from {server_name}: {tools}")
return tools
except TimeoutError:
except TimeoutError as e:
verbose_logger.warning(f"Timeout while listing tools from {server_name}")
return []
except asyncio.CancelledError:
raise MCPServerListError(ServerListFault(tag="timeout"), server_name) from e
except asyncio.CancelledError as e:
verbose_logger.warning(f"Task cancelled while listing tools from {server_name}")
return []
raise MCPServerListError(ServerListFault(tag="internal"), server_name) from e
except ConnectionError as e:
verbose_logger.warning(f"Connection error while listing tools from {server_name}: {str(e)}")
return []
raise MCPServerListError(ServerListFault(tag="unreachable"), server_name) from e
except Exception as e:
auth_info = _extract_upstream_auth_failure(e)
if auth_info is not None and auth_info[0] == 401:
_, www_authenticate = auth_info
verbose_logger.info(f"Upstream auth failure from MCP server {server_name}: HTTP 401")
raise MCPUpstreamAuthError(
status_code=401,
www_authenticate=www_authenticate,
server_name=server_name,
) from e
verbose_logger.warning(f"Error listing tools from {server_name}: {str(e)}")
return []
raise_classified_list_failure(e, server_name)
_SHORT_PREFIX_MAX_REHASH_ATTEMPTS = 1024

View file

@ -19,12 +19,20 @@ import httpx
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
from litellm._logging import verbose_logger
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.faults.list_outcomes import (
classify_list_exception,
list_fault_http_status,
)
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
build_effective_auth_contexts,
)
from litellm.proxy._experimental.mcp_server.utils import (
MCPMissingUserEnvVarsError,
get_server_prefix,
merge_mcp_headers,
)
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
@ -515,20 +523,19 @@ if MCP_AVAILABLE:
# enforced even when no allowlist is set (matches the SSE/HTTP path).
tools = filter_tools_by_allowed_tools(tools, server)
# Filter tools based on user_api_key_auth.object_permission.mcp_tool_permissions
# This provides per-key/team/org control over which tools can be accessed
if (
user_api_key_auth
and user_api_key_auth.object_permission
and user_api_key_auth.object_permission.mcp_tool_permissions
):
# Dict keys may be server_ids OR names/aliases; normalize so lookup
# by concrete server_id resolves name-keyed restrictions too.
allowed_tools_for_server = global_mcp_server_manager.expand_tool_permissions(
user_api_key_auth.object_permission.mcp_tool_permissions
).get(server.server_id)
if allowed_tools_for_server is not None and len(allowed_tools_for_server) > 0:
# Filter tools to only include those in the allowed list
# Filter by the key's effective tool permissions through the same
# primitive the MCP protocol path uses (direct grants, toolset grants,
# and team/agent/org ceilings), so REST listing cannot drift from it
if user_api_key_auth:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
allowed_tools_for_server = await MCPRequestHandler.get_allowed_tools_for_server(
server_id=server.server_id,
user_api_key_auth=user_api_key_auth,
)
if allowed_tools_for_server is not None:
tools = [tool for tool in tools if _tool_name_matches(tool.name, allowed_tools_for_server)]
return _create_tool_response_objects(tools, server)
@ -627,6 +634,16 @@ if MCP_AVAILABLE:
# matching status code and WWW-Authenticate challenge; that is what
# lets standards-compliant MCP clients run the upstream OAuth flow.
raise
except MCPServerListError as e:
fault = classify_list_exception(e)
verbose_logger.info(f"Listing tools from {server.name} failed with a {fault.tag} fault")
raise HTTPException(
status_code=list_fault_http_status(fault),
detail={
"error": fault.tag,
"message": f"Failed to list tools from server {get_server_prefix(server)}",
},
) from e
except Exception as e:
verbose_logger.exception(f"Error getting tools from {server.name}: {e}")
return {
@ -838,7 +855,11 @@ if MCP_AVAILABLE:
list_tools_result.extend(tools_result)
except Exception as e:
verbose_logger.exception(f"Error getting tools from {server.name}: {e}")
errors.append(f"{server.name}: {str(e)}")
errors.append(
f"{get_server_prefix(server)}: {classify_list_exception(e).tag}"
if isinstance(e, (MCPServerListError, MCPUpstreamAuthError))
else f"{get_server_prefix(server)}: {str(e)}"
)
continue
if errors and not list_tools_result:
@ -858,7 +879,10 @@ if MCP_AVAILABLE:
request_path=request.scope.get("_original_path") or request.url.path,
)
except HTTPException as http_exc:
if http_exc.status_code == status.HTTP_404_NOT_FOUND:
if http_exc.status_code == status.HTTP_404_NOT_FOUND or server_id:
# Single-server requests relay the truthful status (a 502/504 upstream fault must
# not masquerade as a 200 empty-success body); only the multi-server aggregate
# keeps the legacy error-dict response shape below.
raise
# Internal access/IP 403s keep the legacy error-dict response shape
# so the existing contract stays intact.

View file

@ -348,6 +348,7 @@ if MCP_AVAILABLE:
CallToolResult,
EmbeddedResource,
ImageContent,
ListToolsResult,
Prompt,
TextContent,
)
@ -356,6 +357,14 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import (
MCPAuthenticatedUser,
)
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
SERVER_OUTCOMES_META_KEY,
AggregateToolListing,
ServerListOk,
ServerOutcome,
classify_list_exception,
outcome_wire_value,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
_caller_authorization_fans_out,
@ -664,9 +673,12 @@ if MCP_AVAILABLE:
########################################################
@server.list_tools()
async def handle_list_tools() -> List[Tool]:
async def handle_list_tools() -> "ListToolsResult | List[Tool]":
"""
List all available tools.
List all available tools, with each server's listing outcome attached to the result's
``_meta`` (SERVER_OUTCOMES_META_KEY) so a broken upstream is distinguishable from a healthy
server with no tools. Returning a ListToolsResult (rather than a bare list) makes the MCP SDK
pass the result through unwrapped, which is what lets the ``_meta`` survive to the client.
Also captures the active session for propagation to callbacks.
"""
from mcp.server.lowlevel.server import request_ctx
@ -709,7 +721,7 @@ if MCP_AVAILABLE:
# Get mcp_servers from context variable
verbose_logger.debug("MCP list_tools - Calling _list_mcp_tools")
tools = await _list_mcp_tools(
listing = await _list_mcp_tools(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
@ -719,8 +731,15 @@ if MCP_AVAILABLE:
log_list_tools_to_spendlogs=True,
list_tools_log_source="mcp_protocol",
)
verbose_logger.info(f"MCP list_tools - Successfully returned {len(tools)} tools")
return tools
verbose_logger.info(f"MCP list_tools - Successfully returned {len(listing.tools)} tools")
if not listing.outcomes:
return listing.tools
outcome_meta = {
SERVER_OUTCOMES_META_KEY: {
key: outcome_wire_value(outcome) for key, outcome in listing.outcomes.items()
}
}
return ListToolsResult.model_validate({"tools": listing.tools, "_meta": outcome_meta})
except Exception as e:
verbose_logger.exception(f"Error in list_tools endpoint: {str(e)}")
# Return empty list instead of failing completely
@ -1746,6 +1765,13 @@ if MCP_AVAILABLE:
_mcp_gateway_initialize_instructions.reset(instructions_token)
_mcp_gateway_server_name.reset(server_name_token)
def _aggregate_server_key(server: MCPServer) -> str:
"""The client-visible key for a server in listing outcomes and spend metadata: the same
display prefix (alias, or the short prefix when that mode is enabled) the caller already
sees on the tool names. Canonical internal server names never key a caller-readable
surface; when the display naming deliberately hides them, the outcome keys must too."""
return get_server_prefix(server) or "unknown"
async def _get_tools_from_mcp_servers(
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_auth_header: Optional[str],
@ -1758,7 +1784,7 @@ if MCP_AVAILABLE:
litellm_trace_id: Optional[str] = None,
request_tags: Optional[list[str]] = None,
client_ip: Optional[str] = None,
) -> List[MCPTool]:
) -> AggregateToolListing:
"""
Helper method to fetch tools from MCP servers based on server filtering criteria.
@ -1770,10 +1796,11 @@ if MCP_AVAILABLE:
oauth2_headers: Optional dict of oauth2 headers
Returns:
List[MCPTool]: Combined list of tools from filtered servers
AggregateToolListing: Combined tools from filtered servers plus each server's
classified listing outcome
"""
if not MCP_AVAILABLE:
return []
return AggregateToolListing(tools=[], outcomes={})
list_tools_start_time = datetime.now()
litellm_logging_obj: Optional[LiteLLMLoggingObj] = None
@ -1858,10 +1885,12 @@ if MCP_AVAILABLE:
async def _fetch_and_filter_server_tools(
server: MCPServer,
) -> List[MCPTool]:
"""Fetch and filter tools from a single server with error handling."""
) -> "tuple[List[MCPTool], ServerOutcome]":
"""Fetch and filter tools from a single server, classifying any failure into that
server's outcome so the aggregate can keep serving the healthy subset without a
broken server masquerading as an empty one."""
if server is None:
return []
return [], ServerListOk(tool_count=0)
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
@ -1931,8 +1960,8 @@ if MCP_AVAILABLE:
verbose_logger.debug(
f"Successfully fetched {len(tools)} tools from server {server.name}, {len(filtered_tools)} after filtering"
)
return filtered_tools
except MCPUpstreamAuthError:
return filtered_tools, ServerListOk(tool_count=len(filtered_tools))
except MCPUpstreamAuthError as e:
# Absorb so one unauthenticated server does not empty every other server's
# tools. Surfacing the upstream 401 to the client as a re-auth challenge is
# intentionally not done here: raising from this list handler cannot produce a
@ -1940,31 +1969,30 @@ if MCP_AVAILABLE:
# error). Single-server routes surface it via the request-scope preemptive
# check in _raise_preemptive_401_for_unauthenticated_servers instead.
verbose_logger.debug(f"MCP list_tools: omitting {server.name}; it needs upstream auth")
return []
return [], classify_list_exception(e)
except Exception as e:
verbose_logger.exception(f"Error getting tools from server {server.name}: {str(e)}")
return []
return [], classify_list_exception(e)
# Fetch tools from all servers in parallel
tasks = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers]
results = await asyncio.gather(*tasks)
# Flatten results into single list
all_tools: List[MCPTool] = [tool for tools in results for tool in tools]
all_tools: List[MCPTool] = [tool for tools, _ in results for tool in tools]
server_outcomes: Dict[str, ServerOutcome] = {
_aggregate_server_key(server): outcome
for server, (_, outcome) in zip(allowed_mcp_servers, results)
if server is not None
}
# If logging is enabled, enrich spend_logs_metadata with counts
if litellm_logging_obj:
per_server_tool_counts: Dict[str, int] = {}
for server, server_tools in zip(allowed_mcp_servers, results):
if server is None:
continue
server_key = (
getattr(server, "server_name", None)
or getattr(server, "alias", None)
or getattr(server, "name", None)
or "unknown"
)
per_server_tool_counts[str(server_key)] = len(server_tools)
per_server_tool_counts: Dict[str, int] = {
_aggregate_server_key(server): len(server_tools)
for server, (server_tools, _) in zip(allowed_mcp_servers, results)
if server is not None
}
metadata_dict = litellm_logging_obj.model_call_details.get("metadata")
if isinstance(metadata_dict, dict):
@ -1975,6 +2003,9 @@ if MCP_AVAILABLE:
spend_meta["allowed_server_count"] = len(allowed_mcp_servers)
spend_meta["tool_count_total"] = len(all_tools)
spend_meta["per_server_tool_counts"] = per_server_tool_counts
spend_meta["per_server_list_outcomes"] = {
key: outcome_wire_value(outcome) for key, outcome in server_outcomes.items()
}
end_time = datetime.now()
try:
@ -1995,7 +2026,7 @@ if MCP_AVAILABLE:
verbose_logger.info(f"Successfully fetched {len(all_tools)} tools total from all MCP servers")
return all_tools
return AggregateToolListing(tools=all_tools, outcomes=server_outcomes)
except Exception as e:
# Only fire failure hook if logging was requested for this list-tools execution
if log_list_tools_to_spendlogs and user_api_key_auth is not None:
@ -2218,43 +2249,6 @@ if MCP_AVAILABLE:
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
return [t for t in tools if strip_known_server_prefix(t.name, server) in allowed_tool_names]
async def _merge_toolset_permissions(
user_api_key_auth: Optional[UserAPIKeyAuth],
) -> Optional[UserAPIKeyAuth]:
"""
Resolve mcp_toolsets on the key's object_permission into tool-level permissions
and merge them (union) into object_permission.mcp_tool_permissions.
Returns the (possibly mutated copy of) user_api_key_auth.
"""
if user_api_key_auth is None:
return None
op = user_api_key_auth.object_permission
if op is None:
return user_api_key_auth
toolset_ids = getattr(op, "mcp_toolsets", None) or []
if not toolset_ids:
return user_api_key_auth
toolset_perms = await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=toolset_ids)
if not toolset_perms:
return user_api_key_auth
# Merge toolset_perms into existing mcp_tool_permissions (union)
existing = dict(op.mcp_tool_permissions or {})
for server_id, tool_names in toolset_perms.items():
existing_tools = existing.get(server_id, [])
merged = list(set(existing_tools) | set(tool_names))
existing[server_id] = merged
# Build updated object_permission with merged tool permissions and server IDs.
# Union the toolset's server IDs into mcp_servers so downstream server-level
# filtering doesn't silently drop servers that the toolset references but that
# aren't already in the key's explicit mcp_servers list.
merged_servers = list(set(op.mcp_servers or []) | set(existing.keys()))
updated_op = op.model_copy(update={"mcp_servers": merged_servers, "mcp_tool_permissions": existing})
return user_api_key_auth.model_copy(update={"object_permission": updated_op})
async def _list_mcp_tools(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
@ -2265,7 +2259,7 @@ if MCP_AVAILABLE:
log_list_tools_to_spendlogs: bool = False,
list_tools_log_source: Optional[str] = None,
client_ip: Optional[str] = None,
) -> List[MCPTool]:
) -> AggregateToolListing:
"""
List all available MCP tools.
@ -2277,19 +2271,14 @@ if MCP_AVAILABLE:
client_ip: Client IP for IP-based server access control
Returns:
List[MCPTool]: Combined list of tools from all accessible servers
AggregateToolListing: Combined tools from all accessible servers plus each server's
classified listing outcome
"""
if not MCP_AVAILABLE:
return []
return AggregateToolListing(tools=[], outcomes={})
# Resolve toolset permissions and merge into the key's object_permission
# so that the existing filter_tools_by_key_team_permissions logic picks them up.
user_api_key_auth = await _merge_toolset_permissions(user_api_key_auth)
# Get tools from managed MCP servers with error handling
managed_tools = []
try:
managed_tools = await _get_tools_from_mcp_servers(
listing = await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
@ -2300,12 +2289,12 @@ if MCP_AVAILABLE:
list_tools_log_source=list_tools_log_source,
client_ip=client_ip,
)
verbose_logger.debug(f"Successfully fetched {len(managed_tools)} tools from managed MCP servers")
verbose_logger.debug(f"Successfully fetched {len(listing.tools)} tools from managed MCP servers")
return listing
except Exception as e:
verbose_logger.exception(f"Error getting tools from managed MCP servers: {str(e)}")
# Continue with empty managed tools list instead of failing completely
return managed_tools
# Continue with an empty listing instead of failing completely
return AggregateToolListing(tools=[], outcomes={})
async def _list_mcp_prompts(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,

View file

@ -91,7 +91,7 @@ async def handle_mcp_tool_search(
from litellm.proxy._experimental.mcp_server.server import _list_mcp_tools
mcp_tools = await _list_mcp_tools(
mcp_listing = await _list_mcp_tools(
user_api_key_auth=user_api_key_dict,
mcp_servers=mcp_servers,
client_ip=client_ip,
@ -100,6 +100,7 @@ async def handle_mcp_tool_search(
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
mcp_tools = mcp_listing.tools
tools = [
{
"name": t.name,

View file

@ -527,6 +527,7 @@ async def common_checks(
request_headers=_safe_get_request_headers(request=request),
request_query_params=_safe_get_request_query_params(request=request),
llm_router=llm_router,
request=request,
)
if route in MODEL_DISCOVERY_ROUTES:

View file

@ -14,6 +14,9 @@ from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH, STANDARD_CUSTOMER_ID_HE
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
from litellm.proxy._types import *
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
)
from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS
from litellm.types.utils import CustomPricingLiteLLMParams
@ -1482,13 +1485,50 @@ def _format_model_candidates(
return candidates
def _request_dispatched_to_pass_through_endpoint(request: Request | None) -> bool:
"""Whether FastAPI resolved this request to a user-defined pass-through handler.
Reads the marker set by ``create_pass_through_route`` off the dispatched endpoint
(``request.scope["endpoint"]``). Because routing has already run by the time auth
dependencies execute, this reflects the handler that actually serves the request:
a custom path colliding with a built-in route resolves to the built-in handler,
which carries no marker, so model-access checks are never wrongly skipped.
"""
if request is None:
return False
scope = getattr(request, "scope", None)
if not isinstance(scope, dict):
return False
endpoint = scope.get("endpoint")
# Identity check against True (not truthiness): the marker is set to the literal
# True, and this keeps a spec'd Mock request (whose attribute access yields truthy
# child mocks) from being misread as a pass-through dispatch.
return getattr(endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, False) is True
def get_model_from_request(
request_data: dict,
route: str,
request_headers: Optional[Mapping[str, Any]] = None,
request_query_params: Optional[Mapping[str, Any]] = None,
llm_router: Optional[Router] = None,
request: Request | None = None,
) -> Optional[Union[str, List[str]]]:
"""Resolve the model(s) a request targets, for model-access and budget checks.
Returns ``None`` when the request was dispatched to a user-defined pass-through
endpoint: its body is forwarded verbatim to the configured upstream, so a
``model`` field there names an upstream model, not a LiteLLM-managed one, and
enforcing key/team model allowlists against it would reject valid requests. The
check reads the FastAPI-resolved endpoint (``request.scope["endpoint"]``), not the
request path, so a custom path that collides with a built-in route never
suppresses model-access checks: on a collision the built-in handler is dispatched
and does not carry the marker. Built-in provider passthrough routes
(``/vertex_ai``, ``/gemini``, ...) are separate handlers and keep model enforcement.
"""
if _request_dispatched_to_pass_through_endpoint(request):
return None
candidates = _extract_model_candidates_from_request(
request_data=request_data,
route=route,

View file

@ -162,6 +162,7 @@ def _get_model_from_request_context(
request_headers=_safe_get_request_headers(request=request),
request_query_params=_safe_get_request_query_params(request=request),
llm_router=llm_router,
request=request,
)

View file

@ -151,6 +151,80 @@ async def _record_streaming_client_disconnect_if_needed(
return True
def _deferred_stream_logging_is_armed(request_data: dict) -> bool:
logging_obj = request_data.get("litellm_logging_obj")
if logging_obj is None:
return False
return (
getattr(logging_obj, "_on_deferred_stream_complete", None) is not None
and getattr(logging_obj, "_deferred_stream_complete_args", None) is not None
)
async def _bill_partial_streamed_spend_on_disconnect(request_data: dict, response: object) -> bool:
"""
A client disconnect throws GeneratorExit/CancelledError into the streaming
generator, so neither the success nor the failure logging callback fires
and the chunks already streamed (plus any sub-call cost folded into the
logging object) would never reach spend tracking. Assemble the partial
response from the wrapper's collected chunks and dispatch success logging
for it; dispatch_success_handlers dedups against a natural end-of-stream
dispatch via has_dispatched_final_stream_success.
Awaited directly by the shielded cleanup rather than scheduled with
create_task: the client is already gone so the extra latency is harmless,
and an unrooted task could be garbage-collected before it bills.
Returns True when a disconnect-time success event owns the request's
max_parallel_requests slot release (one was dispatched here, or one had
already been dispatched for this stream), so the caller can skip the
explicit slot release and avoid a double release. Returns False when no
success event fired (logging disabled, nothing streamed, or assembly
failed) and the caller must release the slot itself.
"""
if litellm.disable_streaming_logging is True:
return False
logging_obj = request_data.get("litellm_logging_obj")
if not isinstance(logging_obj, LiteLLMLoggingObj):
return False
if logging_obj.model_call_details.get("has_dispatched_final_stream_success"):
# A natural end-of-stream success event already fired and released the
# slot; do not bill again, and let the caller skip the slot release.
return True
chunks: object = getattr(response, "chunks", None)
if not isinstance(chunks, list) or not chunks:
return False
verbose_proxy_logger.debug(
"Billing partial streamed spend for %s chunks after client disconnect, litellm_call_id=%s",
len(chunks),
request_data.get("litellm_call_id"),
)
messages: object = getattr(response, "messages", None)
try:
partial_response = litellm.stream_chunk_builder(
chunks=chunks,
messages=messages if isinstance(messages, list) else None,
logging_obj=logging_obj,
)
except Exception as e: # noqa: BLE001 # partial billing is best-effort; never break stream teardown
verbose_proxy_logger.debug("Failed to assemble partial streamed response for disconnect billing: %s", e)
return False
if partial_response is None:
return False
try:
await logging_obj.dispatch_success_handlers(
partial_response,
cache_hit=False,
start_time=None,
end_time=None,
prefer_async_handlers=True,
)
except Exception as e: # noqa: BLE001 # partial billing is best-effort; never break stream teardown
verbose_proxy_logger.debug("Failed to dispatch disconnect billing event: %s", e)
return False
return True
async def _cancel_pending_gather_tasks(tasks: list["asyncio.Task[Any]"]) -> None:
pending_tasks = [task for task in tasks if not task.done()]
for task in pending_tasks:
@ -851,9 +925,12 @@ class ProxyBaseLLMRequestProcessing:
# If conversion fails, use original spend
pass
model_name = ProxyBaseLLMRequestProcessing._get_deployment_model_name(litellm_logging_obj)
headers = {
"x-litellm-call-id": call_id,
"x-litellm-model-id": model_id,
"x-litellm-model-name": model_name,
"x-litellm-cache-key": cache_key,
"x-litellm-model-api-base": (
api_base.split("?")[0] if api_base else None
@ -1322,6 +1399,27 @@ class ProxyBaseLLMRequestProcessing:
model_id = model_info.get("id", "") or ""
return model_id
@staticmethod
def _get_deployment_model_name(
litellm_logging_obj: LiteLLMLoggingObj | None,
) -> str | None:
"""Extract the underlying deployment model string (e.g. ``azure/gpt-4o``).
The router rewrites the response ``model`` field to the model-group alias
the client requested, so neither the response body nor the existing
headers expose the concrete deployment model. The router records it under
``litellm_params`` metadata as ``deployment``, so read it back from there.
"""
litellm_params = getattr(litellm_logging_obj, "litellm_params", None)
if not isinstance(litellm_params, dict):
return None
for key in ("litellm_metadata", "metadata"):
metadata = litellm_params.get(key, {}) or {}
deployment = metadata.get("deployment")
if deployment:
return deployment
return None
@staticmethod
def _response_cost_from_logging_obj(
*,
@ -2575,6 +2673,8 @@ class ProxyBaseLLMRequestProcessing:
response: Any,
stream_completed: bool = False,
client_disconnected: bool = False,
user_api_key_dict: UserAPIKeyAuth | None = None,
proxy_logging_obj: ProxyLogging | None = None,
) -> None:
with anyio.CancelScope(shield=True):
should_record_client_disconnect = client_disconnected or (not stream_completed)
@ -2586,7 +2686,28 @@ class ProxyBaseLLMRequestProcessing:
client_disconnected,
)
if recorded_client_disconnect:
deferred_stream_logging_armed = _deferred_stream_logging_is_armed(request_data)
ProxyLogging._fire_deferred_stream_logging(request_data)
# A disconnect-time success event (the deferred-guardrail flush
# above, or the partial-spend billing below) releases the
# request's max_parallel_requests slot through the limiter's
# own success callback. Release the slot explicitly only when
# no such event fires, so exactly one release happens; two
# concurrent releases would race and double-decrement under the
# limiter's in-memory fallback.
success_event_owns_slot_release = deferred_stream_logging_armed
if not deferred_stream_logging_armed:
success_event_owns_slot_release = await _bill_partial_streamed_spend_on_disconnect(
request_data, response
)
if (
not success_event_owns_slot_release
and proxy_logging_obj is not None
and user_api_key_dict is not None
):
await proxy_logging_obj._arelease_max_parallel_requests_on_disconnect(
user_api_key_dict, request_data
)
if hasattr(response, "aclose"):
try:
@ -2675,12 +2796,13 @@ class ProxyBaseLLMRequestProcessing:
except (asyncio.CancelledError, GeneratorExit):
# Client disconnected mid-stream. CancelledError / GeneratorExit
# are BaseException and bypass the success/failure logging
# callbacks that release the pre-call max_parallel_requests +1;
# release it here. This is the outermost generator Starlette closes
# on disconnect, so the nested iterator hook (which only sees
# GeneratorExit on GC) cannot own the refund.
# callbacks that release the pre-call max_parallel_requests +1.
# Flag the disconnect; the shielded cleanup in `finally` owns the
# slot release so it can coordinate with disconnect-time success
# billing and release exactly once. This is the outermost generator
# Starlette closes on disconnect, so the nested iterator hook (which
# only sees GeneratorExit on GC) cannot own the refund.
if not stream_completed:
proxy_logging_obj._release_max_parallel_requests_on_disconnect(user_api_key_dict)
client_disconnected = True
if not delivered_chunk:
from litellm.proxy.spend_tracking.budget_reservation import (
@ -2723,6 +2845,8 @@ class ProxyBaseLLMRequestProcessing:
response=response,
stream_completed=stream_completed,
client_disconnected=client_disconnected,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
)
@staticmethod

View file

@ -0,0 +1,44 @@
from typing import TYPE_CHECKING
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .singulr import SingulrGuardrail
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def initialize_guardrail(
litellm_params: "LitellmParams",
guardrail: "Guardrail",
):
import litellm
_cb = SingulrGuardrail(
singulr_api_base=getattr(litellm_params, "singulr_api_base", None) or litellm_params.api_base,
singulr_api_key=getattr(litellm_params, "singulr_api_key", None) or litellm_params.api_key,
singulr_application_id=getattr(litellm_params, "singulr_application_id", None),
singulr_guardrail_id=getattr(litellm_params, "singulr_guardrail_id", None),
block_on_error=getattr(litellm_params, "block_on_error", None),
timeout=litellm_params.timeout,
guardrail_name=guardrail.get(
"guardrail_name",
"",
),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
)
litellm.logging_callback_manager.add_litellm_callback(
_cb,
)
return _cb
guardrail_initializer_registry = {
SupportedGuardrailIntegrations.SINGULR.value: initialize_guardrail,
}
guardrail_class_registry = {
SupportedGuardrailIntegrations.SINGULR.value: SingulrGuardrail,
}

View file

@ -0,0 +1,216 @@
import os
from typing import Any
from urllib.parse import urlparse
import httpx
import pydantic
from litellm._logging import verbose_proxy_logger
from litellm.exceptions import GuardrailRaisedException
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.litellm_core_utils.litellm_logging import (
Logging as LiteLLMLoggingObj,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.guardrails.guardrail_hooks.base import (
GuardrailConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.singulr import (
SingulrGuardrailPayload,
SingulrGuardrailRequest,
SingulrGuardrailResponse,
)
from litellm.types.utils import GenericGuardrailAPIInputs
_DEFAULT_API_BASE = "http://localhost:8003"
_GUARD_ENDPOINT = "/api/v1/ai-gateway/litellm"
_DEFAULT_TIMEOUT = 30.0
class SingulrGuardrail(CustomGuardrail):
def __init__(
self,
singulr_api_key: str | None = None,
singulr_api_base: str | None = None,
singulr_application_id: str | None = None,
singulr_guardrail_id: str | None = None,
block_on_error: bool | None = None,
timeout: float | None = None,
**kwargs: Any,
) -> None:
self.singulr_api_key = singulr_api_key or os.environ.get("SINGULR_API_KEY")
self.singulr_api_base = (singulr_api_base or os.environ.get("SINGULR_API_BASE") or _DEFAULT_API_BASE).rstrip(
"/"
)
parsed = urlparse(self.singulr_api_base)
if parsed.scheme == "http" and parsed.hostname not in (
"localhost",
"127.0.0.1",
):
raise ValueError(
f"Singulr: api_base {self.singulr_api_base} uses plain HTTP for a "
"non-local endpoint. Guardrail payloads contain the API token, full "
"conversation content, and the guardrail decision, so this endpoint "
"must use HTTPS."
)
self.singulr_application_id = singulr_application_id or os.environ.get("SINGULR_ENFORCEMENT_ENTITY_ID")
self.singulr_guardrail_id = singulr_guardrail_id or os.environ.get("SINGULR_GUARDRAIL_ID")
if block_on_error is None:
env = os.environ.get("SINGULR_BLOCK_ON_ERROR", "true")
self.block_on_error = env.lower() in ("true", "1", "yes")
else:
self.block_on_error = block_on_error
self.timeout = _DEFAULT_TIMEOUT if timeout is None else timeout
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
)
if "supported_event_hooks" not in kwargs:
kwargs["supported_event_hooks"] = [
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
]
super().__init__(**kwargs)
@staticmethod
def get_config_model() -> type["GuardrailConfigModel"] | None:
from litellm.types.proxy.guardrails.guardrail_hooks.singulr import (
SingulrGuardrailConfigModel,
)
return SingulrGuardrailConfigModel
def _build_payload(
self,
request_data: dict[str, Any],
inputs: GenericGuardrailAPIInputs,
input_type: str,
) -> dict[str, Any]:
if not request_data:
texts = inputs.get("texts", [])
payload = SingulrGuardrailPayload(
input_type=input_type,
is_playground_request=True,
playground_text=texts[0] if texts else None,
)
else:
response = request_data.get("response")
singulr_req_object = SingulrGuardrailRequest(
model=request_data.get("model"),
messages=request_data.get("messages"),
tools=request_data.get("tools"),
model_response=response.model_dump(mode="json") if input_type == "response" and response else None,
litellm_metadata=request_data.get("litellm_metadata"),
)
payload = SingulrGuardrailPayload(
litellm_call_id=request_data.get("litellm_call_id"),
request_data=singulr_req_object,
input_type=input_type,
)
return payload.model_dump(mode="json")
def _build_headers(self) -> dict[str, str]:
return dict(
(header, value)
for header, value in (
("Content-Type", "application/json"),
("X-Singulr-Gateway-Token", self.singulr_api_key),
(
"X-Singulr-Enforcement-Entity-Id",
self.singulr_application_id or "",
),
("X-Singulr-Guardrail-Id", self.singulr_guardrail_id or ""),
)
if value
)
async def _call_api(self, payload: dict[str, Any]) -> SingulrGuardrailResponse | None:
endpoint = f"{self.singulr_api_base}{_GUARD_ENDPOINT}"
verbose_proxy_logger.debug("Singulr: %s", endpoint)
try:
response = await self.async_handler.post(
url=endpoint,
headers=self._build_headers(),
json=payload,
timeout=self.timeout,
)
response.raise_for_status()
result = SingulrGuardrailResponse.model_validate(response.json())
verbose_proxy_logger.debug("Singulr: result=%s", result)
return result
except httpx.HTTPStatusError as exc:
verbose_proxy_logger.error(
"Singulr API returned HTTP %s: %s",
exc.response.status_code,
str(exc),
)
if self.block_on_error:
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message=(f"Singulr API returned HTTP {exc.response.status_code}: {exc.response.text}"),
) from exc
return None
except httpx.TransportError as exc:
verbose_proxy_logger.error("Singulr API unreachable: %s", str(exc))
if self.block_on_error:
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message=f"Singulr API unreachable (block_on_error=True): {exc}",
) from exc
return None
except (ValueError, pydantic.ValidationError) as exc:
verbose_proxy_logger.error("Singulr API returned an invalid response: %s", str(exc))
if self.block_on_error:
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message=f"Singulr API returned an invalid response: {exc}",
) from exc
return None
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: str,
logging_obj: "LiteLLMLoggingObj | None" = None,
) -> GenericGuardrailAPIInputs:
payload = self._build_payload(request_data, inputs, input_type)
if not payload:
return inputs
result = await self._call_api(payload)
if result is None:
return inputs
verbose_proxy_logger.debug(
"Singulr: should_block=%s blocking_due_to=%s",
result.should_block,
result.blocking_due_to,
)
if result.should_block:
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message=f"Blocked by Singulr: {result.blocking_due_to or 'unknown'}",
)
return inputs

View file

@ -0,0 +1,71 @@
from typing import TYPE_CHECKING
import litellm
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .straiker import StraikerGuardrail
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
_OPTIONAL_INIT_FIELDS = (
"timeout",
"max_retries",
"initial_backoff",
"max_backoff",
"unreachable_fallback",
"fail_on_error",
"max_payload_bytes",
"custom_headers",
"metadata",
"verbose",
)
def _get_config_value(litellm_params: "LitellmParams", optional_params: object, attribute_name: str) -> object:
if optional_params is not None:
if isinstance(optional_params, dict):
value = optional_params.get(attribute_name)
else:
value = getattr(optional_params, attribute_name, None)
if value is not None:
return value
return getattr(litellm_params, attribute_name, None)
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
optional_params = getattr(litellm_params, "optional_params", None)
api_key = litellm_params.api_key
if not api_key:
raise ValueError("api_key is required for straiker")
api_base = litellm_params.api_base or "https://api.prod.straiker.ai"
default_app = getattr(litellm_params, "default_app", None) or getattr(litellm_params, "source", None)
source = default_app if isinstance(default_app, str) and default_app else "LiteLLM Gateway"
kwargs: dict[str, object] = {
field: value
for field in _OPTIONAL_INIT_FIELDS
for value in [_get_config_value(litellm_params, optional_params, field)]
if value is not None
}
_callback = StraikerGuardrail(
api_key=api_key,
api_base=api_base if isinstance(api_base, str) else "https://api.prod.straiker.ai",
source=source,
guardrail_name=guardrail.get("guardrail_name", "straiker"),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
**kwargs,
)
litellm.logging_callback_manager.add_litellm_callback(_callback)
return _callback
guardrail_initializer_registry = {
SupportedGuardrailIntegrations.STRAIKER.value: initialize_guardrail,
}
guardrail_class_registry = {
SupportedGuardrailIntegrations.STRAIKER.value: StraikerGuardrail,
}

View file

@ -0,0 +1,541 @@
from __future__ import annotations
import asyncio
import json
import random
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Literal, NoReturn
from urllib.parse import urlsplit
import httpx
from pydantic import ValidationError
from litellm._logging import verbose_proxy_logger
from litellm._version import version as litellm_version
from litellm.exceptions import (
BadRequestError,
GuardrailRaisedException,
ModifyResponseException,
Timeout,
)
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
get_session_id_from_request_data,
log_guardrail_information,
)
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
STRAIKER_WEBHOOK_SCHEMA_VERSION,
StraikerGuardrailConfigModel,
StraikerWebhookApplication,
StraikerWebhookContent,
StraikerWebhookContext,
StraikerWebhookEvent,
StraikerWebhookIdentity,
StraikerWebhookRequest,
StraikerWebhookResponse,
StraikerWebhookStream,
StraikerWebhookUsage,
)
from litellm.types.utils import GenericGuardrailAPIInputs, Usage
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
GUARDRAIL_NAME = "straiker"
DEFAULT_BLOCK_MESSAGE = "Content violates policy"
DEFAULT_API_BASE = "https://api.prod.straiker.ai"
DEFAULT_MAX_PAYLOAD_BYTES = 524288
WEBHOOK_PATH = "/api/v1/detect/webhook"
RETRY_STATUS = frozenset({408, 429, 500, 502, 503, 504})
UNREACHABLE_STATUS = frozenset({502, 503, 504})
_APPLICATION_METADATA_KEYS = frozenset({"agent_id", "app_name"})
_OPAQUE_METADATA_SCALAR_TYPES = (str, int, float, bool)
@dataclass(frozen=True, slots=True)
class _WebhookFailure:
message: str
is_unreachable: bool
def _as_dict(value: object) -> dict:
return value if isinstance(value, dict) else {}
def _merged_metadata(request_data: dict) -> dict:
return {
**_as_dict(request_data.get("metadata")),
**_as_dict(request_data.get("litellm_metadata")),
}
def _as_optional_str(value: object) -> str | None:
return value if isinstance(value, str) and value else None
def _build_webhook_metadata(request_data: dict, default_metadata: dict[str, str]) -> dict[str, object] | None:
out: dict[str, object] = {}
for key, value in _as_dict(request_data.get("metadata")).items():
if key in _APPLICATION_METADATA_KEYS or key.startswith("user_api"):
continue
if key == "session_id":
continue
if isinstance(value, _OPAQUE_METADATA_SCALAR_TYPES):
out[key] = value
out.update(default_metadata)
return out or None
def _extract_identity(request_data: dict) -> StraikerWebhookIdentity:
meta = _merged_metadata(request_data)
return StraikerWebhookIdentity(
litellm_key=_as_optional_str(meta.get("user_api_key_alias"))
or _as_optional_str(meta.get("user_api_key_hash"))
or _as_optional_str(meta.get("user_api_key_token")),
litellm_team=_as_optional_str(meta.get("user_api_key_team_alias"))
or _as_optional_str(meta.get("user_api_key_team_id")),
litellm_user_id=_as_optional_str(meta.get("user_api_key_user_id")),
litellm_user_email=_as_optional_str(meta.get("user_api_key_user_email")),
litellm_org_id=_as_optional_str(meta.get("user_api_key_org_id")),
end_user_id=_as_optional_str(meta.get("user_api_key_end_user_id")),
)
def _resolve_provider(request_data: dict, model: str | None) -> str | None:
litellm_params = _as_dict(request_data.get("litellm_params"))
custom_llm_provider = request_data.get("custom_llm_provider") or litellm_params.get("custom_llm_provider")
if custom_llm_provider:
return custom_llm_provider
if not model:
return None
try:
_, provider, _, _ = get_llm_provider(
model=model,
api_base=request_data.get("api_base") or litellm_params.get("api_base"),
api_key=request_data.get("api_key") or litellm_params.get("api_key"),
)
except BadRequestError:
return None
return provider or None
def _resolve_destination(request_data: dict) -> str | None:
litellm_params = _as_dict(request_data.get("litellm_params"))
api_base = request_data.get("api_base") or litellm_params.get("api_base")
if not isinstance(api_base, str):
return None
try:
return urlsplit(api_base).hostname
except ValueError:
return None
def _resolve_call_surface(logging_obj: LiteLLMLoggingObj | None, request_data: dict) -> str:
call_type = (
(getattr(logging_obj, "call_type", None) if logging_obj is not None else None)
or request_data.get("call_type")
or request_data.get("litellm_call_type")
)
return call_type if isinstance(call_type, str) and call_type else "unknown"
def _response_finish_reason(response: Any) -> str | None:
choices = getattr(response, "choices", None)
if not isinstance(choices, list):
return None
for choice in choices:
reason = getattr(choice, "finish_reason", None)
if isinstance(reason, str) and reason:
return reason
return None
def _build_usage(response: object) -> StraikerWebhookUsage | None:
usage = getattr(response, "usage", None)
if not isinstance(usage, Usage):
return None
input_tokens = usage.prompt_tokens
output_tokens = usage.completion_tokens
if input_tokens is None and output_tokens is None:
return None
return StraikerWebhookUsage(input_tokens=input_tokens, output_tokens=output_tokens)
def _is_streamed_request(request_data: dict) -> bool:
if request_data.get("stream") is True:
return True
body = _as_dict(_as_dict(request_data.get("proxy_server_request")).get("body"))
return body.get("stream") is True
class StraikerGuardrail(CustomGuardrail):
@staticmethod
def get_config_model() -> type[GuardrailConfigModel]:
return StraikerGuardrailConfigModel
@classmethod
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
return [
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
]
def __init__(
self,
api_key: str,
api_base: str = DEFAULT_API_BASE,
source: str = "LiteLLM Gateway",
timeout: float = 5.0,
max_retries: int = 2,
initial_backoff: float = 0.1,
max_backoff: float = 2.0,
unreachable_fallback: Literal["fail_open", "fail_closed"] = "fail_closed",
fail_on_error: bool = True,
max_payload_bytes: int = DEFAULT_MAX_PAYLOAD_BYTES,
custom_headers: dict[str, str] | None = None,
metadata: dict[str, str] | None = None,
verbose: bool = False,
async_handler: httpx.AsyncClient | None = None,
**kwargs: object,
) -> None:
if not api_key:
raise ValueError("api_key must be non-empty")
if unreachable_fallback not in ("fail_open", "fail_closed"):
raise ValueError(f"unreachable_fallback must be 'fail_open' or 'fail_closed'; got {unreachable_fallback!r}")
self.api_key = api_key
self.api_base = api_base.rstrip("/")
self.source = source
self.timeout = float(timeout)
self.max_retries = max(0, int(max_retries))
self.initial_backoff = max(0.0, float(initial_backoff))
self.max_backoff = max(self.initial_backoff, float(max_backoff))
self.unreachable_fallback = unreachable_fallback
self.fail_on_error = fail_on_error
self.max_payload_bytes = int(max_payload_bytes)
self.custom_headers = dict(custom_headers) if custom_headers else {}
self.default_metadata = dict(metadata) if metadata else {}
self.verbose = bool(verbose)
self.streaming_end_of_stream_only = True
self.streaming_buffer_until_moderated = True
self.async_handler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
)
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
super().__init__(**kwargs)
def _webhook_url(self) -> str:
return f"{self.api_base}{WEBHOOK_PATH}"
def _headers(self) -> dict[str, str]:
reserved = {"authorization", "content-type", "x-straiker-webhook-format"}
extra = {k: v for k, v in self.custom_headers.items() if k.lower() not in reserved}
return {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
"X-Straiker-Webhook-Format": "litellm",
**extra,
}
def _build_application(self, request_data: dict) -> StraikerWebhookApplication:
meta = _merged_metadata(request_data)
agent_id = _as_optional_str(meta.get("agent_id"))
return StraikerWebhookApplication(
source=agent_id or self.source,
name=_as_optional_str(meta.get("app_name")),
)
def _build_context(
self,
request_data: dict,
model: str | None,
logging_obj: LiteLLMLoggingObj | None,
) -> StraikerWebhookContext:
return StraikerWebhookContext(
call_surface=_resolve_call_surface(logging_obj, request_data),
model=model,
model_provider=_resolve_provider(request_data, model),
destination=_resolve_destination(request_data),
session_id=get_session_id_from_request_data(request_data),
litellm_call_id=getattr(logging_obj, "litellm_call_id", None) if logging_obj else None,
litellm_trace_id=getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None,
litellm_version=litellm_version,
)
def _build_envelope(
self,
*,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: LiteLLMLoggingObj | None,
) -> StraikerWebhookRequest:
model = inputs.get("model") or request_data.get("model")
call_id = getattr(logging_obj, "litellm_call_id", None) if logging_obj else None
event_id = f"{call_id or 'litellm'}:{input_type}"
content = StraikerWebhookContent(
texts=list(inputs.get("texts") or []),
images=list(inputs.get("images") or []),
structured_messages=inputs.get("structured_messages"),
tools=inputs.get("tools"),
tool_calls=inputs.get("tool_calls"),
)
if input_type == "request":
event = StraikerWebhookEvent(type="pre_call", id=event_id)
return StraikerWebhookRequest(
event=event,
request=content,
context=self._build_context(request_data, model, logging_obj),
identity=_extract_identity(request_data),
application=self._build_application(request_data),
metadata=_build_webhook_metadata(request_data, self.default_metadata),
)
response_obj = request_data.get("response")
content.finish_reason = _response_finish_reason(response_obj)
original_messages = request_data.get("messages")
request_content = StraikerWebhookContent(
structured_messages=original_messages if isinstance(original_messages, list) else None,
)
phase: Literal["none", "assembled"] = "assembled" if _is_streamed_request(request_data) else "none"
event = StraikerWebhookEvent(type="post_call", id=event_id, stream=StraikerWebhookStream(phase=phase))
return StraikerWebhookRequest(
event=event,
request=request_content,
response=content,
context=self._build_context(request_data, model, logging_obj),
identity=_extract_identity(request_data),
application=self._build_application(request_data),
usage=_build_usage(response_obj),
metadata=_build_webhook_metadata(request_data, self.default_metadata),
)
async def _post_webhook(self, payload: dict) -> tuple[StraikerWebhookResponse | None, _WebhookFailure | None]:
try:
body = json.dumps(payload).encode("utf-8")
except (TypeError, ValueError, OverflowError) as error:
return None, _WebhookFailure(f"request serialization failed: {error}", is_unreachable=False)
body_bytes = len(body)
if body_bytes > self.max_payload_bytes:
return None, _WebhookFailure(
f"payload {body_bytes}B exceeds max_payload_bytes {self.max_payload_bytes}",
is_unreachable=False,
)
url = self._webhook_url()
headers = self._headers()
attempts = self.max_retries + 1
last_failure: _WebhookFailure | None = None
if self.verbose:
verbose_proxy_logger.info(
json.dumps(
{
"event": "straiker.webhook_request",
"url": url,
"bytes": body_bytes,
"payload": payload,
},
default=str,
)
)
for attempt in range(attempts):
try:
resp = await self.async_handler.post(url, content=body, headers=headers, timeout=self.timeout)
if resp.status_code == 200:
try:
body = resp.json()
parsed = StraikerWebhookResponse.model_validate(body)
except (ValidationError, json.JSONDecodeError) as ve:
return None, _WebhookFailure(f"invalid response schema: {ve}", is_unreachable=False)
if self.verbose:
verbose_proxy_logger.info(
json.dumps(
{
"event": "straiker.webhook_response",
"status_code": resp.status_code,
"body": body,
},
default=str,
)
)
return parsed, None
last_failure = _WebhookFailure(
f"HTTP {resp.status_code}: {resp.text[:200]}",
is_unreachable=resp.status_code in UNREACHABLE_STATUS,
)
if resp.status_code not in RETRY_STATUS:
return None, last_failure
except (httpx.RequestError, asyncio.TimeoutError, Timeout) as e:
last_failure = _WebhookFailure(f"{type(e).__name__}: {e}", is_unreachable=True)
except (json.JSONDecodeError, TypeError, ValueError) as e:
return None, _WebhookFailure(f"{type(e).__name__}: {e}", is_unreachable=False)
if attempt < attempts - 1:
backoff = min(self.initial_backoff * (2**attempt), self.max_backoff)
await asyncio.sleep(random.uniform(0, backoff))
return None, last_failure or _WebhookFailure("unknown error", is_unreachable=True)
def _record(
self,
*,
request_data: dict,
logging_obj: LiteLLMLoggingObj | None,
parsed: StraikerWebhookResponse,
) -> None:
if not self.verbose:
return
response_obj = request_data.get("response")
hidden = getattr(response_obj, "_hidden_params", None)
if isinstance(hidden, dict):
straiker_hidden = hidden.setdefault("straiker", {})
if isinstance(straiker_hidden, dict):
straiker_hidden.update({"action": parsed.action, "turn_id": parsed.turn_id})
def _fail(
self,
*,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
error: str,
is_unreachable: bool,
) -> GenericGuardrailAPIInputs:
fail_open = (is_unreachable and self.unreachable_fallback == "fail_open") or not self.fail_on_error
verbose_proxy_logger.error(
json.dumps(
{
"event": "straiker.error",
"input_type": input_type,
"error": error,
"fail_open": fail_open,
},
default=str,
)
)
if fail_open:
return inputs
self._block(
request_data=request_data,
input_type=input_type,
message=f"Straiker detection unavailable: {error}",
)
def _block(
self,
*,
request_data: dict,
input_type: Literal["request", "response"],
message: str,
) -> NoReturn:
if input_type == "request":
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name or GUARDRAIL_NAME,
message=message,
should_wrap_with_default_message=False,
)
raise ModifyResponseException(
message=message,
model=request_data.get("model", "unknown") or "unknown",
request_data=request_data,
guardrail_name=self.guardrail_name or GUARDRAIL_NAME,
original_response=request_data.get("response"),
)
@staticmethod
def _intervened_inputs(
inputs: GenericGuardrailAPIInputs,
parsed: StraikerWebhookResponse,
) -> GenericGuardrailAPIInputs:
return_inputs: GenericGuardrailAPIInputs = {}
return_inputs.update(inputs)
if parsed.texts is not None:
return_inputs["texts"] = parsed.texts
return return_inputs
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: LiteLLMLoggingObj | None = None,
) -> GenericGuardrailAPIInputs:
try:
envelope = self._build_envelope(
inputs=inputs,
request_data=request_data,
input_type=input_type,
logging_obj=logging_obj,
)
payload = envelope.model_dump(mode="json", exclude_none=True)
except (ValidationError, TypeError, ValueError) as error:
return self._fail(
inputs=inputs,
request_data=request_data,
input_type=input_type,
error=str(error),
is_unreachable=False,
)
parsed, failure = await self._post_webhook(payload)
if failure is not None:
return self._fail(
inputs=inputs,
request_data=request_data,
input_type=input_type,
error=failure.message,
is_unreachable=failure.is_unreachable,
)
if parsed is None:
return self._fail(
inputs=inputs,
request_data=request_data,
input_type=input_type,
error="empty response from Straiker",
is_unreachable=False,
)
self._record(request_data=request_data, logging_obj=logging_obj, parsed=parsed)
if parsed.schema_version is not None and parsed.schema_version != STRAIKER_WEBHOOK_SCHEMA_VERSION:
verbose_proxy_logger.warning(
json.dumps(
{
"event": "straiker.schema_drift",
"expected": STRAIKER_WEBHOOK_SCHEMA_VERSION,
"received": parsed.schema_version,
}
)
)
if parsed.action == "BLOCKED":
self._block(
request_data=request_data,
input_type=input_type,
message=parsed.blocked_reason or DEFAULT_BLOCK_MESSAGE,
)
if parsed.action == "GUARDRAIL_INTERVENED":
is_streamed_response = input_type == "response" and _is_streamed_request(request_data)
if parsed.texts is None or is_streamed_response:
self._block(
request_data=request_data,
input_type=input_type,
message=parsed.blocked_reason or DEFAULT_BLOCK_MESSAGE,
)
return self._intervened_inputs(inputs, parsed)
return inputs

View file

@ -7,6 +7,7 @@ This is currently in development and not yet ready for production.
import asyncio
import binascii
import os
import uuid
from datetime import datetime
from typing import (
TYPE_CHECKING,
@ -185,6 +186,69 @@ end
return results
"""
PARALLEL_ACQUIRE_SCRIPT = """
-- Atomic check-and-acquire for the max_parallel_requests concurrency gauge.
-- Each gauge key is a sorted set of per-request slot ids scored by acquire
-- time (Redis server clock). In-flight requests are counted by ZCARD after
-- pruning slots older than the slot TTL, so unlike the windowed RPM/TPM
-- counters the gauge is never reset while requests are in flight, a
-- rejected request never occupies a slot, and a slot leaked by a crashed
-- worker self-heals after the slot TTL even under continuous traffic.
--
-- KEYS: one gauge zset key per descriptor.
-- ARGV: per-key triples (limit, slot_ttl_seconds, slot_id).
-- Success: { 0, in_flight_1, ... }. Over-limit: { 1, key_index, in_flight, limit }.
local time_reply = redis.call('TIME')
local now = tonumber(time_reply[1])
for i = 1, #KEYS do
local limit = tonumber(ARGV[(i - 1) * 3 + 1])
local slot_ttl = tonumber(ARGV[(i - 1) * 3 + 2])
redis.call('ZREMRANGEBYSCORE', KEYS[i], '-inf', now - slot_ttl)
local in_flight = redis.call('ZCARD', KEYS[i])
if in_flight + 1 > limit then
return { 1, i, in_flight, limit }
end
end
local results = { 0 }
for i = 1, #KEYS do
local slot_ttl = tonumber(ARGV[(i - 1) * 3 + 2])
local slot_id = ARGV[(i - 1) * 3 + 3]
redis.call('ZADD', KEYS[i], now, slot_id)
redis.call('EXPIRE', KEYS[i], slot_ttl)
table.insert(results, redis.call('ZCARD', KEYS[i]))
end
return results
"""
PARALLEL_RELEASE_SCRIPT = """
-- Release one slot per gauge key by removing this request's slot id.
-- ZREM of an absent member (or key) is a no-op, so a release without a
-- matching acquire (proxy-side rejection, double-fired callback, slot
-- already expired) can never free a slot owned by another request.
-- KEYS: gauge zset keys. ARGV: per-key slot_id.
-- Returns the remaining in-flight count per key.
local results = {}
for i = 1, #KEYS do
redis.call('ZREM', KEYS[i], ARGV[i])
table.insert(results, redis.call('ZCARD', KEYS[i]))
end
return results
"""
PARALLEL_COUNT_SCRIPT = """
-- Read the current in-flight count per gauge key (prunes expired slots
-- first so leaked slots do not inflate the reading).
-- KEYS: gauge zset keys. ARGV: per-key slot_ttl_seconds.
local time_reply = redis.call('TIME')
local now = tonumber(time_reply[1])
local results = {}
for i = 1, #KEYS do
redis.call('ZREMRANGEBYSCORE', KEYS[i], '-inf', now - tonumber(ARGV[i]))
table.insert(results, redis.call('ZCARD', KEYS[i]))
end
return results
"""
TOKEN_INCREMENT_SCRIPT = """
local results = {}
@ -248,6 +312,19 @@ RATE_LIMIT_DESCRIPTORS_KEY = "_litellm_rate_limit_descriptors"
# mirror ``x-ratelimit-*`` headers into the SLP. Streaming exits
# common_request_processing before ``async_post_call_success_hook`` runs.
RATE_LIMIT_RESPONSE_KEY = "_litellm_proxy_rate_limit_response"
# Holds the acquisition the pre-call hook made for this request: the slot id
# plus the gauge counter keys it was registered under. The success/failure
# callbacks release only this exact acquisition: those callbacks also fire
# for requests rejected at pre-call (which never acquired a slot), and an
# id-less release would free a slot still owned by another in-flight request
# — every rejection would then raise effective concurrency above the
# configured limit.
MAX_PARALLEL_SLOT_ACQUIRED_KEY = "_litellm_max_parallel_slot_acquired"
# How long an acquired slot counts toward the in-flight total before it is
# considered leaked (worker crashed without any release callback firing) and
# pruned. Also the longest request duration the gauge can track: a request
# running longer than this stops occupying its slot.
PARALLEL_REQUEST_SLOT_TTL_SECONDS = 3600
# Stash keys live ONLY in metadata channels — never at the top level of the
# request body. Top-level keys are forwarded as body params to upstream
# providers, which reject unknown fields with 400/429 errors.
@ -258,6 +335,7 @@ _LITELLM_STASH_KEYS: Tuple[str, ...] = (
TPM_RESERVATION_RELEASED_KEY,
RATE_LIMIT_DESCRIPTORS_KEY,
RATE_LIMIT_RESPONSE_KEY,
MAX_PARALLEL_SLOT_ACQUIRED_KEY,
)
@ -274,6 +352,17 @@ class RateLimitDescriptor(TypedDict):
rate_limit: Optional[RateLimitDescriptorRateLimitObject]
class ParallelRequestGauge(TypedDict):
counter_key: str
limit: int
descriptor_key: str
class ParallelSlotAcquisition(TypedDict):
slot_id: str
counter_keys: list[str]
class RateLimitStatus(TypedDict):
code: str
current_limit: int
@ -310,10 +399,22 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
self.check_and_increment_by_n_script = (
self.internal_usage_cache.dual_cache.redis_cache.async_register_script(CHECK_AND_INCREMENT_BY_N_SCRIPT)
)
self.parallel_acquire_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
PARALLEL_ACQUIRE_SCRIPT
)
self.parallel_release_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
PARALLEL_RELEASE_SCRIPT
)
self.parallel_count_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
PARALLEL_COUNT_SCRIPT
)
else:
self.batch_rate_limiter_script = None
self.token_increment_script = None
self.check_and_increment_by_n_script = None
self.parallel_acquire_script = None
self.parallel_release_script = None
self.parallel_count_script = None
self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60))
@ -559,7 +660,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
counter_key = keys_to_fetch[i + 1]
counter_value = cache_values[i + 1]
requests_limit = key_metadata[window_key]["requests_limit"]
max_parallel_requests_limit = key_metadata[window_key]["max_parallel_requests_limit"]
tokens_limit = key_metadata[window_key]["tokens_limit"]
# Determine which limit to use for current_limit and limit_remaining
@ -568,9 +668,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
if counter_key.endswith(":requests"):
current_limit = requests_limit
rate_limit_type = "requests"
elif counter_key.endswith(":max_parallel_requests"):
current_limit = max_parallel_requests_limit
rate_limit_type = "max_parallel_requests"
elif counter_key.endswith(":tokens"):
current_limit = tokens_limit
rate_limit_type = "tokens"
@ -694,6 +791,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
parent_otel_span: Optional[Span] = None,
read_only: bool = False,
skip_tpm_check: bool = False,
parallel_slot_id: str | None = None,
) -> RateLimitResponse:
"""
Check if any of the rate limit descriptors should be rate limited.
@ -710,15 +808,122 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
``reserve_tpm_tokens`` reservation path should set this to
avoid the +1-per-key Lua / in-memory increment double-charging
the tokens counter.
``max_parallel_requests`` descriptors are enforced by the dedicated
concurrency-gauge path (``_check_parallel_request_gauges``), never by
the windowed counters. The gauge phase must stay AFTER the windowed
check so a windowed rejection never strands an acquired slot; the
reverse order would leak one gauge slot per RPM/TPM rejection.
``parallel_slot_id`` names the slot an admission registers; callers
that enforce (not read_only) should pass the id they will later
release with — when omitted, a generated slot id is used and the slot
can only be reclaimed by TTL expiry.
"""
current_time = self._get_current_time()
now = current_time.timestamp()
now_int = int(now) # Convert to integer for Redis Lua script
# Collect all keys and their metadata upfront
keys_to_fetch, key_metadata, gauges = self._collect_windowed_keys_and_gauges(
descriptors=descriptors,
skip_tpm_check=skip_tpm_check,
)
windowed_response = RateLimitResponse(overall_code="OK", statuses=[])
if keys_to_fetch:
## CHECK IN-MEMORY CACHE
cache_values = await self.internal_usage_cache.async_batch_get_cache(
keys=keys_to_fetch,
parent_otel_span=parent_otel_span,
local_only=True,
)
if cache_values is not None:
rate_limit_response = self.is_cache_list_over_limit(keys_to_fetch, cache_values, key_metadata)
if rate_limit_response["overall_code"] == "OVER_LIMIT":
return rate_limit_response
## IF under limit in-memory, check Redis
if read_only:
# READ-ONLY MODE: Just read current values without incrementing
cache_values = await self.internal_usage_cache.async_batch_get_cache(
keys=keys_to_fetch,
parent_otel_span=parent_otel_span,
local_only=False, # Check Redis too
)
# For keys that don't exist yet, set them to 0
if cache_values is None:
cache_values = []
for _ in keys_to_fetch:
cache_values.append(str(now_int) if _.endswith(":window") else 0)
elif self.batch_rate_limiter_script is not None:
# NORMAL MODE: Increment counters in Redis
# Group keys by hash tag for Redis cluster compatibility
cache_values = await self._execute_redis_batch_rate_limiter_script(
keys_to_fetch=keys_to_fetch,
now_int=now_int,
)
# update in-memory cache with new values
for i in range(0, len(cache_values), 2):
window_key = keys_to_fetch[i]
counter_key = keys_to_fetch[i + 1]
window_value = cache_values[i]
counter_value = cache_values[i + 1]
await self.internal_usage_cache.async_set_cache(
key=counter_key,
value=counter_value,
ttl=self.window_size,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
await self.internal_usage_cache.async_set_cache(
key=window_key,
value=window_value,
ttl=self.window_size,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
else:
# NORMAL MODE: In-memory sliding window (no Redis)
cache_values = await self.in_memory_cache_sliding_window(
keys=keys_to_fetch,
now_int=now_int,
window_size=self.window_size,
)
windowed_response = self.is_cache_list_over_limit(keys_to_fetch, cache_values, key_metadata)
if windowed_response["overall_code"] == "OVER_LIMIT":
return windowed_response
if not gauges:
return windowed_response
gauge_response = await self._check_parallel_request_gauges(
gauges=gauges,
slot_id=parallel_slot_id or uuid.uuid4().hex,
parent_otel_span=parent_otel_span,
read_only=read_only,
)
return RateLimitResponse(
overall_code=gauge_response["overall_code"],
statuses=[*windowed_response["statuses"], *gauge_response["statuses"]],
)
def _collect_windowed_keys_and_gauges(
self,
descriptors: list[RateLimitDescriptor],
skip_tpm_check: bool,
) -> tuple[list[str], dict[str, dict[str, Any]], list[ParallelRequestGauge]]:
"""
Split descriptors into the windowed (window_key, counter_key) fetch
list with its per-window metadata, and the concurrency gauges for
descriptors carrying a max_parallel_requests limit.
"""
keys_to_fetch: List[str] = []
key_metadata = {} # Store metadata for each key
key_metadata: dict[str, dict[str, Any]] = {}
gauges: list[ParallelRequestGauge] = []
for descriptor in descriptors:
descriptor_key = descriptor["key"]
descriptor_value = descriptor["value"]
@ -732,6 +937,17 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
window_key = f"{{{descriptor_key}:{descriptor_value}}}:window"
if max_parallel_requests_limit is not None:
gauges.append(
ParallelRequestGauge(
counter_key=self.create_rate_limit_keys(
descriptor_key, descriptor_value, "max_parallel_requests"
),
limit=int(max_parallel_requests_limit),
descriptor_key=descriptor_key,
)
)
rate_limit_set = False
if requests_limit is not None:
rpm_key = self.create_rate_limit_keys(descriptor_key, descriptor_value, "requests")
@ -741,12 +957,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
tpm_key = self.create_rate_limit_keys(descriptor_key, descriptor_value, "tokens")
keys_to_fetch.extend([window_key, tpm_key])
rate_limit_set = True
if max_parallel_requests_limit is not None:
max_parallel_requests_key = self.create_rate_limit_keys(
descriptor_key, descriptor_value, "max_parallel_requests"
)
keys_to_fetch.extend([window_key, max_parallel_requests_key])
rate_limit_set = True
if not rate_limit_set:
continue
@ -754,77 +964,252 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
key_metadata[window_key] = {
"requests_limit": (int(requests_limit) if requests_limit is not None else None),
"tokens_limit": int(tokens_limit) if tokens_limit is not None else None,
"max_parallel_requests_limit": (
int(max_parallel_requests_limit) if max_parallel_requests_limit is not None else None
),
"window_size": int(window_size),
"descriptor_key": descriptor_key,
}
return keys_to_fetch, key_metadata, gauges
## CHECK IN-MEMORY CACHE
cache_values = await self.internal_usage_cache.async_batch_get_cache(
keys=keys_to_fetch,
def _gauge_status(self, gauge: ParallelRequestGauge, in_flight: int, code: str) -> RateLimitStatus:
return RateLimitStatus(
code=code,
current_limit=gauge["limit"],
limit_remaining=max(0, gauge["limit"] - in_flight),
rate_limit_type="max_parallel_requests",
descriptor_key=gauge["descriptor_key"],
)
def _gauge_in_flight_from_cache_value(self, raw_value: Any) -> int:
"""
In-flight count from a cached gauge value: a dict of slot_id ->
acquire timestamp when the in-memory registry is authoritative, or
the mirrored integer count from the last Redis script result.
"""
if raw_value is None:
return 0
if isinstance(raw_value, dict):
cutoff = self._get_current_time().timestamp() - PARALLEL_REQUEST_SLOT_TTL_SECONDS
return sum(1 for ts in raw_value.values() if isinstance(ts, (int, float)) and ts >= cutoff)
return max(0, int(raw_value))
async def _check_parallel_request_gauges(
self,
gauges: list[ParallelRequestGauge],
slot_id: str,
parent_otel_span: Span | None = None,
read_only: bool = False,
) -> RateLimitResponse:
"""
Enforce max_parallel_requests as a concurrency gauge over a per-slot
registry: each admitted request registers ``slot_id`` with its
acquire time, and admission requires in_flight + 1 <= limit over the
unexpired slots. Unlike the windowed RPM/TPM counters, the gauge is
never reset while requests are in flight, a rejected request never
occupies a slot, and a slot leaked by a crashed worker is pruned
after PARALLEL_REQUEST_SLOT_TTL_SECONDS even under continuous
traffic. Releases remove exactly this request's slot id, so a
double-fired or unmatched release can never free another request's
slot.
"""
gauge_keys = [gauge["counter_key"] for gauge in gauges]
if read_only:
if self.parallel_count_script is not None:
try:
raw_counts = await self.parallel_count_script(
keys=gauge_keys,
args=[PARALLEL_REQUEST_SLOT_TTL_SECONDS for _ in gauges],
)
counts = [max(0, int(value)) for value in raw_counts]
except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the local mirror, never a 500
verbose_proxy_logger.warning(f"parallel_count_script failed, using local mirror: {str(e)}")
counts = await self._read_local_gauge_counts(gauge_keys, parent_otel_span)
else:
counts = await self._read_local_gauge_counts(gauge_keys, parent_otel_span)
statuses = []
overall_code = "OK"
for gauge, in_flight in zip(gauges, counts):
code = "OVER_LIMIT" if in_flight >= gauge["limit"] else "OK"
if code == "OVER_LIMIT":
overall_code = "OVER_LIMIT"
statuses.append(self._gauge_status(gauge, in_flight, code))
return RateLimitResponse(overall_code=overall_code, statuses=statuses)
local_counts = await self._read_local_gauge_counts(gauge_keys, parent_otel_span)
for gauge, in_flight in zip(gauges, local_counts):
if in_flight >= gauge["limit"]:
return RateLimitResponse(
overall_code="OVER_LIMIT",
statuses=[self._gauge_status(gauge, in_flight, "OVER_LIMIT")],
)
if self.parallel_acquire_script is not None:
try:
raw = await self.parallel_acquire_script(
keys=gauge_keys,
args=[
arg for gauge in gauges for arg in (gauge["limit"], PARALLEL_REQUEST_SLOT_TTL_SECONDS, slot_id)
],
)
except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to in-memory enforcement, never a 500
verbose_proxy_logger.warning(
f"parallel_acquire_script failed, falling back to in-memory gauge: {str(e)}"
)
async with self._check_and_increment_lock:
return await self._acquire_parallel_slots_in_memory(gauges, slot_id, parent_otel_span)
if int(raw[0]) == 1:
gauge = gauges[int(raw[1]) - 1]
return RateLimitResponse(
overall_code="OVER_LIMIT",
statuses=[self._gauge_status(gauge, int(raw[2]), "OVER_LIMIT")],
)
statuses = []
for gauge, in_flight in zip(gauges, raw[1:]):
await self.internal_usage_cache.async_set_cache(
key=gauge["counter_key"],
value=int(in_flight),
ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
statuses.append(self._gauge_status(gauge, int(in_flight), "OK"))
return RateLimitResponse(overall_code="OK", statuses=statuses)
async with self._check_and_increment_lock:
return await self._acquire_parallel_slots_in_memory(gauges, slot_id, parent_otel_span)
async def _read_local_gauge_counts(
self,
gauge_keys: list[str],
parent_otel_span: Span | None = None,
) -> list[int]:
values = await self.internal_usage_cache.async_batch_get_cache(
keys=gauge_keys,
parent_otel_span=parent_otel_span,
local_only=True,
)
if values is None:
return [0 for _ in gauge_keys]
return [self._gauge_in_flight_from_cache_value(value) for value in values]
if cache_values is not None:
rate_limit_response = self.is_cache_list_over_limit(keys_to_fetch, cache_values, key_metadata)
if rate_limit_response["overall_code"] == "OVER_LIMIT":
return rate_limit_response
async def _acquire_parallel_slots_in_memory(
self,
gauges: list[ParallelRequestGauge],
slot_id: str,
parent_otel_span: Span | None = None,
) -> RateLimitResponse:
"""
All-or-nothing in-memory slot-registry acquire. Caller holds the lock.
## IF under limit in-memory, check Redis
if read_only:
# READ-ONLY MODE: Just read current values without incrementing
cache_values = await self.internal_usage_cache.async_batch_get_cache(
keys=keys_to_fetch,
parent_otel_span=parent_otel_span,
local_only=False, # Check Redis too
A cached dict is the authoritative in-memory registry. A cached
integer is the count mirrored from the last successful Redis script
call: when Redis fails over to this path, that mirror still counts
the slots in flight on the Redis side, so it is carried forward as
an integer counter (not discarded as an empty registry, which would
briefly double the admitted concurrency during a Redis outage).
"""
now = self._get_current_time().timestamp()
cutoff = now - PARALLEL_REQUEST_SLOT_TTL_SECONDS
states: list[tuple[dict[str, float] | None, int]] = []
for gauge in gauges:
raw_value = await self.internal_usage_cache.async_get_cache(
key=gauge["counter_key"],
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
if isinstance(raw_value, dict):
registry: dict[str, float] | None = {
key: float(ts) for key, ts in raw_value.items() if isinstance(ts, (int, float)) and ts >= cutoff
}
in_flight = len(registry or {})
elif raw_value is None:
registry = {}
in_flight = 0
else:
registry = None
in_flight = max(0, int(raw_value))
if in_flight + 1 > gauge["limit"]:
return RateLimitResponse(
overall_code="OVER_LIMIT",
statuses=[self._gauge_status(gauge, in_flight, "OVER_LIMIT")],
)
states.append((registry, in_flight))
# For keys that don't exist yet, set them to 0
if cache_values is None:
cache_values = []
for _ in keys_to_fetch:
cache_values.append(str(now_int) if _.endswith(":window") else 0)
elif self.batch_rate_limiter_script is not None:
# NORMAL MODE: Increment counters in Redis
# Group keys by hash tag for Redis cluster compatibility
cache_values = await self._execute_redis_batch_rate_limiter_script(
keys_to_fetch=keys_to_fetch,
now_int=now_int,
statuses = []
for gauge, (registry, in_flight) in zip(gauges, states):
new_value: Union[dict[str, float], int] = (
{**registry, slot_id: now} if registry is not None else in_flight + 1
)
await self.internal_usage_cache.async_set_cache(
key=gauge["counter_key"],
value=new_value,
ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
statuses.append(self._gauge_status(gauge, in_flight + 1, "OK"))
return RateLimitResponse(overall_code="OK", statuses=statuses)
# update in-memory cache with new values
for i in range(0, len(cache_values), 2):
window_key = keys_to_fetch[i]
counter_key = keys_to_fetch[i + 1]
window_value = cache_values[i]
counter_value = cache_values[i + 1]
async def _release_parallel_request_slots(
self,
acquisition: ParallelSlotAcquisition,
parent_otel_span: Span | None = None,
) -> None:
"""
Release the max_parallel_requests slots acquired at pre-call by
removing this request's slot id from every gauge it was registered
under. Removing an absent slot id is a no-op, so a release without a
matching acquire or a double-fired release can never free another
request's slot. The in-memory fallback decrements integer mirror
values (floored at 0) because the mirror carries no per-slot ids.
"""
counter_keys = acquisition["counter_keys"]
slot_id = acquisition["slot_id"]
if not counter_keys or not slot_id:
return
if self.parallel_release_script is not None:
try:
raw = await self.parallel_release_script(
keys=counter_keys,
args=[slot_id for _ in counter_keys],
)
for counter_key, remaining in zip(counter_keys, raw):
await self.internal_usage_cache.async_set_cache(
key=counter_key,
value=max(0, int(remaining)),
ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
return
except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the in-memory release, never a 500
verbose_proxy_logger.warning(
f"parallel_release_script failed, falling back to in-memory release: {str(e)}"
)
async with self._check_and_increment_lock:
for counter_key in counter_keys:
raw_value = await self.internal_usage_cache.async_get_cache(
key=counter_key,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
if isinstance(raw_value, dict):
if slot_id not in raw_value:
continue
new_value: Union[dict[str, float], int] = {
key: ts for key, ts in raw_value.items() if key != slot_id
}
elif raw_value is None:
continue
else:
new_value = max(0, int(raw_value) - 1)
await self.internal_usage_cache.async_set_cache(
key=counter_key,
value=counter_value,
ttl=self.window_size,
value=new_value,
ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
await self.internal_usage_cache.async_set_cache(
key=window_key,
value=window_value,
ttl=self.window_size,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
else:
# NORMAL MODE: In-memory sliding window (no Redis)
cache_values = await self.in_memory_cache_sliding_window(
keys=keys_to_fetch,
now_int=now_int,
window_size=self.window_size,
)
rate_limit_response = self.is_cache_list_over_limit(keys_to_fetch, cache_values, key_metadata)
return rate_limit_response
async def atomic_check_and_increment_by_n(
self,
@ -2027,10 +2412,18 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# shrinking the effective TPM budget by N and causing
# false-positive 429s under bursts. When reservation is disabled,
# this pass enforces TPM directly from the post-call counters.
parallel_counter_keys = [
self.create_rate_limit_keys(d["key"], d["value"], "max_parallel_requests")
for d in descriptors
if (d.get("rate_limit") or {}).get("max_parallel_requests") is not None
]
parallel_slot_id = uuid.uuid4().hex if parallel_counter_keys else None
response = await self.should_rate_limit(
descriptors=descriptors,
parent_otel_span=user_api_key_dict.parent_otel_span,
skip_tpm_check=self.tpm_reservation_enabled,
parallel_slot_id=parallel_slot_id,
)
if response["overall_code"] == "OVER_LIMIT":
@ -2049,6 +2442,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
key=RATE_LIMIT_RESPONSE_KEY,
value=response,
)
if parallel_slot_id is not None:
self._stash_value_in_metadata_channels(
data=data,
key=MAX_PARALLEL_SLOT_ACQUIRED_KEY,
value={
"slot_id": parallel_slot_id,
"counter_keys": parallel_counter_keys,
},
)
# ----------------------------------------------------------------
# TPM token reservation
@ -2108,6 +2510,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
if tpm_response["overall_code"] == "OVER_LIMIT":
acquisition = self._get_parallel_slot_acquisition(kwargs=data)
if acquisition is not None:
await self._release_parallel_request_slots(
acquisition=acquisition,
parent_otel_span=user_api_key_dict.parent_otel_span,
)
self._clear_parallel_slot_marker(data)
self._handle_rate_limit_error(
response=tpm_response,
descriptors=descriptors,
@ -2480,6 +2889,50 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"""True if a prior callback already refunded this request's reservation."""
return bool(cls._lookup_stashed_value(kwargs, standard_logging_metadata, TPM_RESERVATION_RELEASED_KEY))
@classmethod
def _get_parallel_slot_acquisition(
cls,
kwargs: Any,
standard_logging_metadata: dict[str, Any] | None = None,
) -> ParallelSlotAcquisition | None:
"""The slot acquisition this request's pre-call hook made, if any."""
candidate = cls._lookup_stashed_value(kwargs, standard_logging_metadata, MAX_PARALLEL_SLOT_ACQUIRED_KEY)
if not isinstance(candidate, dict):
return None
slot_id = candidate.get("slot_id")
counter_keys = candidate.get("counter_keys")
if not isinstance(slot_id, str) or not slot_id:
return None
if not isinstance(counter_keys, list) or not counter_keys:
return None
if not all(isinstance(key, str) and key for key in counter_keys):
return None
return ParallelSlotAcquisition(slot_id=slot_id, counter_keys=counter_keys)
@staticmethod
def _clear_parallel_slot_marker(data: Any) -> None:
"""
Remove the acquired-slot marker from every metadata channel a sibling
callback might read, so one release per acquire is an invariant even
when multiple callbacks fire for the same request.
"""
if not isinstance(data, dict):
return
for channel in ("metadata", "litellm_metadata"):
channel_dict = data.get(channel)
if isinstance(channel_dict, dict):
channel_dict.pop(MAX_PARALLEL_SLOT_ACQUIRED_KEY, None)
litellm_params = data.get("litellm_params")
if isinstance(litellm_params, dict):
lp_metadata = litellm_params.get("metadata")
if isinstance(lp_metadata, dict):
lp_metadata.pop(MAX_PARALLEL_SLOT_ACQUIRED_KEY, None)
slo = data.get("standard_logging_object")
if isinstance(slo, dict):
slo_meta = slo.get("metadata")
if isinstance(slo_meta, dict):
slo_meta.pop(MAX_PARALLEL_SLOT_ACQUIRED_KEY, None)
@staticmethod
def _mark_reservation_released(data: Any) -> None:
"""
@ -2621,7 +3074,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
standard_logging_object = kwargs.get("standard_logging_object") or {}
standard_logging_metadata = standard_logging_object.get("metadata") or {}
user_api_key = standard_logging_metadata.get("user_api_key_hash")
model_group = get_model_group_from_litellm_kwargs(kwargs)
# Get total tokens from response
@ -2658,20 +3110,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
pipeline_operations: List[RedisPipelineIncrementOperation] = []
# max_parallel_requests is its own counter (api-key only) — always decrement.
if user_api_key:
pipeline_operations.append(
RedisPipelineIncrementOperation(
key=self.create_rate_limit_keys(
key="api_key",
value=user_api_key,
rate_limit_type="max_parallel_requests",
),
increment_value=-1,
ttl=self.window_size,
)
)
# ----------------------------------------------------------------
# TPM reconciliation
# Per-scope behavior:
@ -2719,6 +3157,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
try:
verbose_proxy_logger.debug("INSIDE parallel request limiter ASYNC SUCCESS LOGGING")
standard_logging_object = kwargs.get("standard_logging_object") or {}
standard_logging_metadata = standard_logging_object.get("metadata") or {}
acquisition = self._get_parallel_slot_acquisition(
kwargs=kwargs,
standard_logging_metadata=standard_logging_metadata,
)
if acquisition is not None:
await self._release_parallel_request_slots(
acquisition=acquisition,
parent_otel_span=litellm_parent_otel_span,
)
self._clear_parallel_slot_marker(kwargs)
pipeline_operations = self._build_success_event_pipeline_operations(
kwargs=kwargs,
response_obj=response_obj,
@ -2855,22 +3306,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
litellm_parent_otel_span: Union[Span, None] = _get_parent_otel_span_from_kwargs(kwargs)
standard_logging_object = kwargs.get("standard_logging_object") or {}
standard_logging_metadata = standard_logging_object.get("metadata") or {}
user_api_key = standard_logging_metadata.get("user_api_key_hash")
pipeline_operations: List[RedisPipelineIncrementOperation] = []
if user_api_key:
pipeline_operations.append(
RedisPipelineIncrementOperation(
key=self.create_rate_limit_keys(
key="api_key",
value=user_api_key,
rate_limit_type="max_parallel_requests",
),
increment_value=-1,
ttl=self.window_size,
)
acquisition = self._get_parallel_slot_acquisition(
kwargs=kwargs,
standard_logging_metadata=standard_logging_metadata,
)
if acquisition is not None:
await self._release_parallel_request_slots(
acquisition=acquisition,
parent_otel_span=litellm_parent_otel_span,
)
self._clear_parallel_slot_marker(kwargs)
# Skip the reservation refund if async_post_call_failure_hook
# already released it (proxy-level rejection that also bubbles up
@ -2920,40 +3368,35 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
except Exception as e:
verbose_proxy_logger.exception(f"Error in rate limit failure event: {str(e)}")
async def async_release_max_parallel_requests_on_disconnect(self, user_api_key_dict: UserAPIKeyAuth) -> None:
async def async_release_max_parallel_requests_on_disconnect(
self,
user_api_key_dict: UserAPIKeyAuth,
request_data: dict | None = None,
) -> None:
"""
Release the api-key ``max_parallel_requests`` slot that
``async_pre_call_hook`` reserved, for a request that ended without
``async_pre_call_hook`` acquired, for a request that ended without
either logging callback firing.
The +1 is normally undone by ``async_log_success_event`` (natural
The slot is normally released by ``async_log_success_event`` (natural
stream completion) or ``async_log_failure_event`` (LLM error). When a
client cancels a stream mid-flight, the cancellation surfaces as
``asyncio.CancelledError`` / ``GeneratorExit`` and neither callback
runs, so without this the counter leaks one slot per cancelled stream
until the key wedges at its limit.
runs, so without this the slot leaks per cancelled stream until its
TTL prunes it. ``request_data`` carries the stashed acquisition;
its presence (not the key object's current max_parallel_requests
configuration, which can change mid-request) decides whether there
is anything to release.
"""
if not user_api_key_dict.api_key or user_api_key_dict.max_parallel_requests is None:
acquisition = self._get_parallel_slot_acquisition(kwargs=request_data)
if acquisition is None:
return
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
increment_list=[
RedisPipelineIncrementOperation(
key=self.create_rate_limit_keys(
key="api_key",
value=user_api_key_dict.api_key,
rate_limit_type="max_parallel_requests",
),
increment_value=-1,
# Refresh the window TTL on the decrement, matching the
# failure path. max_parallel_requests is a concurrency
# gauge, not a rolling-window count, so the key must
# outlive in-flight requests rather than expire mid-stream.
ttl=self.window_size,
)
],
litellm_parent_otel_span=None,
await self._release_parallel_request_slots(
acquisition=acquisition,
parent_otel_span=None,
)
self._clear_parallel_slot_marker(request_data)
async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response):
"""
@ -3002,17 +3445,29 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
traceback_str: Optional[str] = None,
) -> None:
"""
Release any TPM reservation when the request is rejected after the
pre-call hook reserved tokens but before the LLM call ran (e.g. a
downstream guardrail/auth hook raised). Without this, those
reservations are stranded — async_log_failure_event is a litellm
completion-level callback and never fires for proxy-side rejections.
Release the parallel-request slot and any TPM reservation when the
request is rejected after the pre-call hook acquired them but before
the LLM call ran (e.g. a downstream guardrail/auth hook raised).
Without this, those resources are stranded — async_log_failure_event
is a litellm completion-level callback and never fires for proxy-side
rejections, so a leaked slot would occupy the gauge for the full
PARALLEL_REQUEST_SLOT_TTL_SECONDS.
Idempotent via TPM_RESERVATION_RELEASED_KEY: if both this hook and
Idempotent: the slot release clears the acquisition marker (and slot
removal is a no-op ZREM on a second run), and the TPM refund is
guarded by TPM_RESERVATION_RELEASED_KEY — if both this hook and
async_log_failure_event end up running in the same flow, only the
first refund applies.
first release/refund applies.
"""
try:
acquisition = self._get_parallel_slot_acquisition(kwargs=request_data)
if acquisition is not None:
await self._release_parallel_request_slots(
acquisition=acquisition,
parent_otel_span=user_api_key_dict.parent_otel_span,
)
self._clear_parallel_slot_marker(request_data)
if self._is_reservation_released(kwargs=request_data):
return
reserved_tokens = self._get_reserved_tokens_from_kwargs(kwargs=request_data)

View file

@ -45,6 +45,7 @@ _EXPLICIT_SESSION_HEADERS = frozenset({"x-litellm-trace-id", "x-litellm-session-
# Session-id values must be non-empty strings of alphanumerics, hyphens, or underscores
# (covers UUIDs and most common session-id formats).
_SESSION_ID_VALUE_RE = re.compile(r"^[a-zA-Z0-9_\-]{8,}$")
_ANTHROPIC_SESSION_ID_VALUE_RE = re.compile(r"^[a-zA-Z0-9_\-]+$")
def _sanitize_for_log(value: Any) -> str:
@ -426,6 +427,30 @@ def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str
)
def _get_anthropic_session_id_from_metadata(metadata: object) -> str | None:
if not isinstance(metadata, dict):
return None
user_id = metadata.get("user_id")
if isinstance(user_id, dict):
session_id = user_id.get("session_id")
if isinstance(session_id, str) and _ANTHROPIC_SESSION_ID_VALUE_RE.fullmatch(session_id):
return session_id
return None
if not isinstance(user_id, str):
return None
session_marker = "_session_"
session_marker_index = user_id.rfind(session_marker)
if session_marker_index == -1:
return None
session_id = user_id[session_marker_index + len(session_marker) :]
if not session_id or not _ANTHROPIC_SESSION_ID_VALUE_RE.fullmatch(session_id):
return None
return session_id
def is_claude_code_user_agent(user_agent: str) -> bool:
"""Claude Code identifies itself as ``claude-cli/<version> ...``; the IDE
extensions and the Agent SDK run through the same CLI and share that prefix."""
@ -935,6 +960,15 @@ class LiteLLMProxyRequestSetup:
data["litellm_session_id"] = chain_id
data["litellm_trace_id"] = chain_id
verbose_proxy_logger.debug(f"Extracted chain_id from header (trace-id/session-id): {chain_id}")
else:
body_metadata = data.get("metadata")
session_id = _get_anthropic_session_id_from_metadata(body_metadata)
if session_id:
metadata_from_headers["session_id"] = session_id
data["litellm_session_id"] = session_id
if isinstance(body_metadata, dict) and isinstance(body_metadata.get("user_id"), dict):
body_metadata["user_id"] = session_id
verbose_proxy_logger.debug("Extracted session_id from Anthropic metadata.user_id")
if isinstance(data[_metadata_variable_name], dict):
data[_metadata_variable_name].update(metadata_from_headers)

View file

@ -723,12 +723,13 @@ if MCP_AVAILABLE:
"""
from litellm.proxy._experimental.mcp_server.server import _list_mcp_tools
tools = await _list_mcp_tools(
listing = await _list_mcp_tools(
user_api_key_auth=user_api_key_dict,
mcp_auth_header=None,
mcp_servers=None,
mcp_server_auth_headers=None,
)
tools = listing.tools
dumped_tools = [dict(tool) for tool in tools]
return {"tools": dumped_tools}

View file

@ -68,6 +68,7 @@ from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
EndpointType,
PassthroughStandardLoggingPayload,
@ -1771,6 +1772,7 @@ def create_pass_through_route(
if hasattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY):
delattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY)
setattr(endpoint_func, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True)
return endpoint_func

View file

@ -48,18 +48,18 @@ def _safe_response_text(httpx_response: httpx.Response) -> str:
class PassThroughEndpointLogging:
def __init__(self):
self.TRACKED_VERTEX_ROUTES = [
self.TRACKED_VERTEX_METHOD_ROUTES = (
"generateContent",
"streamGenerateContent",
"predict",
"rawPredict",
"streamRawPredict",
"search",
"batchPredictionJobs",
"predictLongRunning",
"embedContent",
"batchEmbedContents",
]
)
self.TRACKED_VERTEX_RESOURCE_ROUTES = ("batchPredictionJobs",)
# Anthropic
self.TRACKED_ANTHROPIC_ROUTES = ["/messages", "/v1/messages/batches"]
@ -339,11 +339,10 @@ class PassThroughEndpointLogging:
**kwargs,
)
def is_vertex_route(self, url_route: str):
for route in self.TRACKED_VERTEX_ROUTES:
if route in url_route:
return True
return False
def is_vertex_route(self, url_route: str) -> bool:
if any(f":{method}" in url_route for method in self.TRACKED_VERTEX_METHOD_ROUTES):
return True
return any(resource in url_route for resource in self.TRACKED_VERTEX_RESOURCE_ROUTES)
def is_anthropic_route(self, url_route: str):
for route in self.TRACKED_ANTHROPIC_ROUTES:

View file

@ -802,6 +802,19 @@ class ProxyInitializationHelpers:
),
envvar="MAX_REQUESTS_BEFORE_RESTART_JITTER",
)
@click.option(
"--limit_concurrency",
default=None,
type=click.IntRange(min=1),
help=(
"Set uvicorn's concurrency limit. Uvicorn counts both active tasks and "
"accepted connections and returns HTTP 503 after the limit is reached. "
"Idle connections can consume capacity, so use upstream connection/header "
"timeouts and per-client connection limits. Only applies to uvicorn "
"(ignored under --run_gunicorn / --run_hypercorn / --run_granian)."
),
envvar="LIMIT_CONCURRENCY",
)
@click.option(
"--enforce_prisma_migration_check",
is_flag=True,
@ -870,6 +883,7 @@ def run_server(
timeout_worker_healthcheck,
max_requests_before_restart,
max_requests_before_restart_jitter: Optional[int],
limit_concurrency: Optional[int],
enforce_prisma_migration_check: bool,
use_v2_migration_resolver: bool,
reload: bool,
@ -1243,6 +1257,8 @@ def run_server(
if max_requests_before_restart is not None:
uvicorn_args["limit_max_requests"] = max_requests_before_restart
if run_gunicorn is False and run_hypercorn is False and run_granian is False:
if limit_concurrency is not None:
uvicorn_args["limit_concurrency"] = limit_concurrency
if max_requests_before_restart_jitter is not None:
ProxyInitializationHelpers._apply_uvicorn_max_requests_jitter(
uvicorn_args=uvicorn_args,

View file

@ -1076,9 +1076,10 @@ async def proxy_startup_event(app: FastAPI):
# lazily by the flusher on first tick (see `_state_loaded` flag) so
# hot-reloaded routers also get their persisted priors.
if llm_router is not None and getattr(llm_router, "adaptive_routers", None):
for _ar in llm_router.adaptive_routers.values():
await _ar.load_state_from_db(prisma_client)
_ar._state_loaded = True
for _tagged_routers in llm_router.adaptive_routers.values():
for _tagged in _tagged_routers:
await _tagged.strategy.load_state_from_db(prisma_client)
_tagged.strategy._state_loaded = True
asyncio.create_task(_adaptive_router_flusher_loop())
## [Optional] Initialize dd tracer
@ -3248,16 +3249,18 @@ async def _adaptive_router_flusher_loop():
adaptive_routers = getattr(llm_router, "adaptive_routers", None) or {}
if not adaptive_routers or prisma_client is None:
continue
for ar in adaptive_routers.values():
# Lazy state load: covers adaptive routers registered via
# `/config/reload` after proxy boot.
if not getattr(ar, "_state_loaded", False):
try:
await ar.load_state_from_db(prisma_client)
finally:
ar._state_loaded = True
await ar.queue.flush_state_to_db(prisma_client)
await ar.queue.flush_session_to_db(prisma_client)
for tagged_routers in adaptive_routers.values():
for tagged in tagged_routers:
ar = tagged.strategy
# Lazy state load: covers adaptive routers registered via
# `/config/reload` after proxy boot.
if not getattr(ar, "_state_loaded", False):
try:
await ar.load_state_from_db(prisma_client)
finally:
ar._state_loaded = True
await ar.queue.flush_state_to_db(prisma_client)
await ar.queue.flush_session_to_db(prisma_client)
except asyncio.CancelledError:
raise
except Exception:
@ -3705,22 +3708,22 @@ def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache:
litellm_config_cache.redis_cache = redis_cache
def resolve_complexity_router_plugins(
model_name: str,
complexity_router_config: dict,
def resolve_routing_plugins(
plugin_paths: list,
config_file_path: str | None,
) -> None:
source_label: str,
) -> list:
"""
Resolves `complexity_router_config["plugins"]` dotted-path strings to live
instances via `get_instance_fn` (the same convention `litellm_settings.callbacks`
uses), in place. Raises at config-load time if a path resolves to something that
doesn't implement `RoutingPlugin`, rather than deferring to a confusing
`AttributeError` on the first request that reaches the plugin pipeline.
Resolves a list of routing-plugin entries to live `RoutingPlugin` instances.
Each string entry is resolved through `get_instance_fn` (the same dotted-path
convention `litellm_settings.callbacks` uses, which resolves both local module
files next to the config and modules installed as Python packages); non-string
entries are assumed to already be instances and passed through. Raises at
config-load time if any entry resolves to something that doesn't implement
`RoutingPlugin`, rather than deferring to a confusing `AttributeError` on the
first request that reaches the plugin pipeline. `source_label` names the config
key being resolved so the error points the operator at the right place.
"""
plugin_paths = complexity_router_config.get("plugins")
if not isinstance(plugin_paths, list):
return
resolved_plugins = [
get_instance_fn(value=plugin_path, config_file_path=config_file_path)
if isinstance(plugin_path, str)
@ -3736,12 +3739,31 @@ def resolve_complexity_router_plugins(
getattr(resolved_plugin, "run", None)
):
raise ValueError(
f"complexity_router_config.plugins entry {plugin_path!r} on model {model_name!r} "
f"resolved to {resolved_plugin!r}, which does not implement the RoutingPlugin "
"interface (an async `run(context)` method). Fix the referenced module before "
"starting the proxy."
f"{source_label} entry {plugin_path!r} resolved to {resolved_plugin!r}, which does "
"not implement the RoutingPlugin interface (an async `run(context)` method). Fix the "
"referenced module before starting the proxy."
)
complexity_router_config["plugins"] = resolved_plugins
return resolved_plugins
def resolve_complexity_router_plugins(
model_name: str,
complexity_router_config: dict,
config_file_path: str | None,
) -> None:
"""
Resolves `complexity_router_config["plugins"]` dotted-path strings to live
instances in place, via `resolve_routing_plugins`.
"""
plugin_paths = complexity_router_config.get("plugins")
if not isinstance(plugin_paths, list):
return
complexity_router_config["plugins"] = resolve_routing_plugins(
plugin_paths=plugin_paths,
config_file_path=config_file_path,
source_label=f"complexity_router_config.plugins on model {model_name!r}",
)
class ProxyConfig:
@ -4871,6 +4893,12 @@ class ProxyConfig:
for k, v in router_settings.items():
if k in available_args:
if k == "plugins" and isinstance(v, list):
v = resolve_routing_plugins(
plugin_paths=v,
config_file_path=config_file_path,
source_label="router_settings.plugins",
)
router_params[k] = v
elif k in {"health_check_interval", "health_check_concurrency"}:
raise ValueError(
@ -7373,12 +7401,13 @@ async def async_data_generator(
except (asyncio.CancelledError, GeneratorExit):
# Client disconnected mid-stream. CancelledError / GeneratorExit are
# BaseException, so they bypass the success/failure logging callbacks
# that normally release the pre-call max_parallel_requests +1; release
# it here. This is the outermost generator Starlette closes on
# that normally release the pre-call max_parallel_requests +1. Flag the
# disconnect; the shielded cleanup in `finally` owns the slot release
# so it can coordinate with disconnect-time success billing and release
# exactly once. This is the outermost generator Starlette closes on
# disconnect, so it fires reliably regardless of needs_iterator_wrap
# (a nested iterator hook would only see GeneratorExit on GC).
if not stream_completed:
proxy_logging_obj._release_max_parallel_requests_on_disconnect(user_api_key_dict)
client_disconnected = True
raise
except Exception as e:
@ -7424,6 +7453,8 @@ async def async_data_generator(
response=response,
stream_completed=stream_completed,
client_disconnected=client_disconnected,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
)
@ -16010,7 +16041,11 @@ async def get_adaptive_router_state(
status_code=404,
detail={"error": "No adaptive_router is configured on this proxy."},
)
snapshots = [await ar.get_state_snapshot() for ar in llm_router.adaptive_routers.values()]
snapshots = [
await tagged.strategy.get_state_snapshot()
for tagged_routers in llm_router.adaptive_routers.values()
for tagged in tagged_routers
]
return {"routers": snapshots}

View file

@ -11,14 +11,16 @@ from typing import Any, Dict, Optional, Tuple
import orjson
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
from fastapi.responses import ORJSONResponse
from fastapi.responses import ORJSONResponse, StreamingResponse
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.proxy._types import *
from litellm.proxy.auth.auth_utils import is_request_body_safe
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
@ -604,6 +606,7 @@ async def rag_query(
general_settings,
llm_router,
proxy_config,
select_data_generator,
version,
)
@ -673,6 +676,31 @@ async def rag_query(
**request_data,
)
hidden_params = getattr(response, "_hidden_params", {}) or {}
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=hidden_params.get("litellm_call_id", None) or "",
model_id=hidden_params.get("model_id", None) or "",
cache_key=hidden_params.get("cache_key", None) or "",
api_base=hidden_params.get("api_base", None) or "",
version=version,
response_cost=hidden_params.get("response_cost", None),
request_data=request_data,
)
if isinstance(response, CustomStreamWrapper):
return StreamingResponse(
select_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
request=request,
),
media_type="text/event-stream",
headers=custom_headers,
)
fastapi_response.headers.update(custom_headers)
return response
except HTTPException:

View file

@ -34,16 +34,12 @@ def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any:
module_name = ".".join(parts[:-1])
instance_name = parts[-1]
# If config_file_path is provided, use it to determine the module spec and load the module
module_file_path = None
if config_file_path is not None:
directory = os.path.dirname(config_file_path)
module_file_path = os.path.join(directory, *module_name.split("."))
module_file_path += ".py"
# Check if the file exists before trying to load it
if not os.path.exists(module_file_path):
raise ImportError(f"Could not find module file {module_file_path}")
module_file_path = os.path.join(directory, *module_name.split(".")) + ".py"
if module_file_path is not None and os.path.exists(module_file_path):
spec = importlib.util.spec_from_file_location(module_name, module_file_path) # type: ignore
if spec is None:
raise ImportError(f"Could not find a module specification for {module_file_path}")
@ -52,7 +48,6 @@ def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any:
raise ImportError(f"Could not find a module loader for {module_file_path}")
spec.loader.exec_module(module) # type: ignore
else:
# Dynamically import the module
module = importlib.import_module(module_name)
# Get the instance from the module

View file

@ -19,6 +19,7 @@ from typing import (
Any,
AsyncGenerator,
Awaitable,
Callable,
ClassVar,
Dict,
List,
@ -49,7 +50,7 @@ from litellm.proxy._types import (
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.model_listing import ModelInfoResponse
from litellm.types.utils import CallTypes, CallTypesLiteral
from litellm.types.utils import CallTypes, CallTypesLiteral, ModelInfo
try:
from litellm_enterprise.enterprise_callbacks.send_emails.base_email import (
@ -2583,35 +2584,30 @@ class ProxyLogging:
logging_obj._deferred_stream_complete_args = None
asyncio.create_task(_deferred_cb(*_args))
def _release_max_parallel_requests_on_disconnect(self, user_api_key_dict: UserAPIKeyAuth) -> None:
async def _arelease_max_parallel_requests_on_disconnect(
self,
user_api_key_dict: UserAPIKeyAuth,
request_data: dict | None = None,
) -> None:
"""
Release the api-key max_parallel_requests slot when a streaming
response is cancelled mid-flight (client disconnect). Neither the
success nor failure logging callback fires on the resulting
CancelledError / GeneratorExit, so the pre-call +1 would otherwise
leak.
response is cancelled mid-flight (client disconnect) and no logging
callback fired for it. Neither the success nor failure callback runs on
the resulting CancelledError / GeneratorExit, so the pre-call +1 would
otherwise leak.
Must be called from the outermost streaming generator (the one
Starlette drives and closes on disconnect). A nested iterator-hook
generator only receives GeneratorExit when it is garbage collected,
which is non-deterministic, so the refund cannot live there.
Scheduled fire-and-forget (no await) because awaiting is not
permitted while unwinding a GeneratorExit.
Awaited from the shielded streaming cleanup rather than scheduled
fire-and-forget, so the caller can make it the single owner of the
release: when a disconnect-time success event does fire (partial-spend
billing or a deferred-guardrail flush), that event's own limiter
callback releases the slot and this is not called at all. Two
concurrent releases of the same acquisition would otherwise race and
double-decrement under the limiter's in-memory fallback.
"""
limiter = self.get_proxy_hook("parallel_request_limiter")
if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3):
return
try:
asyncio.create_task(limiter.async_release_max_parallel_requests_on_disconnect(user_api_key_dict))
except RuntimeError:
# No running event loop (e.g. interpreter/loop shutdown); the
# counter's window TTL will reclaim the slot.
verbose_proxy_logger.warning(
"parallel_request_limiter_v3: could not schedule "
"max_parallel_requests release on disconnect; no running "
"event loop. Slot will be reclaimed when its window TTL expires"
)
await limiter.async_release_max_parallel_requests_on_disconnect(user_api_key_dict, request_data)
def _init_response_taking_too_long_task(self, data: Optional[dict] = None):
"""
@ -6101,6 +6097,7 @@ def create_model_info_response(
include_metadata: bool = False,
fallback_type: Optional[str] = None,
llm_router: Optional["Router"] = None,
get_model_info: Callable[[str], ModelInfo] = litellm.get_model_info,
) -> ModelInfoResponse:
"""
Create a standardized OpenAI-compatible model object.
@ -6118,25 +6115,37 @@ def create_model_info_response(
"owned_by": provider,
}
# Surface context-window limits for OpenAI-compatible discovery clients.
# Only emitted when known, so wildcard routes and limitless backends stay clean.
# Limits are best-effort enrichment, so a single malformed deployment degrades
# to the base response rather than 500-ing the whole listing.
try:
model_cost_info: ModelInfo | None = get_model_info(model_id)
except Exception as e:
verbose_proxy_logger.debug(
"create_model_info_response: cost map lookup failed for %s: %s",
model_id,
e,
)
model_cost_info = None
max_input_tokens: int | None = None
max_output_tokens: int | None = None
if model_cost_info is not None:
cost_map_input = model_cost_info.get("max_input_tokens")
if cost_map_input is not None:
max_input_tokens = int(cost_map_input)
cost_map_output = model_cost_info.get("max_output_tokens")
if cost_map_output is not None:
max_output_tokens = int(cost_map_output)
if llm_router is not None:
try:
model_group_info = llm_router.get_model_group_info(model_id)
except Exception as e:
verbose_proxy_logger.debug(
"create_model_info_response: get_model_group_info failed for %s: %s",
model_id,
e,
)
model_group_info = None
if model_group_info is not None:
if model_group_info.max_input_tokens is not None:
base["max_input_tokens"] = int(model_group_info.max_input_tokens)
if model_group_info.max_output_tokens is not None:
base["max_output_tokens"] = int(model_group_info.max_output_tokens)
configured_input, configured_output = llm_router.get_configured_token_limits(model_id)
if configured_input is not None:
max_input_tokens = configured_input
if configured_output is not None:
max_output_tokens = configured_output
if max_input_tokens is not None:
base["max_input_tokens"] = max_input_tokens
if max_output_tokens is not None:
base["max_output_tokens"] = max_output_tokens
if not include_metadata:
return base

View file

@ -11,12 +11,14 @@ __all__ = ["ingest", "aingest", "query", "aquery"]
import asyncio
import contextvars
from contextlib import contextmanager
from functools import partial
from typing import (
TYPE_CHECKING,
Any,
Coroutine,
Dict,
Iterator,
List,
Optional,
Tuple,
@ -27,6 +29,9 @@ from typing import (
import httpx
import litellm
from litellm._internal_context import is_internal_call
from litellm.cost_calculator import vector_store_search_cost
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
from litellm.rag.ingestion.bedrock_ingestion import BedrockRAGIngestion
from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion
@ -188,6 +193,25 @@ async def aingest(
)
@contextmanager
def _suppressed_sub_call_billing() -> Iterator[None]:
"""
Suppress a sub-call's own billing event so the parent aquery event bills it.
Every suppressed sub-call's cost must be folded into the parent event:
into the response's hidden response_cost on the non-streaming path, or via
the logging object's additional_response_cost on the streaming path (the
streamed cost is computed from assembled chunks after this pipeline
returns, so there is no response object to fold into here).
"""
previous = is_internal_call.get()
is_internal_call.set(True)
try:
yield
finally:
is_internal_call.set(previous)
async def _execute_query_pipeline(
model: str,
messages: List[Any],
@ -209,27 +233,46 @@ async def _execute_query_pipeline(
raise ValueError("No query found in messages for RAG query")
# 2. Search vector store
search_response = await litellm.vector_stores.asearch(
vector_store_id=retrieval_config["vector_store_id"],
query=query_text,
max_num_results=retrieval_config.get("top_k", 10),
custom_llm_provider=retrieval_config.get("custom_llm_provider", "openai"),
**kwargs,
)
with _suppressed_sub_call_billing():
search_response = await litellm.vector_stores.asearch(
vector_store_id=retrieval_config["vector_store_id"],
query=query_text,
max_num_results=retrieval_config.get("top_k", 10),
custom_llm_provider=retrieval_config.get("custom_llm_provider", "openai"),
**kwargs,
)
search_provider = retrieval_config.get("custom_llm_provider", "openai")
try:
search_cost = sum(
vector_store_search_cost(
model=search_provider if "/" in search_provider else None,
custom_llm_provider=search_provider,
response=search_response,
)
)
except Exception: # noqa: BLE001 - cost accounting must never break the query path
search_cost = 0.0
rerank_response = None
rerank_cost = 0.0
context_chunks = search_response.get("data", [])
# 3. Optional rerank
if rerank and rerank.get("enabled"):
documents = RAGQuery.extract_documents_from_search(search_response)
if documents:
rerank_response = await litellm.arerank(
model=rerank["model"],
query=query_text,
documents=documents,
top_n=rerank.get("top_n", 5),
)
with _suppressed_sub_call_billing():
rerank_response = await litellm.arerank(
model=rerank["model"],
query=query_text,
documents=documents,
top_n=rerank.get("top_n", 5),
)
rerank_hidden_params = getattr(rerank_response, "_hidden_params", None)
if isinstance(rerank_hidden_params, dict):
rerank_response_cost: float | None = rerank_hidden_params.get("response_cost")
rerank_cost = rerank_response_cost or 0.0
context_chunks = RAGQuery.get_top_chunks_from_rerank(search_response, rerank_response)
# 4. Build context message and call completion
@ -237,28 +280,40 @@ async def _execute_query_pipeline(
modified_messages = messages[:-1] + [context_message] + [messages[-1]]
# Use router if available to properly resolve virtual model names
if router is not None:
response = await router.acompletion(
model=model,
messages=modified_messages,
stream=stream,
**kwargs,
)
else:
response = await litellm.acompletion(
model=model,
messages=modified_messages,
stream=stream,
**kwargs,
)
with _suppressed_sub_call_billing():
if router is not None:
response = await router.acompletion(
model=model,
messages=modified_messages,
stream=stream,
**kwargs,
)
else:
response = await litellm.acompletion(
model=model,
messages=modified_messages,
stream=stream,
**kwargs,
)
# 5. Attach search results to response
sub_call_cost = search_cost + rerank_cost
if not stream and isinstance(response, ModelResponse):
response = RAGQuery.add_search_results_to_response(
response=response,
search_results=search_response,
rerank_results=rerank_response,
)
if sub_call_cost > 0:
hidden_params = getattr(response, "_hidden_params", None)
if isinstance(hidden_params, dict):
completion_response_cost: float | None = hidden_params.get("response_cost")
if completion_response_cost is not None:
hidden_params["response_cost"] = completion_response_cost + sub_call_cost
elif sub_call_cost > 0:
logging_obj: object = kwargs.get("litellm_logging_obj")
if isinstance(logging_obj, LiteLLMLoggingObj):
logging_obj.model_call_details["additional_response_cost"] = sub_call_cost
return response # type: ignore[return-value]

View file

@ -260,7 +260,7 @@ class LiteLLM_Proxy_MCP_Handler:
# names), so use None and let the auth object's mcp_servers do the filtering.
effective_server_filter = None if resolved_toolset_ids else (resolved_mcp_servers or None)
tools = await _get_tools_from_mcp_servers(
listing = await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=effective_server_filter,
@ -270,6 +270,7 @@ class LiteLLM_Proxy_MCP_Handler:
litellm_trace_id=litellm_trace_id,
request_tags=request_tags,
)
tools = listing.tools
allowed_mcp_server_ids = await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth)
allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( # type: ignore[attr-defined]

View file

@ -33,6 +33,7 @@ from typing import (
Optional,
Set,
Tuple,
TypeVar,
Union,
cast,
)
@ -86,7 +87,11 @@ from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler
from litellm.router_strategy.lowest_tpm_rpm import LowestTPMLoggingHandler
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2
from litellm.router_strategy.simple_shuffle import simple_shuffle
from litellm.router_strategy.tag_based_routing import get_deployments_for_tag
from litellm.router_strategy.tag_based_routing import (
_get_tags_from_request_kwargs,
get_deployments_for_tag,
is_valid_deployment_tag,
)
from litellm.router_utils.add_retry_fallback_headers import (
_HiddenParamsHost,
add_fallback_headers_to_response,
@ -175,6 +180,7 @@ from litellm.types.router import (
MockRouterTestingParams,
ModelGroupInfo,
OptionalPreCallChecks,
PreRoutingStrategy,
RetryPolicy,
RouterCacheEnum,
RouterGeneralSettings,
@ -186,6 +192,7 @@ from litellm.types.router import (
RoutingPlugin,
RoutingStrategy,
SearchToolTypedDict,
TaggedPreRoutingStrategy,
)
from litellm.types.services import ServiceTypes
from litellm.types.utils import (
@ -260,6 +267,9 @@ def _cost_value_as_float(value: Union[str, int, float, None]) -> Optional[float]
return None
_PreRoutingStrategyT = TypeVar("_PreRoutingStrategyT")
class RoutingArgs(enum.Enum):
ttl = 60 # 1min (RPM/TPM expire key)
@ -487,10 +497,10 @@ class Router:
self.provider_default_deployment_ids: List[str] = []
self.pattern_router = PatternMatchRouter()
self.team_pattern_routers: Dict[str, PatternMatchRouter] = {} # {"TEAM_ID": PatternMatchRouter}
self.auto_routers: Dict[str, "AutoRouter"] = {}
self.complexity_routers: Dict[str, "ComplexityRouter"] = {}
self.adaptive_routers: Dict[str, "AdaptiveRouter"] = {}
self.quality_routers: Dict[str, "QualityRouter"] = {}
self.auto_routers: dict[str, list[TaggedPreRoutingStrategy["AutoRouter"]]] = {}
self.complexity_routers: dict[str, list[TaggedPreRoutingStrategy["ComplexityRouter"]]] = {}
self.adaptive_routers: dict[str, list[TaggedPreRoutingStrategy["AdaptiveRouter"]]] = {}
self.quality_routers: dict[str, list[TaggedPreRoutingStrategy["QualityRouter"]]] = {}
self.routing_plugins: list[RoutingPlugin] = list(plugins) if plugins else []
# Initialize model_group_alias early since it's used in set_model_list
@ -2037,6 +2047,9 @@ class Router:
logging_obj=model_response.logging_obj,
)
self._async_generator = async_generator
inner_chunks: object = getattr(model_response, "chunks", None)
if isinstance(inner_chunks, list):
self.chunks = inner_chunks
# Preserve hidden params (including litellm_overhead_time_ms) from original response
if hasattr(model_response, "_hidden_params"):
self._hidden_params = model_response._hidden_params.copy()
@ -7568,6 +7581,11 @@ class Router:
return True
return False
@staticmethod
def _deployment_tags(deployment: Deployment) -> tuple[str, ...]:
"""Deployment tags used to disambiguate strategy registries keyed by model_name."""
return tuple(deployment.litellm_params.tags or ())
def init_auto_router_deployment(self, deployment: Deployment):
"""
Initialize the auto-router deployment.
@ -7603,11 +7621,12 @@ class Router:
embedding_model=embedding_model,
litellm_router_instance=self,
)
if deployment.model_name in self.auto_routers:
raise ValueError(
f"Auto-router deployment {deployment.model_name} already exists. Please use a different model name."
)
self.auto_routers[deployment.model_name] = autor_router
self._register_pre_routing_strategy(
registry=self.auto_routers,
deployment=deployment,
strategy=autor_router,
strategy_label="Auto-router",
)
def _is_complexity_router_deployment(self, litellm_params: LiteLLM_Params) -> bool:
"""
@ -7658,20 +7677,54 @@ class Router:
litellm_router_instance=self,
complexity_router_config=complexity_router_config,
)
if deployment.model_name in self.complexity_routers:
raise ValueError(
f"Complexity-router deployment {deployment.model_name} already exists. Please use a different model name."
)
self.complexity_routers[deployment.model_name] = complexity_router
self._register_pre_routing_strategy(
registry=self.complexity_routers,
deployment=deployment,
strategy=complexity_router,
strategy_label="Complexity-router",
)
def _is_adaptive_router_deployment(self, litellm_params: LiteLLM_Params) -> bool:
"""True when this deployment opts in via the `auto_router/adaptive_router` model prefix."""
return litellm_params.model.startswith("auto_router/adaptive_router")
@staticmethod
def _has_registered_strategy(
registry: dict[str, list[TaggedPreRoutingStrategy[_PreRoutingStrategyT]]],
model_name: str,
tags: tuple[str, ...],
) -> bool:
"""True when a strategy for this (model_name, tags) pair is already registered."""
return any(existing.tags == tags for existing in registry.get(model_name, []))
def _register_pre_routing_strategy(
self,
registry: dict[str, list[TaggedPreRoutingStrategy[_PreRoutingStrategyT]]],
deployment: Deployment,
strategy: _PreRoutingStrategyT,
strategy_label: str,
) -> None:
"""
Register `strategy` under `deployment.model_name`, scoped by its tags.
Reusing a `model_name` is allowed when tags differ; a repeat of the same
(model_name, tags) pair is a misconfiguration and is rejected.
"""
tags = self._deployment_tags(deployment)
if self._has_registered_strategy(registry, deployment.model_name, tags):
raise ValueError(
f"{strategy_label} deployment {deployment.model_name} with tags {list(tags)} already exists. "
"Please use a different model name or set different tags."
)
registry[deployment.model_name] = [
*registry.get(deployment.model_name, []),
TaggedPreRoutingStrategy(tags=tags, strategy=strategy),
]
def _finalize_adaptive_router_if_configured(self) -> None:
"""Locate every adaptive-router deployment in the finalized model_list and
build an AdaptiveRouter for each. Safe no-op when none are configured.
Idempotent: skips any deployment whose model_name is already initialized."""
Idempotent: skips any deployment whose (model_name, tags) pair is already
initialized, so hot-reloads don't rebuild routers that would lose state."""
# Drop any adaptive-router hooks left over from a previous Router
# instance (e.g. after `/config/reload` replaced `llm_router`). Without
# this, stale AdaptiveRouterPostCallHook callbacks from the old Router
@ -7694,23 +7747,31 @@ class Router:
litellm_params=(lp if not isinstance(lp, dict) else LiteLLM_Params(**lp)),
model_info=(entry.get("model_info") if isinstance(entry, dict) else entry.model_info),
)
if model_name in self.adaptive_routers:
if self._has_registered_strategy(self.adaptive_routers, model_name, self._deployment_tags(deployment)):
continue
self.init_adaptive_router_deployment(deployment=deployment)
for model_name, complexity_router in self.complexity_routers.items():
if not complexity_router.config.adaptive or model_name in self.adaptive_routers:
continue
adaptive_router = complexity_router._ensure_adaptive_router()
if adaptive_router is not None:
self.adaptive_routers[model_name] = adaptive_router
for model_name, tagged_complexity_routers in self.complexity_routers.items():
for tagged in tagged_complexity_routers:
complexity_router = tagged.strategy
if not complexity_router.config.adaptive:
continue
if self._has_registered_strategy(self.adaptive_routers, model_name, tagged.tags):
continue
adaptive_router = complexity_router._ensure_adaptive_router()
if adaptive_router is not None:
self.adaptive_routers[model_name] = [
*self.adaptive_routers.get(model_name, []),
TaggedPreRoutingStrategy(tags=tagged.tags, strategy=adaptive_router),
]
for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(AdaptiveRouterPostCallHook):
litellm.logging_callback_manager.remove_callback_from_all_lists(callback)
for adaptive_router in self.adaptive_routers.values():
litellm.logging_callback_manager.add_litellm_callback(
AdaptiveRouterPostCallHook(adaptive_router=adaptive_router)
)
for tagged_adaptive_routers in self.adaptive_routers.values():
for tagged in tagged_adaptive_routers:
litellm.logging_callback_manager.add_litellm_callback(
AdaptiveRouterPostCallHook(adaptive_router=tagged.strategy)
)
def init_adaptive_router_deployment(self, deployment: Deployment) -> None:
"""
@ -7763,18 +7824,18 @@ class Router:
if cost is not None:
model_to_cost[name] = float(cost)
if deployment.model_name in self.adaptive_routers:
raise ValueError(
f"Adaptive-router deployment {deployment.model_name} already exists. Please use a different model name."
)
adaptive_router = AdaptiveRouter(
router_name=deployment.model_name,
config=config,
model_to_prefs=model_to_prefs,
model_to_cost=model_to_cost,
)
self.adaptive_routers[deployment.model_name] = adaptive_router
self._register_pre_routing_strategy(
registry=self.adaptive_routers,
deployment=deployment,
strategy=adaptive_router,
strategy_label="Adaptive-router",
)
litellm.logging_callback_manager.add_litellm_callback(
AdaptiveRouterPostCallHook(adaptive_router=adaptive_router)
)
@ -7826,11 +7887,12 @@ class Router:
litellm_router_instance=self,
quality_router_config=quality_router_config,
)
if deployment.model_name in self.quality_routers:
raise ValueError(
f"Quality-router deployment {deployment.model_name} already exists. Please use a different model name."
)
self.quality_routers[deployment.model_name] = quality_router
self._register_pre_routing_strategy(
registry=self.quality_routers,
deployment=deployment,
strategy=quality_router,
strategy_label="Quality-router",
)
def deployment_is_active_for_environment(self, deployment: Deployment) -> bool:
"""
@ -8459,6 +8521,27 @@ class Router:
raise Exception("Model Name invalid - {}".format(type(model)))
return None
def get_configured_token_limits(self, model_name: str) -> "tuple[int | None, int | None]":
"""
Return (max_input_tokens, max_output_tokens) explicitly configured in a concrete
deployment's model_info for model_name, via O(1) index lookup.
Returns (None, None) for wildcard-expanded or unknown names. Unlike
get_model_group_info, this never triggers pattern matching or deep copies, so it
is safe to call per listed model on the /v1/models hot path.
"""
deployment = self.get_deployment_by_model_group_name(model_group_name=model_name)
if deployment is None:
return (None, None)
model_info = deployment.model_info
max_input = model_info.get("max_input_tokens")
max_output = model_info.get("max_output_tokens")
return (
int(max_input) if max_input is not None else None,
int(max_output) if max_output is not None else None,
)
def get_deployment_credentials_with_provider(self, model_id: str) -> Optional[Dict[str, Any]]:
"""
Get API credentials and provider info from a model name in model_list.
@ -10810,6 +10893,35 @@ class Router:
return filtered
def _select_pre_routing_strategy(self, model: str, request_kwargs: Dict) -> "PreRoutingStrategy | None":
"""
Resolve the pre-routing strategy for `model`, disambiguating deployments
that share a `model_name` by matching the request's tags against each
registered strategy's tags before falling back to the first registered.
"""
candidates: list[TaggedPreRoutingStrategy[PreRoutingStrategy]] = [
*self.auto_routers.get(model, []),
*self.complexity_routers.get(model, []),
*self.adaptive_routers.get(model, []),
*self.quality_routers.get(model, []),
]
if not candidates:
return None
if len(candidates) == 1:
return candidates[0].strategy
request_tags = _get_tags_from_request_kwargs(request_kwargs)
if request_tags:
for tagged in candidates:
if tagged.tags and is_valid_deployment_tag(
list(tagged.tags), request_tags, self.tag_filtering_match_any
):
return tagged.strategy
for tagged in candidates:
if "default" in tagged.tags:
return tagged.strategy
return candidates[0].strategy
async def async_pre_routing_hook(
self,
model: str,
@ -10832,12 +10944,7 @@ class Router:
if self.routing_plugins:
await self._run_routing_plugins(model=model, request_kwargs=request_kwargs, messages=messages)
router_strategy = (
self.auto_routers.get(model)
or self.complexity_routers.get(model)
or self.adaptive_routers.get(model)
or self.quality_routers.get(model)
)
router_strategy = self._select_pre_routing_strategy(model=model, request_kwargs=request_kwargs)
if router_strategy is None:
return None

View file

@ -28,6 +28,7 @@ from litellm.types.utils import ModelResponse
from .config import (
DEFAULT_CODE_KEYWORDS,
DEFAULT_ESCALATION_KEYWORDS,
DEFAULT_REASONING_KEYWORDS,
DEFAULT_SIMPLE_KEYWORDS,
DEFAULT_TECHNICAL_KEYWORDS,
@ -173,6 +174,11 @@ class ComplexityRouter(CustomLogger):
self.config.custom_technical_keywords,
)
self.simple_keywords = self.config.simple_keywords or DEFAULT_SIMPLE_KEYWORDS
self.escalation_keywords = (
self.config.escalation_keywords
if self.config.escalation_keywords is not None
else DEFAULT_ESCALATION_KEYWORDS
)
# Lazily built on first semantic request and cached for reuse (route
# embeddings are static, only the prompt is embedded per request). The lock
@ -668,6 +674,53 @@ class ComplexityRouter(CustomLogger):
}
return best_model
def _escalation_triggered(self, user_message: str) -> bool:
"""Whether the prompt asks to escalate to a stronger model.
Matching is a case-sensitive substring test so the default "LITELLM ESCALATE"
only fires on the deliberate, shouted form and not on incidental lowercase
mentions of the word (e.g. "how do I escalate this ticket").
"""
if not self.escalation_keywords:
return False
return any(keyword in user_message for keyword in self.escalation_keywords)
def _tier_for_model(self, model: str) -> ComplexityTier | None:
"""Return the most-severe configured tier whose pool contains this model."""
pools = self._tier_pools()
matched = tuple(ComplexityTier(tier_name) for tier_name, models in pools.items() if model in models)
if not matched:
return None
return max(matched, key=TIER_SEVERITY_ORDER.index)
def _escalate_tier(self, tier: ComplexityTier) -> ComplexityTier:
"""Bump a tier one step up to the next-higher configured tier.
Returns the input tier unchanged when it is already the highest configured
tier, so escalation can never route below the model the user would otherwise
have received.
"""
configured = frozenset(self.config.tiers)
current_index = TIER_SEVERITY_ORDER.index(tier)
higher_tiers = tuple(
candidate for candidate in TIER_SEVERITY_ORDER[current_index + 1 :] if candidate.value in configured
)
return higher_tiers[0] if higher_tiers else tier
def _escalated_pin(self, pinned_model: str) -> str | None:
"""Bump a session's pinned model to the next-higher configured tier.
Returns None when the pin no longer maps to any configured tier, signalling
a full reclassification instead.
"""
pinned_tier = self._tier_for_model(pinned_model)
if pinned_tier is None:
return None
escalated_tier = self._escalate_tier(pinned_tier)
if escalated_tier == pinned_tier:
return pinned_model
return self.get_model_for_tier(escalated_tier)
def _lexical_tier_override(self, user_message: str) -> ComplexityTier | None:
"""When keyword_tier_rules match literally, the most-severe matched tier wins.
@ -910,29 +963,41 @@ class ComplexityRouter(CustomLogger):
if cache_key is not None:
pinned_model = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
if isinstance(pinned_model, str):
# Refresh the TTL on every hit so an active session doesn't lose its
# pin mid-conversation just because it outlives the original write.
await self.litellm_router_instance.cache.async_set_cache(
key=cache_key,
value=pinned_model,
ttl=self.config.session_affinity_ttl_seconds,
)
if self.config.adaptive:
from litellm.router_strategy.adaptive_router.config import (
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
routed_model: str | None = pinned_model
if self.escalation_keywords:
resolved_messages = self._resolve_messages(messages, request_kwargs)
user_message = (
self._extract_user_message_and_system_prompt(resolved_messages)[0]
if resolved_messages
else None
)
if user_message is not None and self._escalation_triggered(user_message):
routed_model = self._escalated_pin(pinned_model)
if routed_model is not None:
# Refresh the TTL on every hit so an active session doesn't lose its
# pin mid-conversation just because it outlives the original write.
await self.litellm_router_instance.cache.async_set_cache(
key=cache_key,
value=routed_model,
ttl=self.config.session_affinity_ttl_seconds,
)
if self.config.adaptive:
from litellm.router_strategy.adaptive_router.config import (
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
)
kwargs_metadata = request_kwargs.setdefault("metadata", {})
if isinstance(kwargs_metadata, dict):
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = pinned_model
verbose_router_logger.info(
f"ComplexityRouter: routing decision cause=session_affinity_pin, routed_model={pinned_model}"
)
has_original_messages = messages is not None and len(messages) > 0
return PreRoutingHookResponse(
model=pinned_model,
messages=messages if has_original_messages else None,
)
kwargs_metadata = request_kwargs.setdefault("metadata", {})
if isinstance(kwargs_metadata, dict):
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = routed_model
cause = "session_affinity_escalation" if routed_model != pinned_model else "session_affinity_pin"
verbose_router_logger.info(
f"ComplexityRouter: routing decision cause={cause}, routed_model={routed_model}"
)
has_original_messages = messages is not None and len(messages) > 0
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
)
response = await self._classify_and_route(
model=model,
@ -1004,13 +1069,17 @@ class ComplexityRouter(CustomLogger):
messages=messages if has_original_messages else None,
)
escalate = self._escalation_triggered(user_message)
override_tier = await self._resolve_keyword_tier_override(user_message, request_kwargs)
if override_tier is not None:
routed_model = await self._pick_model_for_tier(override_tier, messages, resolved_messages, request_kwargs)
cause = "semantic_keyword_match" if self.config.semantic_keyword_matching else "literal_keyword_match"
routed_tier = self._escalate_tier(override_tier) if escalate else override_tier
routed_model = await self._pick_model_for_tier(routed_tier, messages, resolved_messages, request_kwargs)
base_cause = "semantic_keyword_match" if self.config.semantic_keyword_matching else "literal_keyword_match"
cause = f"{base_cause}+escalation" if escalate else base_cause
verbose_router_logger.info(
f"ComplexityRouter: routing decision cause={cause}, "
f"tier={override_tier.value}, routed_model={routed_model}"
f"tier={routed_tier.value}, routed_model={routed_model}"
)
return PreRoutingHookResponse(
model=routed_model,
@ -1018,6 +1087,9 @@ class ComplexityRouter(CustomLogger):
)
tier, score, signals = await self.aclassify(user_message, system_prompt, request_kwargs)
if escalate:
tier = self._escalate_tier(tier)
signals = [*signals, "escalation"]
if self.config.adaptive:
routed_model = self._soft_floor_pick(tier, user_message, request_kwargs)
adaptive = self._ensure_adaptive_router()

View file

@ -162,6 +162,9 @@ DEFAULT_TECHNICAL_KEYWORDS: list[str] = [
# Note: "async", "kubernetes", "docker" are in DEFAULT_CODE_KEYWORDS
]
DEFAULT_ESCALATION_KEYWORDS: list[str] = ["LITELLM ESCALATE"]
DEFAULT_SIMPLE_KEYWORDS: list[str] = [
"what is",
"what's",
@ -339,6 +342,16 @@ class ComplexityRouterConfig(BaseModel):
),
)
escalation_keywords: list[str] | None = Field(
default=None,
description=(
"Case-sensitive phrases a user can include to force a bump to the next-higher "
"complexity tier when they aren't satisfied with results (they can force a stronger "
"model, but not choose which one). Defaults to ['LITELLM ESCALATE'] when unset; "
"set to an empty list to disable."
),
)
# Deterministic keyword -> tier overrides, evaluated before weighted scoring
keyword_tier_rules: list[KeywordTierRule] | None = Field(
default=None,
@ -400,6 +413,13 @@ class ComplexityRouterConfig(BaseModel):
coerced[key] = item
return coerced
@field_validator("escalation_keywords")
@classmethod
def _normalize_escalation_keywords(cls, value: list[str] | None) -> list[str] | None:
if value is None:
return None
return [stripped for keyword in value if (stripped := keyword.strip())]
@model_validator(mode="after")
def _validate_llm_classifier_config(self) -> "ComplexityRouterConfig":
if self.classifier_type == "llm" and self.classifier_llm_config is None:

View file

@ -53,6 +53,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import (
from litellm.types.proxy.guardrails.guardrail_hooks.cisco_ai_defense import (
CiscoAIDefenseGuardrailConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.singulr import (
SingulrGuardrailConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.headroom import (
HeadroomGuardrailConfigModel,
)
@ -125,8 +128,10 @@ class SupportedGuardrailIntegrations(Enum):
RUBRIK = "rubrik"
VIGIL_GUARD = "vigil_guard"
REPELLOAI = "repelloai"
SINGULR = "singulr"
HEADROOM = "headroom"
COMPRESR = "compresr"
STRAIKER = "straiker"
class Role(Enum):
@ -932,6 +937,7 @@ class LitellmParams(
HiddenlayerGuardrailConfigModel,
QostodianNexusConfigModel,
VigilGuardGuardrailConfigModel,
SingulrGuardrailConfigModel,
):
guardrail: str = Field(description="The type of guardrail integration to use")
mode: Union[str, List[str], Mode] = Field(

View file

@ -529,6 +529,7 @@ class ChatCompletionDeltaToolCallChunk(TypedDict, total=False):
class ChatCompletionCachedContent(TypedDict):
type: Literal["ephemeral"]
ttl: NotRequired[Literal["5m", "1h"]]
class ChatCompletionThinkingBlock(TypedDict, total=False):

View file

@ -11,6 +11,14 @@ LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY = "litellm_pass_through_custom_body"
# exact byte/string body, such as AWS SigV4-signed requests.
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY = "litellm_pass_through_raw_body"
# Attribute set on the FastAPI endpoint function of every user-defined pass-through
# route. Auth reads it off the dispatched endpoint (``request.scope["endpoint"]``) to
# decide whether a request body ``model`` names an upstream model rather than a
# LiteLLM-managed one. Keying off the resolved endpoint (not the request path) means a
# custom path that collides with a built-in route never suppresses model-access checks:
# on a collision FastAPI dispatches the built-in handler, which does not carry this flag.
LITELLM_PASS_THROUGH_ENDPOINT_MARKER = "__litellm_pass_through_endpoint__"
class EndpointType(str, Enum):
VERTEX_AI = "vertex-ai"

View file

@ -0,0 +1,63 @@
from typing import Any, Optional
from pydantic import BaseModel, Field
from .base import GuardrailConfigModel
class SingulrGuardrailRequest(BaseModel):
model: Optional[str] = None
messages: Optional[list[dict[str, Any]]] = None
tools: Optional[list[dict[str, Any]]] = None
model_response: Optional[dict[str, Any]] = None
litellm_metadata: Optional[dict[str, Any]] = None
class SingulrGuardrailPayload(BaseModel):
litellm_call_id: Optional[str] = None
request_data: Optional[SingulrGuardrailRequest] = None
input_type: str
is_playground_request: Optional[bool] = None
playground_text: Optional[str] = None
class SingulrGuardrailResponse(BaseModel):
"""Response returned by the Singulr guardrail API."""
should_block: bool = False
blocking_due_to: Optional[str] = None
class SingulrGuardrailConfigModel(GuardrailConfigModel):
singulr_api_key: Optional[str] = Field(
default=None,
description="The Singulr API key. Generate API key from Singulr Platform.",
)
singulr_api_base: Optional[str] = Field(
default=None,
description="The Singulr API base URL. Get base URL from Singulr Platform.",
)
singulr_application_id: Optional[str] = Field(
default=None,
description="The Singulr application ID. Get application ID from Singulr Platform.",
)
singulr_guardrail_id: Optional[str] = Field(
default=None,
description="The Singulr Guardrail ID. Get guardrail ID from Singulr Platform.",
)
block_on_error: Optional[bool] = Field(
default=None,
description=(
"Whether to block requests when the Singulr Guardrails API is unavailable "
"or returns an error. If enabled, requests fail closed. "
"If disabled, requests continue without guardrail enforcement (fail open)."
),
)
@staticmethod
def ui_friendly_name() -> str:
return "Singulr"

View file

@ -0,0 +1,169 @@
from __future__ import annotations
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk
from litellm.types.utils import ChatCompletionMessageToolCall
from .base import GuardrailConfigModel
StraikerWebhookEventType = Literal["pre_call", "post_call"]
StraikerWebhookStreamPhase = Literal["none", "assembled"]
StraikerWebhookAction = Literal["NONE", "BLOCKED", "GUARDRAIL_INTERVENED"]
STRAIKER_WEBHOOK_SCHEMA_VERSION = "1"
class StraikerWebhookStream(BaseModel):
phase: StraikerWebhookStreamPhase = "none"
index: int | None = None
class StraikerWebhookEvent(BaseModel):
type: StraikerWebhookEventType
id: str
stream: StraikerWebhookStream = Field(default_factory=StraikerWebhookStream)
class StraikerWebhookContent(BaseModel):
model_config = ConfigDict(extra="ignore")
texts: list[str] = Field(default_factory=list)
images: list[str] = Field(default_factory=list)
structured_messages: list[AllMessageValues] | None = None
tools: list[dict[str, object]] | None = None
tool_calls: list[ChatCompletionToolCallChunk] | list[ChatCompletionMessageToolCall] | None = None
finish_reason: str | None = None
class StraikerWebhookUsage(BaseModel):
input_tokens: int | None = None
output_tokens: int | None = None
class StraikerWebhookContext(BaseModel):
call_surface: str
model: str | None = None
model_provider: str | None = None
destination: str | None = None
session_id: str | None = None
litellm_call_id: str | None = None
litellm_trace_id: str | None = None
litellm_version: str | None = None
class StraikerWebhookIdentity(BaseModel):
litellm_key: str | None = None
litellm_team: str | None = None
litellm_user_id: str | None = None
litellm_user_email: str | None = None
litellm_org_id: str | None = None
end_user_id: str | None = None
class StraikerWebhookApplication(BaseModel):
source: str
name: str | None = None
class StraikerWebhookRequest(BaseModel):
schema_version: str = STRAIKER_WEBHOOK_SCHEMA_VERSION
event: StraikerWebhookEvent
request: StraikerWebhookContent
response: StraikerWebhookContent | None = None
context: StraikerWebhookContext
identity: StraikerWebhookIdentity
application: StraikerWebhookApplication
usage: StraikerWebhookUsage | None = None
metadata: dict[str, object] | None = None
class StraikerWebhookResponse(BaseModel):
model_config = ConfigDict(extra="allow")
action: StraikerWebhookAction = "NONE"
blocked_reason: str | None = None
texts: list[str] | None = None
schema_version: str | None = None
turn_id: str | None = Field(default=None, alias="turnId")
class StraikerGuardrailConfigModelOptionalParams(BaseModel):
timeout: float | None = Field(
default=5.0,
gt=0.0,
description="Per-attempt HTTP timeout in seconds.",
)
max_retries: int | None = Field(
default=2,
ge=0,
description="Retries on transient HTTP (408/429/5xx) and network errors.",
)
initial_backoff: float | None = Field(
default=0.1,
ge=0.0,
description="Initial retry backoff in seconds.",
)
max_backoff: float | None = Field(
default=2.0,
ge=0.0,
description="Maximum retry backoff in seconds.",
)
unreachable_fallback: Literal["fail_open", "fail_closed"] | None = Field(
default="fail_closed",
description="Behavior when Straiker is unreachable after retries.",
)
fail_on_error: bool | None = Field(
default=True,
description=(
"Behavior on any guardrail error, not just unreachability. True (default) blocks "
"the request on error; False logs and allows the request to proceed."
),
)
max_payload_bytes: int | None = Field(
default=524288,
gt=0,
description="Maximum serialized webhook payload size sent to Straiker.",
)
custom_headers: dict[str, str] | None = Field(
default=None,
description="Additional HTTP headers sent to Straiker, excluding Authorization and the webhook-format header.",
)
metadata: dict[str, str] | None = Field(
default=None,
description=(
"Default metadata key/values added to the webhook metadata bag on every request. "
"On key conflict with request-derived metadata, these configured values win."
),
)
verbose: bool | None = Field(
default=False,
description="Log webhook request/response payloads and record action/turn_id in response hidden params.",
)
class StraikerGuardrailConfigModel(GuardrailConfigModel[StraikerGuardrailConfigModelOptionalParams]):
api_key: str = Field(
min_length=1,
description="Straiker DefendAI environment API key (Bearer token). Env: STRAIKER_API_KEY.",
json_schema_extra={"secret": True},
)
api_base: str | None = Field(
default="https://api.prod.straiker.ai",
description="Straiker API base URL. Use the regional variant for non-US tenants.",
)
default_app: str | None = Field(
default="LiteLLM Gateway",
description=(
"Default application registered in the Straiker Defend Console. "
"Overridden per-request by metadata.agent_id when present."
),
)
@staticmethod
def ui_friendly_name() -> str:
return "Straiker"

View file

@ -5,7 +5,18 @@ litellm.Router Types - includes RouterConfig, UpdateRouterConfig, ModelInfo etc
import datetime
import enum
from dataclasses import dataclass
from typing import Any, Dict, List, Literal, Optional, Tuple, Union, get_type_hints
from typing import (
Any,
Dict,
Generic,
List,
Literal,
Optional,
Tuple,
TypeVar,
Union,
get_type_hints,
)
import httpx
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
@ -830,6 +841,31 @@ class PreRoutingHookResponse(BaseModel):
messages: Optional[List[Dict[str, Any]]]
_PreRoutingStrategyT_co = TypeVar("_PreRoutingStrategyT_co", covariant=True)
@dataclass(frozen=True, slots=True)
class TaggedPreRoutingStrategy(Generic[_PreRoutingStrategyT_co]):
"""A pre-routing strategy paired with the deployment `tags` it was registered under."""
tags: tuple[str, ...]
strategy: _PreRoutingStrategyT_co
@runtime_checkable
class PreRoutingStrategy(Protocol):
"""Structural interface shared by the auto / complexity / adaptive / quality routers."""
async def async_pre_routing_hook(
self,
model: str,
request_kwargs: dict[str, Any],
messages: list[dict[str, Any]] | None = None,
input: "str | list[Any] | None" = None,
specific_deployment: bool | None = False,
) -> "PreRoutingHookResponse | None": ...
class RoutingContext(BaseModel):
"""
Passed through a Router's `plugins` pipeline before the routing decision is made.

View file

@ -266,6 +266,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
"audio_transcription",
"responses",
"ocr",
"realtime",
]
]
tpm: Optional[int]
@ -402,6 +403,11 @@ class CallTypes(str, Enum):
vector_store_search = "vector_store_search"
avector_store_search = "avector_store_search"
ingest = "ingest"
aingest = "aingest"
query = "query"
aquery = "aquery"
#########################################################
# Container Call Types
#########################################################

View file

@ -3196,6 +3196,12 @@ def get_optional_params_embeddings(
non_default_params=non_default_params, optional_params={}, kwargs=kwargs
)
elif custom_llm_provider == "vertex_ai" or custom_llm_provider == "gemini":
# OpenAI SDKs (and litellm's own client) send encoding_format="float"
# by default; float lists are exactly what the vertex API returns, so
# the param is a no-op — don't reject the provider default. Other
# values (e.g. "base64") stay on the unsupported-param path below.
if non_default_params.get("encoding_format") == "float":
non_default_params.pop("encoding_format")
supported_params = get_supported_openai_params(
model=model,
custom_llm_provider="vertex_ai",

View file

@ -3451,7 +3451,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2.2e-05,
"output_cost_per_token": 2.64e-06,
"supports_audio_input": true,
@ -3470,7 +3470,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.00022,
"output_cost_per_token": 2.2e-05,
"supports_audio_input": true,
@ -3489,7 +3489,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2.2e-05,
"supported_modalities": [
@ -4687,7 +4687,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -4707,7 +4707,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -4739,7 +4739,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -4771,7 +4771,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -4832,7 +4832,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.0002,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -4850,7 +4850,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supported_modalities": [
@ -7922,7 +7922,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2.2e-05,
"output_cost_per_token": 2.64e-06,
"supports_audio_input": true,
@ -7941,7 +7941,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.00022,
"output_cost_per_token": 2.2e-05,
"supports_audio_input": true,
@ -7960,7 +7960,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2.2e-05,
"supported_modalities": [
@ -16272,7 +16272,7 @@
"supports_vision": false
},
"fireworks_ai/accounts/fireworks/models/glm-5p2": {
"cache_read_input_token_cost": 2.6e-07,
"cache_read_input_token_cost": 1.4e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
@ -16686,7 +16686,7 @@
"supports_vision": false
},
"fireworks_ai/glm-5p2": {
"cache_read_input_token_cost": 2.6e-07,
"cache_read_input_token_cost": 1.4e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
@ -22169,7 +22169,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -22188,7 +22188,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -22282,7 +22282,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -22300,7 +22300,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -22318,7 +22318,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -24513,7 +24513,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24545,7 +24545,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24577,7 +24577,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24610,7 +24610,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 2.4e-05,
"regional_processing_uplift_multiplier_eu": 1.1,
@ -24645,7 +24645,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"regional_processing_uplift_multiplier_eu": 1.1,
@ -24678,7 +24678,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -24710,7 +24710,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -43694,7 +43694,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -43727,7 +43727,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [

View file

@ -61,7 +61,7 @@ proxy = [
"boto3>=1.43.1,<2.0",
"azure-identity>=1.25.2,<2.0",
"azure-storage-blob>=12.28.0,<13.0",
"mcp>=1.26.0,<2.0",
"mcp>=1.28.1,<2.0",
"litellm-proxy-extras==0.4.78",
"litellm-enterprise==0.1.51",
"RestrictedPython>=8.1,<9.0",
@ -91,7 +91,7 @@ extra_proxy = [
"google-cloud-iam>=2.19.1,<3.0",
# Not in PyPI proxy extra.
"resend>=2.23.0,<3.0",
"redisvl>=0.4.1,<1.0; python_version < '3.14'",
"redisvl>=0.4.1,<1.0",
"a2a-sdk>=1.1.0,<2.0",
]
utils = [
@ -136,7 +136,7 @@ proxy-runtime = [
"mangum>=0.17.0,<1.0",
"azure-ai-contentsafety>=1.0.0,<2.0",
"azure-storage-file-datalake>=12.20.0,<13.0",
"pypdf>=6.12.0,<7.0; python_version < '3.14'",
"pypdf>=6.12.0,<7.0",
"llm-sandbox>=0.3.39,<1.0",
"detect-secrets>=1.5.0,<2.0",
]
@ -181,7 +181,7 @@ dev = [
"pytest-rerunfailures==15.1",
"pytest-cov==5.0.0",
"parameterized==0.9.0",
"openapi-core==0.22.0; python_version < '3.14'",
"openapi-core==0.22.0",
"pytest-timeout==2.4.0",
"vcrpy==8.2.1",
"pytest-recording==0.13.4",

28
router_plugins.json Normal file
View file

@ -0,0 +1,28 @@
[
{
"name": "TEMPLATE: copy this block for a new plugin, then delete this entry",
"description": "One line on what the plugin does and the routing signal it publishes.",
"author": "Plugin author's name.",
"repo": "https://github.com/<owner>/<repo> (public source repository).",
"commit": "Full 40-char git SHA to pin when the plugin is not yet on PyPI; omit once 'pypi' is set.",
"version": "Plugin release version, e.g. 1.0.0.",
"pypi": "PyPI spec pinned to a version, e.g. my-plugin==1.0.0, or null if unpublished.",
"litellm_version": "Minimum compatible litellm version, e.g. >=1.94.0.",
"entrypoint": "Dotted import path to the plugin instance, e.g. my_plugin.plugin.instance.",
"license": "SPDX license id, e.g. MIT.",
"tags": ["searchable", "keywords"]
},
{
"name": "language-detector",
"description": "Detects the user's language and publishes a routing signal.",
"author": "Jean Nuñez",
"repo": "https://github.com/jeann2013/language-detector",
"commit": "9e712819269173fc25a16f59ca3e9890f7864ac1",
"version": "1.0.0",
"pypi": null,
"litellm_version": ">=1.94.0",
"entrypoint": "litellm_plugin_language_detector.plugin.language_detector_plugin",
"license": "MIT",
"tags": ["language", "classification", "routing"]
}
]

View file

@ -13,6 +13,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
- `realtime/` - realtime websocket sessions, including the pipecat audio path
- `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `ratelimit/` (rpm/tpm blocks, window reset, pacing headers on live traffic), `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window), and `spend_tracking/` (spend logging and cost attribution on `/spend/*`)
- `management/` - key/team/user/organization management routes: create/update/delete persistence via the info routes, team membership, and llm-only-key route denials; also the dashboard UI behavior on top of them, driven through the proxy-served UI at /ui with playwright (optional dep behind importorskip)
- `mcp/` - the MCP server surface over api_key auth: an admin registers an upstream MCP server through the management API and grants keys access via `object_permission.mcp_servers`, then the suite asserts tool listing and calling honor that permission (a key without the grant sees none of the server's tools and is refused a `tools/call` with a 403)
- `logging/` - logging-integration delivery (datadog and friends)
- `security/` - secret handling and log-leak protection
- `router/` - routing and reliability behavior (fallbacks, cooldowns)
@ -51,7 +52,7 @@ The shape is layered so tests stay declarative
Each suite provides its own `client` fixture (see `llm_translation/passthrough_client.py`), a frozen dataclass that holds the shared `Gateway` and adds suite-specific routes. Cleanup runs through that same `Gateway`, so whatever keys or customers your test creates get torn down by the `resources` fixture
Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The skip-vs-fail split is deliberate: a test marked `e2e` skips when no proxy answers its liveness probe, but once a request reaches the proxy any wrong behavior is a hard failure, never a skip
Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The harness hard-fails and never skips: a test marked `e2e` fails when no proxy answers its liveness probe, and once a request reaches the proxy any wrong behavior is likewise a hard failure, so a missing proxy turns the run red instead of being mistaken for a pass
Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure coverage of the harness itself carries no marker and runs regardless. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache
@ -131,7 +132,7 @@ Quota Management - behavior features (entity- or config-driven caps and their ac
quota_management.<behavior>.<variant>.<assertion>
behavior : ratelimit | budget | spend_tracking
variant : <ratelimit> rpm | tpm | priority_generous | priority_strict
<budget> key | internal_user | end_user | organization | team_member | tag
<budget> key | internal_user | end_user | organization | team | team_member | tag
| model_max | soft | key_multi_window | team_multi_window
| fallback | spend_counter
<spend_tracking> chat_completions | stream | embeddings | cache_hit | key_rollup
@ -139,7 +140,8 @@ quota_management.<behavior>.<variant>.<assertion>
| spend_calculate | pagination
assertion : blocks_over_limit | resets_after_window | headers_report_remaining | picks_under_tpm
| blocks_then_resets | resets_windows_independently | alerts_without_blocking
| isolates_per_model | routes_to_fallback | reseed_matches_db | logs_cost | zero_cost
| isolates_per_model | isolates_per_member | enforced_across_keys | routes_to_fallback
| reseed_matches_db | logs_cost | zero_cost
| matches_sum_of_logs | loses_no_spend | attributes_spend | writes_own_rows
| writes_failure_row | returns_cost | keeps_total
e.g. quota_management.ratelimit.rpm.blocks_over_limit exercised_on=[chat_completions, messages]
@ -173,7 +175,7 @@ other.<area>.<case>.<assertion>
```
## Hard Rules
- no monkeypatching, mock tests or unit tests of any kind. if a contributor asks you to write an end to end test, do NOT stage a unit test with it. if you find a product gap, call it out in the PR description
- no monkeypatching or mock tests, and never substitute a unit test for e2e feature coverage: a product feature is proven end to end against a live proxy, not with a unit test. if a contributor asks you to write an end to end test, do NOT stage a unit test of the feature with it; if you find a product gap, call it out in the PR description. tests that cover the harness itself are the exception and are allowed (for example `coverage_registry/test_collector.py`, which unit-tests the coverage collector): they carry no `e2e` marker, exercise harness plumbing rather than a product feature, and run whether or not a proxy is up
- use model management endpoints to create new models for a test. this could be in a conftest / inline for each test. ask the user what they want.

View file

@ -54,7 +54,7 @@ The suites run against a live proxy, so bring one up first. `docker-compose.yml`
docker compose down -v
```
Tests marked `@pytest.mark.e2e` skip when no proxy answers `/health/liveliness`, so a run that reports everything skipped means the stack isn't up, not that anything passed
Tests marked `@pytest.mark.e2e` hard-fail when no proxy answers `/health/liveliness`, so a run that goes red with `No live proxy` at setup means the stack isn't up; they never skip for a missing proxy, so an absent stack can't be mistaken for a pass
## What a complete test looks like
@ -132,7 +132,7 @@ The shape is layered so tests stay declarative
Each suite provides its own `client` fixture (see `llm_translation/passthrough_client.py`), a frozen dataclass that holds the shared `Gateway` and adds suite-specific routes. Cleanup runs through that same `Gateway`, so whatever keys or customers your test creates get torn down by the `resources` fixture
Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The skip-vs-fail split is deliberate: a test marked `e2e` skips when no proxy answers its liveness probe, but once a request reaches the proxy any wrong behavior is a hard failure, never a skip
Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The harness hard-fails and never skips: a test marked `e2e` fails when no proxy answers its liveness probe, and once a request reaches the proxy any wrong behavior is likewise a hard failure, so a missing proxy turns the run red instead of being mistaken for a pass
Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure coverage of the harness itself carries no marker and runs regardless. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache

View file

@ -1,4 +1,4 @@
"""Access-control suite client fixture; lifecycle/skip/marker live in the parent conftest."""
"""Access-control suite client fixture; lifecycle/liveness gate/marker live in the parent conftest."""
import pytest

View file

@ -1,6 +1,6 @@
"""Batches suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
live in the parent tests/e2e/conftest.py. BatchClient holds the shared Gateway, so
the `resources` fixture cleans up keys through it; tests register file deletes and
batch cancels via `resources.defer(...)`.

View file

@ -1,247 +0,0 @@
"""Bob the builder: on a red e2e run, ask Devin to fix the failing tests.
Wired as a ``pytest_sessionfinish`` step (see ``conftest.py``). When the run went
red and remediation is enabled, it hands the failing tests plus their captured
tracebacks to Devin *through the LiteLLM proxy's own MCP gateway* -- the same
gateway + master key the suite already uses -- so Devin files a Linear ticket per
failure and opens fix PRs. Nothing new ships in the runner pod: the proxy already
registers the ``devin`` MCP server and holds ``DEVIN_API_KEY``, injecting it
upstream, so this process only needs the proxy key it always has.
Opt-in via ``E2E_DEVIN_REMEDIATION=1`` so a normal local ``pytest tests/e2e`` run
never spawns a Devin session. ``DEVIN_DRY_RUN=1`` prints the prompt it would send
and makes no call. Everything is best-effort: any error here is logged and
swallowed so the run's exit status still reflects the tests, not remediation.
"""
from __future__ import annotations
import hashlib
import os
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Protocol, cast
import pytest
from pydantic import BaseModel, ConfigDict
from e2e_config import MASTER_KEY, PROXY_BASE_URL
from e2e_http import Success
from transport import HttpTransport
REMEDIATION_ENV = "E2E_DEVIN_REMEDIATION"
_LIST_PATH = "/mcp-rest/tools/list"
_CALL_PATH = "/mcp-rest/tools/call"
@dataclass(frozen=True, slots=True)
class Failure:
"""One failed test: its pytest node id and the captured failure text."""
nodeid: str
detail: str
@dataclass(frozen=True, slots=True)
class Config:
server: str
create_tool: str
linear_team: str
target_repo: str
target_ref: str
max_failures: int
max_detail_chars: int
tags: tuple[str, ...]
dry_run: bool
class _NoParams(BaseModel):
pass
class _McpToolInfo(BaseModel):
model_config = ConfigDict(extra="allow")
server_name: str | None = None
alias: str | None = None
class _McpTool(BaseModel):
model_config = ConfigDict(extra="allow")
name: str
mcp_info: _McpToolInfo | None = None
class _McpToolsList(BaseModel):
model_config = ConfigDict(extra="allow")
tools: tuple[_McpTool, ...] = ()
class _DevinSessionArgs(BaseModel):
prompt: str
title: str
tags: list[str]
class _ToolCallBody(BaseModel):
name: str
arguments: _DevinSessionArgs
class _ToolCallResult(BaseModel):
model_config = ConfigDict(extra="allow")
class _Report(Protocol):
@property
def nodeid(self) -> str: ...
@property
def longreprtext(self) -> str: ...
class _TerminalReporter(Protocol):
stats: Mapping[str, Sequence[_Report]]
def _env(name: str, default: str) -> str:
value = os.environ.get(name, "").strip()
return value or default
def load_config() -> Config:
raw_tags = _env("DEVIN_TAGS", "e2e,stage")
return Config(
server=_env("DEVIN_MCP_SERVER", "devin"),
create_tool=_env("DEVIN_SESSION_TOOL", "devin_session_create"),
linear_team=_env("DEVIN_LINEAR_TEAM", "LIT"),
target_repo=_env("DEVIN_TARGET_REPO", "BerriAI/litellm"),
target_ref=_env("DEVIN_TARGET_REF", "litellm_internal_staging"),
max_failures=int(_env("DEVIN_MAX_FAILURES", "50")),
max_detail_chars=int(_env("DEVIN_MAX_DETAIL_CHARS", "3000")),
tags=tuple(t.strip() for t in raw_tags.split(",") if t.strip()),
dry_run=_env("DEVIN_DRY_RUN", "0") == "1",
)
def collect_failures(session: pytest.Session, max_detail_chars: int) -> tuple[Failure, ...]:
"""Pull the failed and errored tests (with their tracebacks) off the run's
terminal reporter. Returns empty when nothing failed or the reporter is
absent (e.g. a skipped, proxy-less session)."""
plugin: object = session.config.pluginmanager.getplugin("terminalreporter")
if plugin is None:
return ()
reporter = cast(_TerminalReporter, plugin)
reports = (*reporter.stats.get("failed", ()), *reporter.stats.get("error", ()))
return tuple(
Failure(nodeid=r.nodeid, detail=r.longreprtext.strip()[-max_detail_chars:]) for r in reports
)
def dedup_tag(failures: tuple[Failure, ...]) -> str:
"""Stable short tag identifying this exact set of failing tests, so repeated
nightly runs on the same failures reference one body of work."""
joined = "\n".join(sorted(f.nodeid for f in failures))
return "e2e-fail-" + hashlib.sha256(joined.encode()).hexdigest()[:12]
def _revision() -> str:
for candidate in (Path(__file__).parent / ".litellm-revision", Path("/app/e2e/.litellm-revision")):
try:
return candidate.read_text(encoding="utf-8").strip()
except OSError:
continue
return _env("E2E_REVISION", "unknown")
def build_prompt(cfg: Config, failures: tuple[Failure, ...], tag: str) -> str:
shown = failures[: cfg.max_failures]
header = (
f"The LiteLLM end-to-end suite failed on the "
f"{_env('E2E_ENVIRONMENT', 'stage')} proxy. Source repo {cfg.target_repo} "
f"at revision {_revision()} (branch {cfg.target_ref}). {len(failures)} "
f"test(s) failed"
+ (f"; the first {len(shown)} are shown" if len(shown) < len(failures) else "")
+ ".\n\n"
)
task = (
"For each failing test below:\n"
f"1. Open a Linear ticket under the {cfg.linear_team} team describing the "
"failure (test id, the assertion/error, likely cause), unless an open "
"ticket for that same test already exists -- do not create duplicates.\n"
f"2. Fix it in {cfg.target_repo}, branching off {cfg.target_ref} and "
"following the repo's CONTRIBUTING and CLAUDE.md conventions (meaningful "
"regression coverage, conventional commits, run the suite locally), then "
"open a PR that references the Linear ticket.\n"
"3. Prefer one focused PR per failing test; if several share a root cause, "
"group them and say so.\n"
f"Before starting, search existing sessions/PRs tagged '{tag}' or "
"referencing these test ids and continue that work instead of restarting.\n\n"
"Failing tests and their captured output:\n"
)
blocks = [f"### {i}. {f.nodeid}\n```\n{f.detail}\n```\n" for i, f in enumerate(shown, start=1)]
return header + task + "\n".join(blocks)
def _resolve_tool_name(transport: HttpTransport, cfg: Config) -> str | None:
"""Find Devin's create-session tool on the gateway. The proxy prefixes tools
with the server alias, so match by suffix and (when present) the owning
server."""
result = transport.get(
_LIST_PATH, headers=transport.master, params=_NoParams(), response_type=_McpToolsList
)
if not isinstance(result, Success):
print(f"bob_the_builder: could not list gateway MCP tools: {result}")
return None
for tool in result.data.tools:
owner = tool.mcp_info.server_name or tool.mcp_info.alias if tool.mcp_info else None
if (owner is None or owner == cfg.server) and (
tool.name == cfg.create_tool or tool.name.endswith(cfg.create_tool)
):
return tool.name
print(
f"bob_the_builder: no '{cfg.create_tool}' tool for server '{cfg.server}' on the gateway; "
f"saw {[t.name for t in result.data.tools]}"
)
return None
def remediate(session: pytest.Session) -> None:
"""Entry point called from ``pytest_sessionfinish``. No-op unless remediation
is enabled and the run actually had failures."""
if os.environ.get(REMEDIATION_ENV) != "1":
return
cfg = load_config()
failures = collect_failures(session, cfg.max_detail_chars)
if not failures:
return
tag = dedup_tag(failures)
title = f"Fix {len(failures)} failing LiteLLM e2e test(s) [{tag}]"
prompt = build_prompt(cfg, failures, tag)
args = _DevinSessionArgs(prompt=prompt, title=title, tags=[*cfg.tags, tag])
if cfg.dry_run:
print("bob_the_builder: DRY RUN -- would create a Devin session:")
print(f" server : {cfg.server}\n tool : {cfg.create_tool}\n title : {title}")
print(f" tags : {args.tags}\n---- prompt ----\n{prompt}")
return
try:
transport = HttpTransport(base_url=PROXY_BASE_URL, master_key=MASTER_KEY)
tool_name = _resolve_tool_name(transport, cfg)
if tool_name is None:
return
result = transport.post(
_CALL_PATH,
headers=transport.master,
json=_ToolCallBody(name=tool_name, arguments=args),
response_type=_ToolCallResult,
)
if isinstance(result, Success):
print(f"bob_the_builder: created Devin session for {len(failures)} failure(s) [{tag}]")
print(result.data.model_dump_json())
else:
print(f"bob_the_builder: Devin session call failed: {result}")
except Exception as exc: # noqa: BLE001 - remediation must never fail the run
print(f"bob_the_builder: remediation error (ignored): {exc}")

View file

@ -16,12 +16,12 @@ cover "OpenAI plus the big three clouds":
carries only the open-weight
gpt-oss MaaS models
The openai and azure_openai columns run unconditionally, like every
other live column: the environments that run the suite carry
`OPENAI_API_KEY` and `AZURE_API_BASE` + `AZURE_API_KEY` pointing at a
resource with gpt-5.6 deployments. The bedrock_mantle column is
opt-in via `COMPAT_MANTLE_CELLS=1` because the AWS account is still
waiting on the Bedrock Mantle allowlist for the `openai.gpt-5.6-*`
The azure_openai column runs unconditionally when Azure gpt-5.6
deployments exist. The openai column is opt-in via
`COMPAT_OPENAI_GPT_CELLS=1` because under the full stage suite those
cells routinely burn minutes on Claude CLI timeouts. The bedrock_mantle
column is opt-in via `COMPAT_MANTLE_CELLS=1` because the AWS account is
still waiting on the Bedrock Mantle allowlist for the `openai.gpt-5.6-*`
models; until the flag is set each Mantle cell skips and its matrix
cell publishes as `not_tested` instead of a credential-shaped red.
The `vertex_ai_gpt` column needs no flag either way: its cells report
@ -35,6 +35,7 @@ import os
import pytest
MANTLE_CELLS_ENV = "COMPAT_MANTLE_CELLS"
OPENAI_GPT_CELLS_ENV = "COMPAT_OPENAI_GPT_CELLS"
VERTEX_AI_GPT_NOT_APPLICABLE_REASON = (
"GCP Vertex AI does not offer OpenAI's closed-weight GPT-5.6 family "
@ -59,3 +60,19 @@ def skip_unless_mantle_cells_enabled() -> None:
f"Bedrock Mantle GPT-5.6 cells are opt-in; set {MANTLE_CELLS_ENV}=1 "
"once the AWS account is allowlisted for the openai.gpt-5.6-* models"
)
def skip_unless_openai_gpt_cells_enabled() -> None:
"""Skip OpenAI GPT-5.6 columns unless `COMPAT_OPENAI_GPT_CELLS` opts them in.
Under the full stage suite these cells routinely hit 120s Claude CLI
timeouts and rate-limit-shaped retries across Sol/Terra/Luna, burning
~8+ minutes per cell without a stable green. Opt in when exercising
the OpenAI GPT translation path in isolation.
"""
if os.environ.get(OPENAI_GPT_CELLS_ENV, "").strip().lower() in {"1", "true", "yes"}:
return
pytest.skip(
f"OpenAI GPT-5.6 cells are opt-in; set {OPENAI_GPT_CELLS_ENV}=1 "
"to run them (stage suite timeouts under concurrent load)"
)

View file

@ -23,6 +23,7 @@ green if all three pass.
from __future__ import annotations
from claude_code._basic_messaging import run_basic_messaging_cell
from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled
OPENAI_MODELS = [
"gpt-5-6-sol-openai",
@ -34,6 +35,7 @@ OPENAI_MODELS = [
def test_basic_messaging_non_streaming_openai(compat_result):
"""Drive the `claude` CLI against the LiteLLM proxy and assert a
non-empty reply from each GPT-5.6 tier."""
skip_unless_openai_gpt_cells_enabled()
run_basic_messaging_cell(
compat_result=compat_result,
models=OPENAI_MODELS,

View file

@ -25,6 +25,7 @@ green if all three pass.
from __future__ import annotations
from claude_code._basic_messaging import run_basic_messaging_cell
from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled
OPENAI_MODELS = [
"gpt-5-6-sol-openai",
@ -36,6 +37,7 @@ OPENAI_MODELS = [
def test_basic_messaging_streaming_openai(compat_result):
"""Drive the `claude` CLI against the LiteLLM proxy and assert a
non-empty streamed reply from each GPT-5.6 tier."""
skip_unless_openai_gpt_cells_enabled()
run_basic_messaging_cell(
compat_result=compat_result,
models=OPENAI_MODELS,

View file

@ -29,6 +29,7 @@ from typing import Any, Mapping, Sequence
import pytest
from claude_code._env import require_proxy
from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled
from claude_code.cli_driver import (
ClaudeCLIError,
failure_diagnostic,
@ -69,6 +70,7 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool:
def test_tool_use_openai(compat_result):
"""Drive the `claude` CLI against the LiteLLM proxy and assert a
tool call was emitted on the wire by each GPT-5.6 tier."""
skip_unless_openai_gpt_cells_enabled()
proxy = require_proxy(compat_result)
outcomes = run_claude_models_parallel(

View file

@ -30,6 +30,7 @@ from typing import Any, Mapping, Sequence
import pytest
from claude_code._env import require_proxy
from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled
from claude_code.cli_driver import (
ClaudeCLIError,
failure_diagnostic,
@ -87,6 +88,7 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int:
def test_tool_use_streaming_openai(compat_result):
skip_unless_openai_gpt_cells_enabled()
proxy = require_proxy(compat_result)
outcomes = run_claude_models_parallel(

View file

@ -15,14 +15,14 @@ shared fixtures build on it.
import functools
import sys
from collections.abc import Generator, Iterator
from collections.abc import Iterator
from pathlib import Path
import pytest
import requests
from e2e_config import CONTROL_PLANE_BASE_URL, PROXY_BASE_URL
from e2e_result_reporter import covers_from_item, format_e2e_result_line, result_from_pytest
from junit_properties import attach_result_properties
from lifecycle import GatewayProvider, ResourceManager
@ -40,6 +40,17 @@ def pytest_configure(config: pytest.Config) -> None:
)
def pytest_collection_modifyitems(items: list[pytest.Item]) -> None:
"""Attach the two custom signals (suite package and covered cell ids) to every
test's user_properties so the standard JUnit report (`--junitxml`) records them
as `<property>` entries, on every outcome including skips and setup errors.
Downstream (Loki/Grafana) reads outcome and duration from the standard report
and these properties for package rollups and coverage drill-down. See
junit_properties.py."""
for item in items:
attach_result_properties(item)
def _liveness_reason(label: str, base_url: str) -> str | None:
"""None if `base_url` answers its liveness probe, else a failure reason."""
try:
@ -86,30 +97,6 @@ def pytest_runtest_call(item: pytest.Item) -> None:
item.session.stash[_E2E_TEST_RAN] = True
@pytest.hookimpl(wrapper=True, tryfirst=True)
def pytest_runtest_makereport(
item: pytest.Item, call: pytest.CallInfo[object]
) -> Generator[None, pytest.TestReport, pytest.TestReport]:
"""Emit one structured E2E_RESULT line per finished test for Loki/Grafana.
Status-history panels should aggregate by package (and optional covers), not
scrape pytest progress basenames. See e2e_result_reporter.py.
"""
report = yield
result = result_from_pytest(
nodeid=str(report.nodeid),
when=str(report.when),
failed=bool(report.failed),
skipped=bool(report.skipped),
passed=bool(report.passed),
duration_seconds=float(report.duration),
covers=covers_from_item(item),
)
if result is not None:
print(format_e2e_result_line(result), flush=True)
return report
def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
"""Once the whole e2e session is done (all suites), truncate the spend logs so
the DB doesn't accumulate test rows. Sessions where no e2e test body ran leave
@ -132,13 +119,6 @@ def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
if spend_dir in sys.path:
sys.path.remove(spend_dir)
try:
from bob_the_builder import remediate
remediate(session)
except Exception as exc: # noqa: BLE001 - remediation is best-effort
print(f"devin remediation best-effort failed: {exc}")
@pytest.fixture
def resources(client: GatewayProvider) -> Iterator[ResourceManager]:

View file

@ -7,14 +7,20 @@
- {id: quota_management.ratelimit.priority_generous.picks_under_tpm, module: quota_management, tier: P1, behavior: ratelimit, variant: priority_generous, assertions: [picks_under_tpm], exercised_on: [chat_completions, messages], source: "dynamic_rate_limiter_v3.py:36-52", rationale: "Generous mode (<80% sat) allows priority borrowing"}
- {id: quota_management.ratelimit.priority_strict.picks_under_tpm, module: quota_management, tier: P1, behavior: ratelimit, variant: priority_strict, assertions: [picks_under_tpm], exercised_on: [chat_completions, messages], source: "dynamic_rate_limiter_v3.py:53-71", rationale: "Strict mode (>=80% sat) enforces priority fairness"}
- {id: quota_management.budget.key.blocks_over_limit, module: quota_management, tier: P0, behavior: budget, variant: key, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "A key's max_budget blocks further paid calls once spend crosses it"}
- {id: quota_management.budget.team.blocks_over_limit, module: quota_management, tier: P0, behavior: budget, variant: team, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "A team's max_budget blocks every key on the team once combined spend crosses it, including keys that spent nothing themselves"}
- {id: quota_management.budget.internal_user.blocks_over_limit, module: quota_management, tier: P1, behavior: budget, variant: internal_user, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "An internal user's max_budget governs personal keys"}
- {id: quota_management.budget.internal_user.enforced_across_keys, module: quota_management, tier: P1, behavior: budget, variant: internal_user, assertions: [enforced_across_keys], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "An internal user's max_budget governs every personal key it owns; a second untouched key is blocked once the shared user budget is exhausted"}
- {id: quota_management.budget.end_user.blocks_over_limit, module: quota_management, tier: P1, behavior: budget, variant: end_user, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "A customer (end-user) max_budget blocks calls attributed via user="}
- {id: quota_management.budget.organization.blocks_over_limit, module: quota_management, tier: P1, behavior: budget, variant: organization, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "An organization's max_budget blocks keys under its teams"}
- {id: quota_management.budget.team_member.blocks_over_limit, module: quota_management, tier: P1, behavior: budget, variant: team_member, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "A member's per-team budget blocks independently of the team budget"}
- {id: quota_management.budget.team_member.isolates_per_member, module: quota_management, tier: P1, behavior: budget, variant: team_member, assertions: [isolates_per_member], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "One team member's exhausted per-team budget does not block a different member on the same team"}
- {id: quota_management.budget.tag.blocks_over_limit, module: quota_management, tier: P1, behavior: budget, variant: tag, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "router_strategy/budget_limiter.py", rationale: "Proxy-level tag budgets block tagged requests at the cap"}
- {id: quota_management.budget.model_max.isolates_per_model, module: quota_management, tier: P1, behavior: budget, variant: model_max, assertions: [isolates_per_model], exercised_on: [chat_completions], source: "proxy/hooks/model_max_budget_limiter.py", rationale: "model_max_budget caps one model without touching a sibling's budget"}
- {id: quota_management.budget.soft.alerts_without_blocking, module: quota_management, tier: P1, behavior: budget, variant: soft, assertions: [alerts_without_blocking], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "soft_budget alerts but never blocks traffic"}
- {id: quota_management.budget.key.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: key, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_duration zeroes key spend after the window; a blocked key serves again"}
- {id: quota_management.budget.team.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: team, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_duration zeroes a team's spend after the window; every key on the team serves again"}
- {id: quota_management.budget.organization.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: organization, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "An org budget resets after its window; keys under the org serve again"}
- {id: quota_management.budget.internal_user.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: internal_user, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "An internal user's budget resets after its window; their personal and team-member keys serve again"}
- {id: quota_management.budget.team_member.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: team_member, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "Member per-team budget reset keeps advancing window after window"}
- {id: quota_management.budget.key_multi_window.blocks_then_resets, module: quota_management, tier: P1, behavior: budget, variant: key_multi_window, assertions: [blocks_then_resets], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_limits enforce within a short window and serve again in the next"}
- {id: quota_management.budget.key_multi_window.resets_windows_independently, module: quota_management, tier: P2, behavior: budget, variant: key_multi_window, assertions: [resets_windows_independently], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "Each window of a multi-window budget resets on its own schedule"}

View file

@ -1,5 +1,7 @@
# local setup to run e2e tests
configs:
mcp_upstream_server:
file: ../mcp_tests/mcp_e2e_upstream_server.py
litellm_config:
content: |
general_settings:
@ -131,7 +133,27 @@ services:
target: /app/config.yaml
command: ["--config", "/app/config.yaml", "--port", "4000"]
# throwaway db
# deterministic self-hosted upstream MCP server (FastMCP add/multiply over
# streamable-http), reachable by the litellm container at mcp-upstream:8090/mcp.
# Not a depends_on of litellm on purpose: only the mcp suite needs it, and it
# boots long before the proxy is live, so it must not gate the other suites'
# stack. The suite registers it through /v1/mcp/server at test time.
mcp-upstream:
image: ghcr.io/berriai/litellm:main-latest
entrypoint: ["python3", "/app/mcp_upstream_server.py"]
environment:
MCP_HOST: 0.0.0.0
MCP_PORT: "8090"
configs:
- source: mcp_upstream_server
target: /app/mcp_upstream_server.py
healthcheck:
test: ["CMD", "python3", "-c", "import socket; socket.create_connection(('127.0.0.1', 8090), 2).close()"]
interval: 3s
timeout: 3s
retries: 40
# throwaway db
db:
image: postgres:16
environment:

View file

@ -1,144 +0,0 @@
"""Structured e2e result lines for Loki / Grafana status history.
Pytest progress lines are a bad dashboard source: they only expose file basenames,
break under quiet modes, and force status-history rows to explode with suite growth.
Each finished test emits one logfmt line:
E2E_RESULT package=logging file=test_langfuse_e2e.py outcome=failed
duration_ms=1234 node_id=logging/test_langfuse_e2e.py::TestX::test_y
covers=logging.langfuse.team.success
Grafana package status-history queries max(fail) by package over E2E_RESULT lines.
Drill-down uses node_id / covers in Explore, not status-history cardinality.
"""
from __future__ import annotations
from collections.abc import Iterable, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Literal, Protocol, runtime_checkable
Outcome = Literal["passed", "failed", "error", "skipped"]
@dataclass(frozen=True, slots=True)
class E2EResult:
package: str
file: str
outcome: Outcome
duration_ms: int
node_id: str
covers: tuple[str, ...]
@runtime_checkable
class _MarkerArgs(Protocol):
args: Sequence[object]
@runtime_checkable
class _ItemWithCovers(Protocol):
def iter_markers(self, name: str) -> Iterable[object]: ...
def package_from_nodeid(nodeid: str) -> str:
"""Top-level suite package under tests/e2e/, or 'root' for top-level files.
Pytest nodeids are relative to the invocation cwd. Repo-root runs look like
`tests/e2e/logging/...`; suite-cwd runs look like `logging/...`. Strip the
`tests/e2e` prefix so package is the suite dir either way.
"""
path_part = nodeid.split("::", 1)[0].replace("\\", "/")
parts = tuple(p for p in path_part.split("/") if p and p != ".")
if len(parts) >= 3 and parts[0] == "tests" and parts[1] == "e2e":
parts = parts[2:]
if len(parts) <= 1:
return "root"
return parts[0]
def file_from_nodeid(nodeid: str) -> str:
path_part = nodeid.split("::", 1)[0].replace("\\", "/")
return Path(path_part).name
def covers_from_item(item: object) -> tuple[str, ...]:
"""Read @pytest.mark.covers cell ids from a pytest Item."""
if not isinstance(item, _ItemWithCovers):
return ()
return tuple(
dict.fromkeys(
arg
for marker in item.iter_markers(name="covers")
if isinstance(marker, _MarkerArgs)
for arg in marker.args
if isinstance(arg, str) and arg
)
)
def outcome_from_report(when: str, failed: bool, skipped: bool, passed: bool) -> Outcome | None:
"""Map pytest TestReport fields to a terminal outcome. None if not final."""
if when == "setup" and skipped:
return "skipped"
if when == "setup" and failed:
return "error"
if when != "call":
return None
if skipped:
return "skipped"
if failed:
return "failed"
if passed:
return "passed"
return "failed"
def _logfmt_escape(value: str) -> str:
if value == "":
return '""'
needs_quote = any(ch.isspace() or ch in "\"=\\" for ch in value)
if not needs_quote:
return value
escaped = value.replace("\\", "\\\\").replace('"', '\\"')
return f'"{escaped}"'
def format_e2e_result_line(result: E2EResult) -> str:
covers = ",".join(result.covers)
fields = (
("package", result.package),
("file", result.file),
("outcome", result.outcome),
("duration_ms", str(result.duration_ms)),
("node_id", result.node_id),
("covers", covers),
)
body = " ".join(f"{key}={_logfmt_escape(value)}" for key, value in fields)
return f"E2E_RESULT {body}"
def result_from_pytest(
*,
nodeid: str,
when: str,
failed: bool,
skipped: bool,
passed: bool,
duration_seconds: float,
covers: tuple[str, ...] = (),
) -> E2EResult | None:
outcome = outcome_from_report(when=when, failed=failed, skipped=skipped, passed=passed)
if outcome is None:
return None
duration_ms = max(0, int(round(duration_seconds * 1000)))
return E2EResult(
package=package_from_nodeid(nodeid),
file=file_from_nodeid(nodeid),
outcome=outcome,
duration_ms=duration_ms,
node_id=nodeid,
covers=covers,
)

View file

@ -1,66 +0,0 @@
# Grafana: package status history for e2e
Dashboard: [LiteLLM E2E](https://berriai.grafana.net/d/mup2cfn/litellm-e2e) (`mup2cfn`).
The old **test suite status history** panel scraped pytest progress lines and
grouped by **file basename** (`test_foo.py`). That does not scale: multi-class
files collapse to one bit, and full `node_id` cardinality melts status-history.
## Emitter
After each test finishes, `tests/e2e/conftest.py` prints one logfmt line:
```
E2E_RESULT package=logging file=test_langfuse_e2e.py outcome=failed duration_ms=1500 node_id="logging/..." covers=cell.id
```
## Panel: package status history (replace panel 11)
**Type:** Status history
**Interval:** 15m (or 1h for multi-day ranges)
**Description:** Per top-level package under `tests/e2e/`: red if any test failed or errored in the bucket.
```logql
max by (package) (
max_over_time(
{service_name="litellm-e2e"}
|= "E2E_RESULT"
| logfmt
| outcome != ""
| label_format result=`{{ if or (eq .outcome "failed") (eq .outcome "error") }}1{{ else }}0{{ end }}`
| unwrap result
[$__interval]
)
)
```
Value mappings: `0` → Pass (green), `1` → Fail (red).
If `service_name` is missing on older scrapes, use:
```logql
{cluster="berrie-litellm-stage", pod=~"litellm-e2e-.+"}
```
instead of `{service_name="litellm-e2e"}`.
## Panel: failed tests (logs drill-down)
```logql
{service_name="litellm-e2e"} |= "E2E_RESULT" | logfmt | outcome=~"failed|error"
```
Show fields: `package`, `file`, `node_id`, `covers`, `duration_ms`.
## Panel (optional): filter by package variable
Dashboard variable `package` (custom or from label_values on E2E_RESULT):
```logql
{service_name="litellm-e2e"} |= "E2E_RESULT" | logfmt | package=`$package` | outcome=~"failed|error"
```
## Do not
- Put full `node_id` as the status-history series key (cardinality).
- Rely on `::S+ PASSED` progress regex as the primary signal once E2E_RESULT is live.

View file

@ -0,0 +1,59 @@
"""Custom per-test signals for the standard JUnit reporter.
The e2e suite ships results to Loki/Grafana from a standard pytest JUnit report
(`--junitxml=e2e-report.xml`), not a bespoke log line. JUnit already records
outcome, duration, and node id for every `<testcase>`; the only signals it cannot
derive on its own are the normalized suite package and the coverage-registry cell
ids a test covers. Those ride along as JUnit `<property>` entries via each item's
`user_properties`, attached in `conftest.py::pytest_collection_modifyitems`.
"""
from __future__ import annotations
from collections.abc import Iterable
import pytest
def package_from_nodeid(nodeid: str) -> str:
"""Top-level suite package under tests/e2e/, or 'root' for top-level files.
Pytest nodeids are relative to the invocation cwd. Repo-root runs look like
`tests/e2e/logging/...`; suite-cwd runs look like `logging/...`. Strip the
`tests/e2e` prefix so package is the suite dir either way.
"""
path_part = nodeid.split("::", 1)[0].replace("\\", "/")
raw = tuple(p for p in path_part.split("/") if p and p != ".")
parts = raw[2:] if len(raw) >= 3 and raw[0] == "tests" and raw[1] == "e2e" else raw
if len(parts) <= 1:
return "root"
return parts[0]
def dedupe_covers(marker_args: Iterable[tuple[object, ...]]) -> tuple[str, ...]:
"""Flatten @pytest.mark.covers arg lists into unique, order-preserving cell
ids, dropping anything that is not a non-empty string."""
return tuple(dict.fromkeys(arg for args in marker_args for arg in args if isinstance(arg, str) and arg))
def covers_from_item(item: pytest.Item) -> tuple[str, ...]:
"""Read @pytest.mark.covers cell ids off a pytest Item, order-preserving."""
return dedupe_covers(marker.args for marker in item.iter_markers(name="covers"))
def result_properties(item: pytest.Item) -> tuple[tuple[str, str], ...]:
"""The custom signals a standard reporter cannot derive: the normalized suite
package and the comma-joined coverage-registry cell ids this test covers."""
return (
("package", package_from_nodeid(item.nodeid)),
("covers", ",".join(covers_from_item(item))),
)
def attach_result_properties(item: pytest.Item) -> None:
"""Attach result_properties to an item's user_properties, idempotently: a
second call is a no-op, so a collection that runs the hook more than once
never emits duplicate <property> entries."""
if any(name == "package" for name, _ in item.user_properties):
return
item.user_properties.extend(result_properties(item))

View file

@ -1,6 +1,6 @@
"""LLM-translation suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
live in the parent tests/e2e/conftest.py. PassthroughClient holds the shared
Gateway, so the `resources` fixture cleans up keys this suite creates.
"""

View file

@ -49,10 +49,10 @@ kept commented out in `PROVIDERS` until they pass end-to-end here; re-enable the
uncommenting their entry.
Every provider is provisioned and asserted; the suite never skips a provider. Per
`tests/e2e/CLAUDE.md` the only sanctioned skip is the whole-suite proxy-liveness
skip, so a provider whose credentials or upstream realtime model are missing on the
gateway is a hard failure, not a skip. Give the gateway each provider's credentials
to turn its tests green.
`tests/e2e/CLAUDE.md` there is no sanctioned skip: the whole-suite proxy-liveness
probe hard-fails when no proxy answers, and a provider whose credentials or upstream
realtime model are missing on the gateway is likewise a hard failure, not a skip.
Give the gateway each provider's credentials to turn its tests green.
## Running
@ -63,5 +63,5 @@ the deployments itself), then
uv run pytest tests/e2e/llm_translation/realtime/ -v
```
The whole suite skips only when no proxy answers `GET /health/liveliness` at
The whole suite hard-fails at setup when no proxy answers `GET /health/liveliness` at
`LITELLM_PROXY_URL` (default `http://localhost:4000`).

View file

@ -1,6 +1,6 @@
"""Realtime suite's `client` and `realtime_models` fixtures.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
live in the parent tests/e2e/conftest.py. RealtimeClient holds the shared Gateway,
so the `resources` fixture cleans up keys this suite creates.

View file

@ -6,10 +6,10 @@ schema: the session lifecycle, the canonical response event sequence with a
reconstructed transcript and usage, and a full tool-call round-trip (call ->
tool result -> a follow-up response that uses the result).
One GA-speaking client validates every provider; only the model alias changes. A
provider whose realtime alias is not configured on the proxy skips (skip on
environment); once it is configured, a protocol failure is a hard failure. See
REALTIME_COVERAGE_MATRIX.md.
One GA-speaking client validates every provider; only the model alias changes.
Every provider is provisioned at session start, so a missing realtime alias is a
hard failure, not a skip; once configured, a protocol failure is likewise a hard
failure. See REALTIME_COVERAGE_MATRIX.md.
"""
import pytest

View file

@ -7,9 +7,9 @@ references the proxy resolves at call time, so adding a provider is a new type
rather than another inline body. Start the proxy with the Rust OCR path enabled:
Each case creates its deployment, drives a real /v1/ocr call, and asserts a
well-formed OCR document comes back. Per the e2e "skip on environment, fail on
behavior" rule, a case skips when no proxy answers but fails (never skips) once a
request reaches it: the proxy fetches each provider's referenced secrets, so a
well-formed OCR document comes back. Per the e2e hard-fail contract, a case
fails when no proxy answers and also fails once a request reaches it: the proxy
fetches each provider's referenced secrets, so a
missing credential surfaces as a live provider error rather than silent green.
"""

View file

@ -1,6 +1,6 @@
"""Management suite fixtures: the client plus a logged-in dashboard page.
Lifecycle/skip/marker live in the parent conftest. The browser fixtures drive
Lifecycle/liveness gate/marker live in the parent conftest. The browser fixtures drive
the dashboard the proxy serves at /ui, so browser tests exercise exactly what an
end user sees. playwright is an optional dependency loaded behind importorskip
inside the fixture, so the API tests in this suite collect and run without it:

16
tests/e2e/mcp/conftest.py Normal file
View file

@ -0,0 +1,16 @@
"""MCP suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness handling, and the
`e2e`/`covers` markers live in the parent tests/e2e/conftest.py. McpClient holds
the shared Gateway, so the `resources` fixture tears down whatever this suite
creates (keys via the Gateway, MCP servers via the deferred cleanups).
"""
import pytest
from mcp_client import McpClient, build_client
@pytest.fixture(scope="session")
def client() -> McpClient:
return build_client()

153
tests/e2e/mcp/mcp_client.py Normal file
View file

@ -0,0 +1,153 @@
"""Client for the MCP e2e suite: admin server registration plus the api_key tool
surface.
An admin registers an upstream MCP server through the management API
(`/v1/mcp/server`, persisted in the DB) and grants a virtual key access to it via
`object_permission.mcp_servers`. Keys then reach the server through the REST bridge
the proxy exposes for api_key auth (`/mcp-rest/tools/list`, `/mcp-rest/tools/call`),
which `user_api_key_auth` gates the same way the JSON-RPC `/mcp` surface does. The
request/response bodies are co-located here because only this suite speaks MCP.
"""
from __future__ import annotations
from dataclasses import dataclass
from pydantic import BaseModel, ConfigDict, Field, RootModel
from e2e_gateway import Gateway, build_gateway
from e2e_http import Headers, NoBody, Result, unwrap
from models import KeyGenerateBody, ObjectPermission
class ApiKeyHeaders(Headers):
x_litellm_api_key: str = Field(serialization_alias="x-litellm-api-key")
class McpServerNewBody(BaseModel):
server_name: str
alias: str
url: str
transport: str = "http"
class McpServerNewResponse(BaseModel):
server_id: str
class McpServerRow(BaseModel):
server_id: str
alias: str | None = None
url: str | None = None
class McpServersListResponse(RootModel[list[McpServerRow]]):
pass
class McpToolMcpInfo(BaseModel):
server_id: str | None = None
alias: str | None = None
class McpToolEntry(BaseModel):
name: str
description: str | None = None
mcp_info: McpToolMcpInfo | None = None
class McpToolsListResponse(BaseModel):
tools: list[McpToolEntry] = []
error: str | None = None
message: str | None = None
def tool_names_for_server(self, server_id: str) -> frozenset[str]:
return frozenset(
tool.name
for tool in self.tools
if tool.mcp_info is not None and tool.mcp_info.server_id == server_id
)
class McpCallToolBody(BaseModel):
name: str
arguments: dict[str, int]
server_id: str
class McpCallContent(BaseModel):
type: str | None = None
text: str | None = None
class McpCallToolResponse(BaseModel):
model_config = ConfigDict(populate_by_name=True)
content: list[McpCallContent] = []
is_error: bool | None = Field(default=None, alias="isError")
@property
def first_text(self) -> str | None:
return self.content[0].text if self.content else None
@dataclass(frozen=True, slots=True)
class McpClient:
gateway: Gateway
def register_server(self, *, server_name: str, alias: str, url: str) -> str:
return unwrap(
self.gateway.transport.post(
"/v1/mcp/server",
headers=self.gateway.transport.master,
json=McpServerNewBody(server_name=server_name, alias=alias, url=url),
response_type=McpServerNewResponse,
)
).server_id
def delete_server(self, server_id: str) -> None:
_ = self.gateway.transport.delete(
f"/v1/mcp/server/{server_id}",
headers=self.gateway.transport.master,
json=NoBody(),
response_type=NoBody,
)
def registered_servers(self) -> list[McpServerRow]:
return unwrap(
self.gateway.transport.get(
"/v1/mcp/server",
headers=self.gateway.transport.master,
params=NoBody(),
response_type=McpServersListResponse,
)
).root
def generate_key(self, *, user_id: str, mcp_servers: list[str] | None) -> str:
object_permission = (
ObjectPermission(mcp_servers=mcp_servers) if mcp_servers is not None else None
)
return self.gateway.generate_key(
KeyGenerateBody(models=[], user_id=user_id, object_permission=object_permission)
)
def list_tools(self, key: str) -> Result[McpToolsListResponse]:
return self.gateway.transport.get(
"/mcp-rest/tools/list",
headers=ApiKeyHeaders(x_litellm_api_key=key),
params=NoBody(),
response_type=McpToolsListResponse,
)
def call_tool(
self, key: str, *, server_id: str, name: str, arguments: dict[str, int]
) -> Result[McpCallToolResponse]:
return self.gateway.transport.post(
"/mcp-rest/tools/call",
headers=ApiKeyHeaders(x_litellm_api_key=key),
json=McpCallToolBody(name=name, arguments=arguments, server_id=server_id),
response_type=McpCallToolResponse,
)
def build_client() -> McpClient:
return McpClient(gateway=build_gateway())

View file

@ -0,0 +1,103 @@
"""Live e2e: a virtual key without MCP access is denied an MCP server's tools.
An admin registers an upstream MCP server through the management API (persisted in
the DB, picked up without a restart) and queues its deletion. Two keys are created
against that one server: one granted access through `object_permission.mcp_servers`
and one with no MCP grant at all. The permitted key is the control that proves the
upstream is alive and the tool is callable, so a failure on the denied key is an
authorization denial rather than a dead server. The denied key must then see none
of the server's tools on `tools/list` and must be refused with a 403 on
`tools/call`.
Both the recorded state (the server is registered; the permitted key resolves its
tools) and the enforced behavior (the unpermitted key sees nothing and is blocked)
are asserted, so a regression that leaks tools to an ungranted key or drops the
call-time permission check fails here.
"""
import os
import pytest
from e2e_config import unique_marker
from e2e_http import UnknownApiError, unwrap
from lifecycle import ResourceManager
from mcp_client import McpClient
pytestmark = pytest.mark.e2e
MCP_UPSTREAM_URL = os.environ.get("E2E_MCP_UPSTREAM_URL", "http://mcp-upstream:8090/mcp")
MATH_TOOLS = frozenset({"add", "multiply"})
def _register_math_server(client: McpClient, resources: ResourceManager) -> str:
name = f"e2e_math_{unique_marker()}"
server_id = client.register_server(server_name=name, alias=name, url=MCP_UPSTREAM_URL)
resources.defer(lambda: client.delete_server(server_id))
return server_id
def _key(client: McpClient, resources: ResourceManager, *, mcp_servers: list[str] | None) -> str:
label = "allowed" if mcp_servers else "denied"
key = client.generate_key(user_id=f"e2e-mcp-{label}-{unique_marker()}", mcp_servers=mcp_servers)
resources.defer(lambda: client.gateway.delete_key(key))
return key
def _assert_registered(client: McpClient, server_id: str) -> None:
registered = {row.server_id for row in client.registered_servers()}
assert server_id in registered, f"registered server {server_id} absent from /v1/mcp/server: {registered}"
class TestMcpKeyWithoutAccessIsDenied:
@pytest.mark.covers("mcp.list_tools.api_key.denied_without_permission")
def test_list_tools_denied_without_permission(
self, client: McpClient, resources: ResourceManager
) -> None:
server_id = _register_math_server(client, resources)
_assert_registered(client, server_id)
permitted_key = _key(client, resources, mcp_servers=[server_id])
denied_key = _key(client, resources, mcp_servers=None)
permitted_tools = unwrap(client.list_tools(permitted_key)).tool_names_for_server(server_id)
assert MATH_TOOLS <= permitted_tools, (
f"granted key did not see the server's tools (upstream dead or grant not applied): "
f"{permitted_tools}"
)
denied_tools = unwrap(client.list_tools(denied_key)).tool_names_for_server(server_id)
assert denied_tools == frozenset(), (
f"ungranted key saw the server's tools; tools/list leaked across the permission "
f"boundary: {denied_tools}"
)
@pytest.mark.covers("mcp.call_tool.api_key.denied_without_permission")
def test_call_tool_denied_without_permission(
self, client: McpClient, resources: ResourceManager
) -> None:
server_id = _register_math_server(client, resources)
_assert_registered(client, server_id)
permitted_key = _key(client, resources, mcp_servers=[server_id])
denied_key = _key(client, resources, mcp_servers=None)
permitted_tools = unwrap(client.list_tools(permitted_key)).tool_names_for_server(server_id)
assert "add" in permitted_tools, (
f"granted key did not discover the add tool (upstream dead or grant not applied): "
f"{permitted_tools}"
)
permitted_call = unwrap(
client.call_tool(permitted_key, server_id=server_id, name="add", arguments={"a": 3, "b": 4})
)
assert permitted_call.is_error is not True, f"granted key's tool call errored: {permitted_call}"
assert permitted_call.first_text == "7", (
f"granted key's add(3, 4) did not return 7 (upstream not reachable): {permitted_call}"
)
match client.call_tool(denied_key, server_id=server_id, name="add", arguments={"a": 3, "b": 4}):
case UnknownApiError(status_code=403, body=body):
assert "access_denied" in body, f"403 was not an MCP access denial: {body}"
case other:
pytest.fail(f"ungranted key's tool call was not refused with 403 access_denied: {other}")

View file

@ -39,6 +39,10 @@ class KeyMetadata(BaseModel):
logging: list[KeyLoggingCallback] | None = None
class ObjectPermission(BaseModel):
mcp_servers: list[str] | None = None
class KeyGenerateBody(BaseModel):
models: list[str] = []
duration: str | None = None
@ -57,6 +61,7 @@ class KeyGenerateBody(BaseModel):
rpm_limit: int | None = None
allowed_routes: list[str] | None = None
metadata: KeyMetadata | None = None
object_permission: ObjectPermission | None = None
class KeyGenerateResponse(BaseModel):

View file

@ -33,12 +33,26 @@ _TEAM_READY_SLEEP_SECONDS = 0.4
class UserNewBody(BaseModel):
max_budget: float
budget_duration: str | None = None
class UserNewResponse(BaseModel):
user_id: str
class UserInfoParams(BaseModel):
user_id: str
class UserInfoRow(BaseModel):
spend: float | None = None
max_budget: float | None = None
class UserInfoResponse(BaseModel):
user_info: UserInfoRow | None = None
class UserDeleteBody(BaseModel):
user_ids: list[str]
@ -51,6 +65,7 @@ class CustomerNewBody(BaseModel):
class OrgNewBody(BaseModel):
organization_alias: str
max_budget: float
budget_duration: str | None = None
class OrgNewResponse(BaseModel):
@ -61,6 +76,14 @@ class OrgDeleteBody(BaseModel):
organization_ids: list[str]
class OrgInfoParams(BaseModel):
organization_id: str
class OrgInfoResponse(BaseModel):
budget_id: str | None = None
class TeamMember(BaseModel):
role: str
user_id: str
@ -69,6 +92,7 @@ class TeamMember(BaseModel):
class TeamNewBody(BaseModel):
team_alias: str
max_budget: float | None = None
budget_duration: str | None = None
organization_id: str | None = None
budget_limits: list[BudgetWindow] | None = None
@ -244,12 +268,12 @@ class BudgetClient:
# ---- internal user --------------------------------------------------
def create_user(self, *, max_budget: float) -> str:
def create_user(self, *, max_budget: float, budget_duration: str | None = None) -> str:
return unwrap(
self.gateway.transport.post(
"/user/new",
headers=self.gateway.transport.master,
json=UserNewBody(max_budget=max_budget),
json=UserNewBody(max_budget=max_budget, budget_duration=budget_duration),
response_type=UserNewResponse,
)
).user_id
@ -262,6 +286,19 @@ class BudgetClient:
response_type=NoBody,
)
def user_info(self, user_id: str) -> UserInfoRow | None:
result = self.gateway.transport.get(
"/user/info",
headers=self.gateway.transport.master,
params=UserInfoParams(user_id=user_id),
response_type=UserInfoResponse,
)
match result:
case Success(data=data):
return data.user_info
case _:
return None
# ---- customer / end-user -------------------------------------------
def create_customer(self, customer_id: str, *, max_budget: float) -> str:
@ -275,16 +312,36 @@ class BudgetClient:
# ---- organization ---------------------------------------------------
def create_org(self, *, max_budget: float, alias: str) -> str:
def create_org(self, *, max_budget: float, alias: str, budget_duration: str | None = None) -> str:
return unwrap(
self.gateway.transport.post(
"/organization/new",
headers=self.gateway.transport.master,
json=OrgNewBody(organization_alias=alias, max_budget=max_budget),
json=OrgNewBody(
organization_alias=alias,
max_budget=max_budget,
budget_duration=budget_duration,
),
response_type=OrgNewResponse,
)
).organization_id
def org_budget_id(self, org_id: str) -> str | None:
"""The id of the budget row backing an org; its budget_reset_at is read via
budget_info (LIT-4570: /organization/new stores budget_duration without
scheduling budget_reset_at, so the reset job's first tick schedules it)."""
result = self.gateway.transport.get(
"/organization/info",
headers=self.gateway.transport.master,
params=OrgInfoParams(organization_id=org_id),
response_type=OrgInfoResponse,
)
match result:
case Success(data=data):
return data.budget_id
case _:
return None
def delete_org(self, org_id: str) -> None:
_ = self.gateway.transport.delete(
"/organization/delete",
@ -300,6 +357,7 @@ class BudgetClient:
*,
alias: str,
max_budget: float | None = None,
budget_duration: str | None = None,
organization_id: str | None = None,
budget_limits: list[BudgetWindow] | None = None,
) -> str:
@ -310,6 +368,7 @@ class BudgetClient:
json=TeamNewBody(
team_alias=alias,
max_budget=max_budget,
budget_duration=budget_duration,
organization_id=organization_id,
budget_limits=budget_limits,
),

View file

@ -1,6 +1,6 @@
"""Budgets suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
live in the parent tests/e2e/conftest.py. BudgetClient holds the shared Gateway,
so the `resources` fixture cleans up keys through it; tests register entity deletes
via `resources.defer(...)`.

View file

@ -4,7 +4,7 @@ Each entity is an E2ECase (lifecycle.E2ECase) driven by run_case: init() creates
the budgeted entity + a key, run() drives spend until a `budget_exceeded` block,
teardown() deletes everything init() created (always runs, even on failure/skip).
Covers the entities with no prior live coverage - internal user, end-user,
organization, team member. See BUDGET_TEST_COVERAGE_MATRIX.md.
organization, team member - plus key and team. See BUDGET_TEST_COVERAGE_MATRIX.md.
A non-budget error fails hard (never a skip); if calls never get blocked, budget
enforcement is broken -> fail.
@ -18,16 +18,17 @@ import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from e2e_http import StreamingResponse, require_successful_call
from lifecycle import run_case
pytestmark = pytest.mark.e2e
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> None:
"""Send paid calls until the entity's budget blocks one. Key/user/org/member
block within a couple calls off real-time reservation counters; the end-user
budget enforces off table spend that lands on the batch write, so it takes a
few more. A non-budget error fails hard (never a skip)."""
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> StreamingResponse:
"""Send paid calls until the entity's budget blocks one; return the blocked
response so callers can assert on its shape. Key/user/org/member block within
a couple calls off real-time reservation counters; the end-user budget
enforces off table spend that lands on the batch write, so it takes a few
more. A non-budget error fails hard (never a skip)."""
for _ in range(40):
result = client.chat(
key,
@ -37,7 +38,7 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") ->
user=user or None,
)
if is_budget_block(result):
return
return result
require_successful_call(result)
time.sleep(2)
pytest.fail("budget never enforced within the call budget")
@ -69,18 +70,85 @@ class _BudgetCase:
class KeyBudgetCase(_BudgetCase):
"""A bare key (no team_id / user_id) carrying its own max_budget, so only the
key-level budget can be the thing that blocks. The refusal must be a 429
budget_exceeded; any other error already fails via _assert_budget_blocks."""
def init(self) -> None:
self.key = self.client.generate_key(max_budget=3e-6)
self._undo.append(lambda: self.client.delete_key(self.key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
class TeamBudgetCase(_BudgetCase):
"""An admin caps a whole team: two keys under a tiny-budget team, neither with
a key-level budget. Key A is driven until the team cap blocks it; key B's very
first call must then be refused too, proving the cap sits on the team, not the
key that spent. Both refusals must be 429 budget_exceeded."""
def init(self) -> None:
team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", max_budget=3e-6
)
self._undo.append(lambda: self.client.delete_team(team_id))
self.key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self.key))
self._sibling_key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self._sibling_key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
sibling = self.client.chat(
self._sibling_key,
"claude-haiku-4-5",
f"spend {unique_marker()}",
max_tokens=16,
)
assert is_budget_block(sibling) and sibling.status_code == 429, (
f"a sibling key on the capped team must get the same 429 budget_exceeded, "
f"got {sibling.status_code}: {sibling.body[:200]}"
)
class InternalUserBudgetCase(_BudgetCase):
"""A user's max_budget follows the person, not the key. The capped user holds
two personal keys (no team, no key budgets) plus a team-member key on an
uncapped team; once the first personal key is refused, the other two must be
refused as well - a second key is not a fresh allowance, and since #32005 the
user budget draws down team keys too. All refusals must be 429 budget_exceeded."""
def init(self) -> None:
user_id = self.client.create_user(max_budget=3e-6)
self._undo.append(lambda: self.client.delete_user(user_id))
# personal key (no team) -> the user budget governs
self.key = self.client.generate_key(user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self.key))
self._second_key = self.client.generate_key(user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self._second_key))
team_id = self.client.create_team(alias=f"e2e-budget-team-{unique_marker()}")
self._undo.append(lambda: self.client.delete_team(team_id))
self.client.add_team_member(team_id, user_id)
self._team_key = self.client.generate_key(team_id=team_id, user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self._team_key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
for label, key in (("second personal key", self._second_key), ("team-member key", self._team_key)):
result = self.client.chat(key, "claude-haiku-4-5", f"spend {unique_marker()}", max_tokens=16)
assert is_budget_block(result) and result.status_code == 429, (
f"the {label} of a user over budget must get the same 429 budget_exceeded, "
f"got {result.status_code}: {result.body[:200]}"
)
class EndUserBudgetCase(_BudgetCase):
@ -97,34 +165,67 @@ class EndUserBudgetCase(_BudgetCase):
class OrganizationBudgetCase(_BudgetCase):
"""Org carries the tiny budget; the team under it and the key carry none, so
the org is the only entity that can block (the historically weak link). The
refusal must be a 429 budget_exceeded that names the org as the blocker."""
def init(self) -> None:
# Org carries the tiny budget; the team under it has none, so a block here
# is org-level enforcement (the historically weak link).
org_id = self.client.create_org(
self._org_id = self.client.create_org(
max_budget=3e-6, alias=f"e2e-budget-org-{unique_marker()}"
)
self._undo.append(lambda: self.client.delete_org(org_id))
self._undo.append(lambda: self.client.delete_org(self._org_id))
team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", organization_id=org_id
alias=f"e2e-budget-team-{unique_marker()}", organization_id=self._org_id
)
self._undo.append(lambda: self.client.delete_team(team_id))
self.key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self.key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
assert f"Organization={self._org_id}" in blocked.body, (
f"refusal must name the org as the blocker, got: {blocked.body[:200]}"
)
class TeamMemberBudgetCase(_BudgetCase):
"""Member A's per-team budget is tiny while the team and both members' user
budgets are roomy (100.0), so the only cap that can trip is A's: a block
proves member-level enforcement and must be a 429 budget_exceeded. Teammate
B, uncapped on the same team, must keep serving after A is cut off, proving
the member cap does not leak onto the team or its members."""
def init(self) -> None:
# Member's per-team budget is tiny while the team has a large budget, so a
# block proves member-level (not team-level) enforcement.
team_id = self.client.create_team(
self._team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", max_budget=100.0
)
self._undo.append(lambda: self.client.delete_team(team_id))
user_id = self.client.create_user(max_budget=100.0)
self._undo.append(lambda: self.client.delete_user(user_id))
self.client.add_team_member(team_id, user_id, max_budget_in_team=3e-6)
self.key = self.client.generate_key(team_id=team_id, user_id=user_id)
self._undo.append(lambda: self.client.delete_team(self._team_id))
self._member_id = self.client.create_user(max_budget=100.0)
self._undo.append(lambda: self.client.delete_user(self._member_id))
self.client.add_team_member(self._team_id, self._member_id, max_budget_in_team=3e-6)
self.key = self.client.generate_key(team_id=self._team_id, user_id=self._member_id)
self._undo.append(lambda: self.client.delete_key(self.key))
teammate_id = self.client.create_user(max_budget=100.0)
self._undo.append(lambda: self.client.delete_user(teammate_id))
self.client.add_team_member(self._team_id, teammate_id)
self._teammate_key = self.client.generate_key(team_id=self._team_id, user_id=teammate_id)
self._undo.append(lambda: self.client.delete_key(self._teammate_key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
teammate = self.client.chat(
self._teammate_key,
"claude-haiku-4-5",
f"spend {unique_marker()}",
max_tokens=16,
)
require_successful_call(teammate)
def _case_id(case_cls: Type[_BudgetCase]) -> str:
@ -138,6 +239,10 @@ def _case_id(case_cls: Type[_BudgetCase]) -> str:
KeyBudgetCase,
marks=pytest.mark.covers("quota_management.budget.key.blocks_over_limit"),
),
pytest.param(
TeamBudgetCase,
marks=pytest.mark.covers("quota_management.budget.team.blocks_over_limit"),
),
pytest.param(
InternalUserBudgetCase,
marks=pytest.mark.covers("quota_management.budget.internal_user.blocks_over_limit"),

View file

@ -1,12 +1,4 @@
"""Live e2e: a key budget resets (zeroes spend) after its budget_duration.
Short budget_duration (30s) + the fast-rescheduled reset job: a key blocked for
exceeding its max_budget starts succeeding again once the duration elapses and the
reset job zeroes key.spend. Closes the reset-zeroing gap in
BUDGET_TEST_COVERAGE_MATRIX.md (reset_budget_for_litellm_keys), which the unit
suite covers but no live test did - distinct from the per-window reset in
test_multi_window_budget_e2e.py.
"""
"""Live e2e: an entity blocked over its max_budget serves again after its budget_duration window."""
import time
@ -19,42 +11,107 @@ from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
TINY_CAP = 3e-6
WINDOW = "30s"
RESET_DEADLINE_SECONDS = 150
def _call(client: BudgetClient, key: str):
return client.chat(
key, "claude-haiku-4-5", f"reset {unique_marker()}", max_tokens=16
)
return client.chat(key, "claude-haiku-4-5", f"reset {unique_marker()}", max_tokens=16)
@pytest.mark.covers("quota_management.budget.key.resets_after_window")
def test_key_budget_resets_after_duration(
client: BudgetClient, resources: ResourceManager
) -> None:
key = client.generate_key(max_budget=3e-6, budget_duration="30s")
resources.defer(lambda: client.delete_key(key))
# 1. exceed the budget -> litellm returns budget_exceeded
blocked = False
for _ in range(20):
def _drive_to_block(client: BudgetClient, key: str) -> None:
"""Spend until the cap blocks a call, staying under one window so the block
is observed before the reset job can fire; fail hard if enforcement never trips."""
for _ in range(12):
result = _call(client, key)
if is_budget_block(result):
blocked = True
break
return
require_successful_call(result)
time.sleep(2)
assert blocked, "key budget never enforced"
pytest.fail("budget never enforced before the window could reset")
# 2. once the 30s duration elapses + the reset job runs, key.spend zeroes and
# calls flow again. The window is wall-clock-aligned, so the reset lands up to
# a window later, then the rescheduler (~15-20s) zeroes the spend; allow
# generous headroom over that. A stuck rescheduler is caught by the wait-loop
# timeout, not this elapsed bound.
start = time.monotonic()
while time.monotonic() < start + 150:
def _poll_until_serves_again(client: BudgetClient, key: str) -> None:
"""Poll past the window until the blocked key serves again; every refusal must
stay a budget block, so a crashed reset path or provider error fails loudly."""
deadline = time.monotonic() + RESET_DEADLINE_SECONDS
while time.monotonic() < deadline:
time.sleep(5)
result = _call(client, key)
if result.ok:
assert time.monotonic() - start < 120, "reset too slow for a 30s budget"
return
assert is_budget_block(result), f"non-budget error: {result.body[:200]}"
pytest.fail("key budget never reset within 150s")
if not is_budget_block(result):
pytest.fail(f"non-budget error during reset wait: HTTP {result.status_code}: {result.body[:200]}")
pytest.fail(f"budget never reset within {RESET_DEADLINE_SECONDS}s")
class TestBudgetResetDiagonal:
@pytest.mark.covers("quota_management.budget.key.resets_after_window")
def test_bare_key_budget_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
key = client.generate_key(max_budget=TINY_CAP, budget_duration=WINDOW)
resources.defer(lambda: client.delete_key(key))
_drive_to_block(client, key)
_poll_until_serves_again(client, key)
@pytest.mark.covers("quota_management.budget.team.resets_after_window")
def test_team_budget_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
team_id = client.create_team(
alias=f"e2e-team-reset-{unique_marker()}", max_budget=TINY_CAP, budget_duration=WINDOW
)
resources.defer(lambda: client.delete_team(team_id))
key = client.generate_key(team_id=team_id)
resources.defer(lambda: client.delete_key(key))
_drive_to_block(client, key)
_poll_until_serves_again(client, key)
@pytest.mark.covers("quota_management.budget.organization.resets_after_window")
def test_org_budget_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
org_id = client.create_org(
max_budget=TINY_CAP, alias=f"e2e-org-reset-{unique_marker()}", budget_duration=WINDOW
)
resources.defer(lambda: client.delete_org(org_id))
team_id = client.create_team(alias=f"e2e-org-team-{unique_marker()}", organization_id=org_id)
resources.defer(lambda: client.delete_team(team_id))
key = client.generate_key(team_id=team_id)
resources.defer(lambda: client.delete_key(key))
budget_id = client.org_budget_id(org_id)
assert budget_id, "org created without a budget row"
deadline = time.monotonic() + 30
while not any(row.budget_reset_at for row in client.budget_info(budget_id)):
if time.monotonic() > deadline:
pytest.fail("org budget window never scheduled by the reset job")
time.sleep(2)
_drive_to_block(client, key)
_poll_until_serves_again(client, key)
@pytest.mark.covers("quota_management.budget.internal_user.resets_after_window")
def test_personal_key_user_budget_resets_after_window(
self, client: BudgetClient, resources: ResourceManager
) -> None:
user_id = client.create_user(max_budget=TINY_CAP, budget_duration=WINDOW)
resources.defer(lambda: client.delete_user(user_id))
key = client.generate_key(user_id=user_id)
resources.defer(lambda: client.delete_key(key))
_drive_to_block(client, key)
_poll_until_serves_again(client, key)
@pytest.mark.covers("quota_management.budget.internal_user.resets_after_window")
def test_team_member_key_user_budget_resets_after_window(
self, client: BudgetClient, resources: ResourceManager
) -> None:
user_id = client.create_user(max_budget=TINY_CAP, budget_duration=WINDOW)
resources.defer(lambda: client.delete_user(user_id))
team_id = client.create_team(alias=f"e2e-user-team-reset-{unique_marker()}")
resources.defer(lambda: client.delete_team(team_id))
client.add_team_member(team_id, user_id, max_budget_in_team=100.0)
key = client.generate_key(team_id=team_id, user_id=user_id)
resources.defer(lambda: client.delete_key(key))
_drive_to_block(client, key)
_poll_until_serves_again(client, key)

View file

@ -23,14 +23,16 @@ pytestmark = pytest.mark.e2e
WINDOW_SECONDS = 30 # the tight window; calls succeed again only after it elapses
# Prefer the OpenAI cheap model for this polling test: under the full stage suite
# Claude chat latency + ALB target idle timeout (~60s) can surface as awselb 502
# HTML mid-wait, which is not a budget signal. gpt-5.5 + 1 token stays well under
# that ceiling so the wait loop measures window reset, not provider/ALB timeout.
# HTML mid-wait, which is not a budget signal. gpt-5.5 stays well under that
# ceiling so the wait loop measures window reset, not provider/ALB timeout.
# max_tokens must be >1: gpt-5.5 refuses completions that hit the output limit
# mid-message when capped at 1 token.
MODEL = CHEAP_OPENAI_MODEL
def _call(client: BudgetClient, key: str):
return client.chat(
key, MODEL, f"window {unique_marker()}", max_tokens=1
key, MODEL, f"window {unique_marker()}", max_tokens=16
)

View file

@ -0,0 +1,119 @@
"""Live e2e: per-team-member budgets are enforced independently between members.
Two members share one team that has a large team budget. The tight member is capped
at a tiny per-team budget and spends past it; the roomy member has plenty of room.
Once the tight member is blocked with budget_exceeded, the roomy member still serves
on the same team, its calls land in the spend logs under its own user id, and the
tight member stays blocked. A shared or leaky member counter would either block the
roomy member too or let the tight member back through once its peer spent.
"""
import time
from collections.abc import Iterator
from dataclasses import dataclass
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import Success, require_successful_call
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage
pytestmark = pytest.mark.e2e
MODEL = "gpt-5.5"
TEAM_BUDGET = 100.0
TIGHT_MEMBER_BUDGET = 3e-6
ROOMY_MEMBER_BUDGET = 100.0
ROOMY_BURST = 3
@dataclass(frozen=True, slots=True)
class _Pair:
team_id: str
tight_user_id: str
roomy_user_id: str
tight_key: str
roomy_key: str
@pytest.fixture(scope="class")
def pair(client: BudgetClient) -> Iterator[_Pair]:
"""One team with a large budget and two members on it: a tight member capped at
a tiny per-team budget and a roomy member with headroom, each with their own key.
Shared across the class and torn down LIFO best-effort when it finishes."""
resources = ResourceManager(client=client.gateway)
try:
marker = unique_marker()
team_id = client.create_team(alias=f"e2e-member-iso-{marker}", max_budget=TEAM_BUDGET)
resources.defer(lambda: client.delete_team(team_id))
tight_user = client.create_user(max_budget=TEAM_BUDGET)
resources.defer(lambda: client.delete_user(tight_user))
roomy_user = client.create_user(max_budget=TEAM_BUDGET)
resources.defer(lambda: client.delete_user(roomy_user))
client.add_team_member(team_id, tight_user, max_budget_in_team=TIGHT_MEMBER_BUDGET)
client.add_team_member(team_id, roomy_user, max_budget_in_team=ROOMY_MEMBER_BUDGET)
tight_key = client.generate_key(team_id=team_id, user_id=tight_user)
resources.defer(lambda: client.delete_key(tight_key))
roomy_key = client.generate_key(team_id=team_id, user_id=roomy_user)
resources.defer(lambda: client.delete_key(roomy_key))
yield _Pair(
team_id=team_id,
tight_user_id=tight_user,
roomy_user_id=roomy_user,
tight_key=tight_key,
roomy_key=roomy_key,
)
finally:
resources.teardown()
def _roomy_send(client: BudgetClient, key: str) -> str:
"""One roomy-member call that must go through; returns its request id."""
match client.gateway.chat(
key,
ChatBody(
model=MODEL,
messages=[ChatMessage(role="user", content=f"roomy {unique_marker()}")],
max_tokens=16,
),
):
case Success(data=response):
assert response.id is not None, "roomy member call returned no id"
return response.id
case other:
pytest.fail(f"roomy member call failed while a peer was over budget: {other}")
class TestTeamMemberBudgetIsolation:
@pytest.mark.covers("quota_management.budget.team_member.isolates_per_member")
def test_blocked_member_does_not_block_peer(self, client: BudgetClient, pair: _Pair) -> None:
blocked = False
for _ in range(40):
result = client.chat(pair.tight_key, MODEL, f"tight {unique_marker()}", max_tokens=16)
if is_budget_block(result):
blocked = True
break
require_successful_call(result)
time.sleep(2)
assert blocked, "tight member's per-team budget never enforced"
sent = frozenset(_roomy_send(client, pair.roomy_key) for _ in range(ROOMY_BURST))
assert is_budget_block(
client.chat(pair.tight_key, MODEL, f"tight {unique_marker()}", max_tokens=16)
), "tight member stopped being blocked once the peer spent"
rows = client.gateway.poll_logs_for_key(
pair.roomy_key, predicate=lambda rs: bool(sent & {r.request_id for r in rs})
)
logged = [row for row in rows if row.request_id in sent]
assert logged, "none of the roomy member's calls reached the spend logs"
for row in logged:
assert row.user == pair.roomy_user_id, (
f"roomy call {row.request_id} logged under user {row.user}, not {pair.roomy_user_id}"
)
assert row.team_id == pair.team_id, (
f"roomy call {row.request_id} logged under team {row.team_id}, not {pair.team_id}"
)

View file

@ -0,0 +1,79 @@
"""Live e2e: a per-user max_budget is enforced across ALL of that user's keys.
An internal user's budget governs every personal key it owns, not only the one
that happened to spend it down. One user with a tiny max_budget owns two keys:
driving the first key to a budget_exceeded block then makes a fresh, untouched
second key of the same user (which carries no budget of its own, so nothing but the
shared user budget can block it) reject the same way, and the user's recorded spend
has crossed the cap. A key-scoped-only budget would leave the second key serving.
"""
import time
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import StreamingResponse, require_successful_call
from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
MODEL = "gpt-5.5"
TINY_CAP = 3e-6
RECORDED_SPEND_DEADLINE_SECONDS = 90
SECOND_KEY_BLOCK_ATTEMPTS = 6
def _call(client: BudgetClient, key: str) -> StreamingResponse:
return client.chat(key, MODEL, f"across {unique_marker()}", max_tokens=16)
def _drive_to_block(client: BudgetClient, key: str, subject: str) -> None:
for _ in range(40):
result = _call(client, key)
if is_budget_block(result):
return
require_successful_call(result)
time.sleep(2)
pytest.fail(f"user budget never enforced on {subject} within the call budget")
def _expect_prompt_block(client: BudgetClient, key: str, subject: str) -> None:
"""The shared user budget is already exhausted before this key makes a single
call, so a key with no budget of its own must be rejected promptly. The small
bounded retry only absorbs spend-propagation lag between the two keys; it is far
below the spend a key-scoped budget would need to accumulate to block itself, so
a block here can only come from the shared user budget."""
for _ in range(SECOND_KEY_BLOCK_ATTEMPTS):
result = _call(client, key)
if is_budget_block(result):
return
require_successful_call(result)
time.sleep(2)
pytest.fail(
f"{subject} was not blocked by the shared user budget within {SECOND_KEY_BLOCK_ATTEMPTS} calls"
)
class TestUserBudgetAcrossKeys:
@pytest.mark.covers("quota_management.budget.internal_user.enforced_across_keys")
def test_user_budget_blocks_a_second_key(self, client: BudgetClient, resources: ResourceManager) -> None:
user_id = client.create_user(max_budget=TINY_CAP)
resources.defer(lambda: client.delete_user(user_id))
first_key = client.generate_key(user_id=user_id)
resources.defer(lambda: client.delete_key(first_key))
second_key = client.generate_key(user_id=user_id)
resources.defer(lambda: client.delete_key(second_key))
_drive_to_block(client, first_key, "the first key")
_expect_prompt_block(client, second_key, "the second key")
deadline = time.monotonic() + RECORDED_SPEND_DEADLINE_SECONDS
while time.monotonic() < deadline:
info = client.user_info(user_id)
if info is not None and (info.spend or 0.0) >= TINY_CAP:
return
time.sleep(5)
pytest.fail(f"user spend never reached the {TINY_CAP} cap in the recorded state")

View file

@ -1,6 +1,6 @@
"""Quota-management suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
live in the parent tests/e2e/conftest.py. QuotaClient holds the shared Gateway,
so the `resources` fixture cleans up keys through it.
"""

View file

@ -80,5 +80,5 @@ proxy + SpendLogs rows. Status: `covered` / `partial` / `gap`.
`proxy_batch_write_at` (~60s) means rows land late; every read polls to a deadline.
Fresh scoped key per test (isolation, xdist-safe, cleaned up). Assert invariants
(`spend > 0`, `total == prompt + completion`, aggregate == sum), not literal
$/token values, so pricing drift is not a failure. Skip on environment (no proxy /
no provider key), fail on behavior (a real 2xx call with a wrong/missing row).
$/token values, so pricing drift is not a failure. Hard-fail when no proxy
answers, fail on behavior (a real 2xx call with a wrong/missing row).

View file

@ -1,6 +1,6 @@
"""Spend-tracking suite's `client` fixture and driver-model registration.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
live in the parent tests/e2e/conftest.py. SpendClient exposes the shared Gateway
(GatewayProvider), so the `resources` fixture cleans up keys and customers this
suite creates.

View file

@ -1,6 +1,6 @@
"""Router suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
live in the parent tests/e2e/conftest.py. ComplexityRouterClient holds the shared
Gateway, so the `resources` fixture cleans up keys this suite creates.

View file

@ -172,7 +172,7 @@ def anthropic_messages():
"content": [
{
"type": "text",
"text": "Here is the full text of a complex legal agreement" * 400,
"text": "Here is the full text of a complex legal agreement" * 500,
"cache_control": {"type": "ephemeral"},
}
],

View file

@ -0,0 +1,40 @@
"""Deterministic upstream MCP server for the mcp e2e suite.
A tiny FastMCP server exposing `add` and `multiply` over streamable-http so the
suite has a self-hosted, offline upstream to register and exercise. DNS-rebinding
protection is turned off because the litellm container reaches this over the
compose network by service name (`mcp-upstream:8090`), not localhost, and the
stack is an isolated throwaway. Bind host/port come from MCP_HOST/MCP_PORT.
"""
import os
from mcp.server.fastmcp import FastMCP
from mcp.server.transport_security import TransportSecuritySettings
mcp: FastMCP = FastMCP(
"e2e-math",
host=os.getenv("MCP_HOST", "0.0.0.0"),
port=int(os.getenv("MCP_PORT", "8090")),
transport_security=TransportSecuritySettings(enable_dns_rebinding_protection=False),
)
@mcp.tool()
def add(a: int, b: int) -> int:
"""Add two integers"""
return a + b
@mcp.tool()
def multiply(a: int, b: int) -> int:
"""Multiply two integers"""
return a * b
def main() -> None:
mcp.run(transport="streamable-http")
if __name__ == "__main__":
main()

View file

@ -935,8 +935,8 @@ async def test_get_tools_from_mcp_servers():
mcp_auth_header=mock_auth_header,
mcp_servers=["server1"],
)
assert len(result) == 1, "Should only return tools from server1"
assert result[0].name == "tool1", "Should return tool from server1"
assert len(result.tools) == 1, "Should only return tools from server1"
assert result.tools[0].name == "tool1", "Should return tool from server1"
# Test Case 2: Without specific MCP servers
# Create a different mock manager for the second test case
@ -978,9 +978,9 @@ async def test_get_tools_from_mcp_servers():
mcp_auth_header=mock_auth_header,
mcp_servers=None,
)
assert len(result) == 2, "Should return tools from all servers"
assert len(result.tools) == 2, "Should return tools from all servers"
assert (
result[0].name == "tool1" and result[1].name == "tool2"
result.tools[0].name == "tool1" and result.tools[1].name == "tool2"
), "Should return tools from all servers"
#
@ -1015,8 +1015,8 @@ async def test_get_tools_from_mcp_servers():
mcp_auth_header=mock_auth_header,
mcp_servers=["group-a"],
)
assert len(result) == 1, "Should only return tools from server3"
assert result[0].name == "tool1", "Should return tool from server1"
assert len(result.tools) == 1, "Should only return tools from server3"
assert result.tools[0].name == "tool1", "Should return tool from server1"
except AssertionError as e:
pytest.fail(f"Test failed: {str(e)}")
@ -2436,11 +2436,12 @@ async def test_filter_tools_by_allowed_tools_integration():
mock_client_constructor,
):
# Call _get_tools_from_mcp_servers which should apply the filtering
filtered_tools = await _get_tools_from_mcp_servers(
listing = await _get_tools_from_mcp_servers(
user_api_key_auth=mock_user_auth,
mcp_auth_header="Bearer test_token",
mcp_servers=None, # Get from all servers
)
filtered_tools = listing.tools
# Verify that only allowed tools are returned
assert (
@ -2549,11 +2550,12 @@ async def test_filter_tools_by_disallowed_tools_integration():
mock_client_constructor,
):
# Call _get_tools_from_mcp_servers which should apply the filtering
filtered_tools = await _get_tools_from_mcp_servers(
listing = await _get_tools_from_mcp_servers(
user_api_key_auth=mock_user_auth,
mcp_auth_header="Bearer test_token",
mcp_servers=None, # Get from all servers
)
filtered_tools = listing.tools
# Verify that only safe tools are returned (dangerous tools filtered out)
assert (
@ -2650,11 +2652,12 @@ async def test_filter_tools_no_restrictions_integration():
mock_client_constructor,
):
# Call _get_tools_from_mcp_servers which should apply the filtering
filtered_tools = await _get_tools_from_mcp_servers(
listing = await _get_tools_from_mcp_servers(
user_api_key_auth=mock_user_auth,
mcp_auth_header="Bearer test_token",
mcp_servers=None, # Get from all servers
)
filtered_tools = listing.tools
# Should return all tools when no restrictions
assert (

View file

@ -1820,7 +1820,7 @@ def test_init_auto_router_deployment_success(mock_auto_router, model_list):
# Verify the auto-router was added to the router's auto_routers dict
assert "test-auto-router" in router.auto_routers
assert router.auto_routers["test-auto-router"] == mock_auto_router_instance
assert router.auto_routers["test-auto-router"][0].strategy == mock_auto_router_instance
@patch("litellm.router_strategy.auto_router.auto_router.AutoRouter")
@ -1833,7 +1833,11 @@ def test_init_auto_router_deployment_duplicate_model_name(mock_auto_router, mode
mock_auto_router.return_value = mock_auto_router_instance
# Add an existing auto-router
router.auto_routers["test-auto-router"] = mock_auto_router_instance
from litellm.types.router import TaggedPreRoutingStrategy
router.auto_routers["test-auto-router"] = [
TaggedPreRoutingStrategy(tags=(), strategy=mock_auto_router_instance)
]
# Try to add another auto-router with the same name
litellm_params = LiteLLM_Params(
@ -1849,7 +1853,7 @@ def test_init_auto_router_deployment_duplicate_model_name(mock_auto_router, mode
)
with pytest.raises(
ValueError, match="Auto-router deployment test-auto-router already exists"
ValueError, match="Auto-router deployment test-auto-router with tags .* already exists"
):
router.init_auto_router_deployment(deployment)

View file

@ -2,7 +2,9 @@ import copy
import datetime
import json
import os
import subprocess
import sys
import textwrap
import unittest
from typing import List, Optional, Tuple
from unittest.mock import ANY, MagicMock, Mock, patch
@ -1533,3 +1535,242 @@ class TestApplyToAnthropicMessagesRequest:
sys_blocks = sum(1 for b in (result_sys or []) if isinstance(b, dict) and b.get("cache_control") is not None)
total_blocks = sys_blocks + sum(AnthropicCacheControlHook._count_cache_control_blocks(m) for m in result_msgs)
assert total_blocks <= 4
class TestEnableAnthropicPromptCaching:
"""Auto-injected default breakpoints via litellm.enable_anthropic_prompt_caching."""
MESSAGES: List[AllMessageValues] = [
{"role": "system", "content": "a long system prompt"},
{"role": "user", "content": "first turn"},
{"role": "assistant", "content": "a reply"},
{"role": "user", "content": "latest turn"},
]
def _points(self, model="claude-sonnet-4-5", provider="anthropic", messages=None, system=None, tools=None):
return AnthropicCacheControlHook.get_default_injection_points(
messages=copy.deepcopy(self.MESSAGES) if messages is None else messages,
system=system,
model=model,
custom_llm_provider=provider,
tools=tools,
)
def test_disabled_by_default(self):
assert litellm.enable_anthropic_prompt_caching is False
assert self._points() == []
def test_injects_system_and_trailing_turn(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert self._points() == [
{"location": "message", "role": "system", "index": None, "control": {"type": "ephemeral"}},
{"location": "message", "role": None, "index": -1, "control": {"type": "ephemeral"}},
]
def test_bedrock_claude_is_injected(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
points = self._points(model="us.anthropic.claude-sonnet-4-5-20250929-v1:0", provider="bedrock")
assert [p["index"] for p in points] == [None, -1]
@pytest.mark.parametrize("model, provider", [("gpt-4o", "openai"), ("gemini-2.0-flash", "gemini")])
def test_non_anthropic_providers_never_injected(self, monkeypatch, model, provider):
"""These report supports_prompt_caching=True but never consume cache_control markers."""
from litellm.utils import supports_prompt_caching
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert supports_prompt_caching(model=model, custom_llm_provider=provider) is True
assert self._points(model=model, provider=provider) == []
def test_model_without_caching_support_not_injected(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert self._points(model="anthropic.claude-3-5-sonnet-20240620-v1:0", provider="bedrock") == []
def test_stands_down_when_client_sent_cache_control(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
messages = [
{"role": "system", "content": [{"type": "text", "text": "s", "cache_control": {"type": "ephemeral"}}]},
{"role": "user", "content": "latest turn"},
]
assert self._points(messages=messages) == []
def test_stands_down_when_system_block_has_cache_control(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
system = [{"type": "text", "text": "s", "cache_control": {"type": "ephemeral"}}]
assert self._points(messages=[{"role": "user", "content": "hi"}], system=system) == []
@staticmethod
def _tools(count: int, cached: bool) -> List[dict]:
tool: dict = {"type": "function", "function": {"name": "t", "description": "d", "parameters": {}}}
if cached:
tool["cache_control"] = {"type": "ephemeral"}
return [{**tool, "function": {**tool["function"], "name": f"t{i}"}} for i in range(count)]
def test_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
"""Caching just the tool definitions is a normal client pattern, and those
breakpoints count toward the provider's four-block limit. Three of them plus
our two would be five, which Anthropic rejects outright."""
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert self._points(tools=self._tools(3, cached=True)) == []
def test_injects_when_tools_carry_no_cache_control(self, monkeypatch):
"""Tools alone must not suppress injection; only client-marked ones do."""
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert [p["index"] for p in self._points(tools=self._tools(3, cached=False))] == [None, -1]
@pytest.mark.parametrize("tools", [None, []])
def test_absent_tools_do_not_suppress_injection(self, monkeypatch, tools):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert [p["index"] for p in self._points(tools=tools)] == [None, -1]
def test_seed_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
"""Same guard on the /chat/completions seeding path."""
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
params: dict = {}
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=params,
messages=copy.deepcopy(self.MESSAGES),
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
tools=self._tools(3, cached=True),
)
assert "cache_control_injection_points" not in params
def test_v1_messages_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
"""Same guard on the /v1/messages path, where tools reach the hook directly."""
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
messages = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
copy.deepcopy(messages),
"sys",
{},
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
tools=self._tools(3, cached=True),
)
assert result_sys == "sys"
assert result_msgs == messages
def test_default_ttl_is_anthropics_five_minute_cache(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert all(p["control"] == {"type": "ephemeral"} for p in self._points())
@pytest.mark.parametrize("ttl", ["5m", "1h"])
def test_ttl_override_applied(self, monkeypatch, ttl):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
monkeypatch.setattr(litellm, "anthropic_prompt_caching_ttl", ttl)
assert all(p["control"] == {"type": "ephemeral", "ttl": ttl} for p in self._points())
def test_seed_does_not_override_configured_points(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
configured = [{"location": "message", "role": "user", "index": 0}]
params = {"cache_control_injection_points": configured}
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=params,
messages=copy.deepcopy(self.MESSAGES),
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert params["cache_control_injection_points"] is configured
def test_seed_adds_defaults_when_enabled(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
params: dict = {}
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=params,
messages=copy.deepcopy(self.MESSAGES),
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert [p["index"] for p in params["cache_control_injection_points"]] == [None, -1]
def test_seed_is_noop_when_disabled(self):
params: dict = {}
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=params,
messages=copy.deepcopy(self.MESSAGES),
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert params == {}
def test_v1_messages_applies_defaults_end_to_end(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
messages = [
{"role": "user", "content": [{"type": "text", "text": "first"}]},
{"role": "assistant", "content": [{"type": "text", "text": "reply"}]},
{"role": "user", "content": [{"type": "text", "text": "latest"}]},
]
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
messages,
"a system prompt",
{},
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert result_sys == [{"type": "text", "text": "a system prompt", "cache_control": {"type": "ephemeral"}}]
assert result_msgs[-1]["content"][-1]["cache_control"] == {"type": "ephemeral"}
assert "cache_control" not in result_msgs[0]["content"][-1]
def test_v1_messages_is_noop_when_disabled(self):
messages = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
messages,
"sys",
{},
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert result_sys == "sys"
assert result_msgs == messages
class TestAnthropicPromptCachingEnvVars:
"""Both settings are read from the environment at import, so an admin can enable
auto-caching without a config file. Each case re-imports litellm in a subprocess
so the env is read fresh without contaminating this process's module graph.
"""
@staticmethod
def _import_litellm_with_env(env_override: dict) -> Tuple[bool, Optional[str]]:
env = os.environ.copy()
env.pop("LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING", None)
env.pop("LITELLM_ANTHROPIC_PROMPT_CACHING_TTL", None)
env.update(env_override)
script = textwrap.dedent(
"""
import json, litellm
print(json.dumps([litellm.enable_anthropic_prompt_caching, litellm.anthropic_prompt_caching_ttl]))
"""
)
result = subprocess.run(
[sys.executable, "-c", script], capture_output=True, text=True, env=env, timeout=300
)
assert result.returncode == 0, result.stderr
enabled, ttl = json.loads(result.stdout.strip().splitlines()[-1])
return enabled, ttl
def test_unset_env_leaves_auto_caching_off(self):
assert self._import_litellm_with_env({}) == (False, None)
@pytest.mark.parametrize("value", ["true", "True", "TRUE"])
def test_env_enables_auto_caching_case_insensitively(self, value):
enabled, _ = self._import_litellm_with_env({"LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING": value})
assert enabled is True
@pytest.mark.parametrize("value", ["false", "0", "yes", ""])
def test_env_only_enables_on_true(self, value):
enabled, _ = self._import_litellm_with_env({"LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING": value})
assert enabled is False
@pytest.mark.parametrize("value", ["5m", "1h"])
def test_ttl_env_is_applied(self, value):
_, ttl = self._import_litellm_with_env({"LITELLM_ANTHROPIC_PROMPT_CACHING_TTL": value})
assert ttl == value
@pytest.mark.parametrize("value", ["10m", "1H", "3600", "ephemeral"])
def test_unsupported_ttl_env_falls_back_to_provider_default(self, value):
"""An unparseable TTL must fall back to Anthropic's 5m default, never reach the provider verbatim."""
_, ttl = self._import_litellm_with_env({"LITELLM_ANTHROPIC_PROMPT_CACHING_TTL": value})
assert ttl is None

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