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

This commit is contained in:
Devin AI 2026-07-18 03:07:08 +00:00
commit b6e94b3a4a
233 changed files with 14162 additions and 4011 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

@ -19,12 +19,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"
@ -233,21 +227,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"
@ -264,17 +243,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"
@ -310,7 +278,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",
@ -608,15 +575,6 @@ dependencies = [
"hashbrown",
]
[[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"
@ -718,15 +676,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"
@ -803,29 +752,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",
@ -834,19 +780,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",
@ -854,9 +799,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",
@ -866,13 +811,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",
]
@ -1321,9 +1265,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"
@ -1559,12 +1503,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

@ -15,8 +15,8 @@ repository = "https://github.com/BerriAI/litellm"
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", "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

@ -168,7 +168,7 @@ fn aocr(
.await
.map_err(core_error_to_pyerr)?;
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

@ -2,7 +2,7 @@ import os
import sys
from typing import List, Literal, Optional
from litellm.litellm_core_utils.env_utils import get_env_int
from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_or_none
DEFAULT_HEALTH_CHECK_PROMPT = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"))
AZURE_DEFAULT_RESPONSES_API_VERSION = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview"))
@ -269,9 +269,18 @@ TOOL_POLICY_CACHE_TTL_SECONDS = int(os.getenv("TOOL_POLICY_CACHE_TTL_SECONDS", 6
MAX_SIZE_IN_MEMORY_QUEUE = int(os.getenv("MAX_SIZE_IN_MEMORY_QUEUE", int(LITELLM_ASYNCIO_QUEUE_MAXSIZE * 0.8)))
MAX_IN_MEMORY_QUEUE_FLUSH_COUNT = int(os.getenv("MAX_IN_MEMORY_QUEUE_FLUSH_COUNT", 1000))
###############################################################################################
MINIMUM_PROMPT_CACHE_TOKEN_COUNT = int(
os.getenv("MINIMUM_PROMPT_CACHE_TOKEN_COUNT", 1024)
) # minimum number of tokens to cache a prompt by Anthropic
# Providers will not cache a prefix below a minimum size. That minimum is per-model, not global:
# Anthropic's ranges from 512 to 4096 depending on the model, and can differ per platform for the
# same model. The real minimum is resolved from `prompt_cache_min_tokens` in the model cost map;
# this value is only the fallback for models the cost map has no entry for, and doubles as a global
# escape hatch when `MINIMUM_PROMPT_CACHE_TOKEN_COUNT` is explicitly set.
MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE: int | None = get_env_int_or_none("MINIMUM_PROMPT_CACHE_TOKEN_COUNT")
DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT = 1024
MINIMUM_PROMPT_CACHE_TOKEN_COUNT = (
MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE
if MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE is not None
else DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT
)
DEFAULT_TRIM_RATIO = float(
os.getenv("DEFAULT_TRIM_RATIO", 0.75)
) # default ratio of tokens to trim from the end of a prompt

View file

@ -17,6 +17,7 @@ from litellm.llms.base_llm.google_genai.transformation import (
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import CallTypes
from litellm.utils import ProviderConfigManager, client
if TYPE_CHECKING:
@ -39,6 +40,11 @@ base_llm_http_handler = BaseLLMHTTPHandler()
#################################################
def _mark_async_entrypoint(logging_obj: LiteLLMLoggingObj | None, marker: str, is_async: bool) -> None:
if logging_obj is not None:
logging_obj.model_call_details.setdefault("litellm_params", {})[marker] = is_async
class GenerateContentSetupResult(BaseModel):
"""Internal Type - Result of setting up a generate content call"""
@ -315,6 +321,8 @@ def generate_content(
try:
_is_async = kwargs.pop("agenerate_content", False)
_mark_async_entrypoint(kwargs.get("litellm_logging_obj"), CallTypes.agenerate_content.value, _is_async)
# Handle generationConfig parameter from kwargs for backward compatibility
if "generationConfig" in kwargs and config is None:
config = kwargs.pop("generationConfig")
@ -403,6 +411,8 @@ async def agenerate_content_stream(
try:
kwargs["agenerate_content_stream"] = True
_mark_async_entrypoint(kwargs.get("litellm_logging_obj"), CallTypes.agenerate_content_stream.value, True)
# Handle generationConfig parameter from kwargs for backward compatibility
if "generationConfig" in kwargs and config is None:
config = kwargs.pop("generationConfig")
@ -497,6 +507,8 @@ def generate_content_stream(
# Remove any async-related flags since this is the sync function
_is_async = kwargs.pop("agenerate_content_stream", False)
_mark_async_entrypoint(kwargs.get("litellm_logging_obj"), CallTypes.agenerate_content_stream.value, _is_async)
# Handle generationConfig parameter from kwargs for backward compatibility
if "generationConfig" in kwargs and config is None:
config = kwargs.pop("generationConfig")

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

@ -19,3 +19,19 @@ def get_env_int(env_var: str, default: int) -> int:
return int(raw)
except (ValueError, TypeError):
return default
def get_env_int_or_none(env_var: str) -> int | None:
"""Parse an environment variable as an integer, returning None when it is unset or unusable.
Use this instead of `get_env_int` when callers must distinguish "explicitly configured"
from "left at the default", for example when an override should take precedence over a
value resolved from somewhere else.
"""
raw = os.getenv(env_var)
if raw is None:
return None
try:
return int(raw.strip())
except (ValueError, TypeError):
return None

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(
@ -1531,6 +1534,9 @@ class Logging(LiteLLMLoggingBaseClass):
and litellm_params.get(CallTypes.aimage_generation.value, False) is not True
and litellm_params.get(CallTypes.atranscription.value, False) is not True
and litellm_params.get(CallTypes.allm_passthrough_route.value, False) is not True
and litellm_params.get(CallTypes.aanthropic_messages.value, False) is not True
and litellm_params.get(CallTypes.agenerate_content.value, False) is not True
and litellm_params.get(CallTypes.agenerate_content_stream.value, False) is not True
)
def _is_assembled_stream_success(self, result=None) -> bool:

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

@ -36,6 +36,7 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import CallTypes
from litellm.utils import ProviderConfigManager, client
from ..utils import is_reasoning_auto_summary_enabled
@ -236,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)
@ -425,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)
@ -463,6 +468,9 @@ def anthropic_messages_handler(
"model": original_model,
"custom_llm_provider": custom_llm_provider,
}
litellm_logging_obj.model_call_details.setdefault("litellm_params", {})[CallTypes.aanthropic_messages.value] = (
is_async
)
# Check if stream was converted for WebSearch interception
# This is set in the async wrapper above when stream=True is converted to stream=False

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

@ -1731,18 +1731,18 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"""
Check if the candidate token count is inclusive of the thinking token count
if prompttokencount + candidatesTokenCount == totalTokenCount, then the candidate token count is inclusive of the thinking token count
if promptTokenCount + candidatesTokenCount + toolUsePromptTokenCount == totalTokenCount, then the candidate token count is inclusive of the thinking token count
else the candidate token count is exclusive of the thinking token count
Addresses - https://github.com/BerriAI/litellm/pull/10141#discussion_r2052272035
"""
if usage_metadata.get("promptTokenCount", 0) + usage_metadata.get(
"candidatesTokenCount", 0
) == usage_metadata.get("totalTokenCount", 0):
return True
else:
return False
non_thinking_tokens = (
usage_metadata.get("promptTokenCount", 0)
+ usage_metadata.get("candidatesTokenCount", 0)
+ usage_metadata.get("toolUsePromptTokenCount", 0)
)
return non_thinking_tokens == usage_metadata.get("totalTokenCount", 0)
@staticmethod
def _calculate_usage(
@ -1888,12 +1888,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
response_tokens_details = CompletionTokensDetailsWrapper()
response_tokens_details.reasoning_tokens = reasoning_tokens
tool_use_prompt_tokens = usage_metadata.get("toolUsePromptTokenCount") or None
prompt_tokens_details = PromptTokensDetailsWrapper(
cached_tokens=cached_tokens,
audio_tokens=prompt_audio_tokens,
text_tokens=prompt_text_tokens,
image_tokens=prompt_image_tokens,
video_tokens=prompt_video_tokens,
tool_use_tokens=tool_use_prompt_tokens,
)
completion_tokens = response_tokens or completion_response["usageMetadata"].get("candidatesTokenCount", 0)
@ -1901,7 +1904,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
completion_tokens = reasoning_tokens + completion_tokens
## GET USAGE ##
usage = Usage(
prompt_tokens=usage_metadata.get("promptTokenCount", 0),
prompt_tokens=usage_metadata.get("promptTokenCount", 0) + (tool_use_prompt_tokens or 0),
completion_tokens=completion_tokens,
total_tokens=usage_metadata.get("totalTokenCount", 0),
prompt_tokens_details=prompt_tokens_details,

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

File diff suppressed because it is too large Load diff

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,
@ -3582,48 +3571,70 @@ if MCP_AVAILABLE:
# preemptive challenge and let downstream authorization
# return 403.
continue
if server and server.auth_type == MCPAuth.oauth2 and not oauth2_headers:
# For per-user OAuth servers, only skip the pre-emptive 401 when
# a stored token actually exists for this user+server pair.
# If no stored token exists, fail fast with 401 so clients can
# kick off PKCE/interactive OAuth flow immediately.
if server.needs_user_oauth_token:
if getattr(server, "delegate_auth_to_upstream", False) is True:
# Delegate-auth servers run upstream PKCE: challenge with
# the proxied resource_metadata (RFC 9728), not the
# gateway authorization_uri below which would authorize
# against the gateway instead of the upstream IdP.
www_authenticate = _get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": www_authenticate},
)
# The v2 resolver owns the existence check, so every authorization_code
# resolution (egress and this discovery challenge) runs through it.
if server and server.auth_type == MCPAuth.oauth2:
# The challenge decision is per oauth2 sub-mode, not per header:
# gateway-managed modes (M2M and interactive authorization_code)
# never receive a client-supplied upstream token, so a bearer in
# Authorization is a LiteLLM key (surfaced here as oauth2_headers)
# and must not suppress the challenge. Only the delegate mode
# treats a present bearer as the upstream token. The sub-mode is
# resolved the same way egress resolves it, via
# effective_oauth2_flow: an unstamped (null oauth2_flow) row with
# the M2M shape resolves to client_credentials, so the bare
# has_client_credentials column is never trusted here.
if MCPServerManager.effective_oauth2_flow(server) == "client_credentials":
# M2M: the gateway mints its own token at egress from the
# stored client credentials, so there is nothing to challenge.
continue
if getattr(server, "delegate_auth_to_upstream", False) is not True:
# Gateway-managed interactive (authorization_code): the only
# thing that authorizes egress is a stored per-user token, so
# challenge whenever one is absent, regardless of any bearer.
# The v2 resolver owns the existence check, so every
# authorization_code resolution (egress and this discovery
# challenge) runs through it.
if await global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth):
continue
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
_path = scope.get("_original_path") or scope.get("path", "") or ""
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
_path = scope.get("_original_path") or scope.get("path", "") or ""
# Pick the well-known AS-metadata form that matches the inbound route
# so strict RFC 9728 §3.2 clients can resolve it correctly.
if _path.startswith(f"/mcp/{server_name}"):
_as_url = f"{base_url}/.well-known/oauth-authorization-server/mcp/{server_name}"
else:
_as_url = f"{base_url}/.well-known/oauth-authorization-server/{server_name}"
authorization_uri = f'Bearer authorization_uri="{_as_url}"'
# Pick the well-known AS-metadata form that matches the inbound route
# so strict RFC 9728 §3.2 clients can resolve it correctly.
if _path.startswith(f"/mcp/{server_name}"):
_as_url = f"{base_url}/.well-known/oauth-authorization-server/mcp/{server_name}"
else:
_as_url = f"{base_url}/.well-known/oauth-authorization-server/{server_name}"
authorization_uri = f'Bearer authorization_uri="{_as_url}"'
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": authorization_uri},
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": authorization_uri},
)
if not oauth2_headers:
# Delegate-auth servers run upstream PKCE: a present bearer is
# the upstream token, so only challenge when it is absent, with
# the proxied resource_metadata (RFC 9728), not the gateway
# authorization_uri above which would authorize against the
# gateway instead of the upstream IdP.
www_authenticate = _get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": www_authenticate},
)
# Delegate server with a bearer present: it is the upstream token,
# so admit the session and move to the next target. Every oauth2
# sub-mode is terminal here (continue or raise) so no oauth2 server
# reaches the token_exchange / pass-through blocks below.
continue
# token_exchange (OBO): the caller supplied no subject token. Challenge at connect
# (transport level, where WWW-Authenticate survives) with the RFC 9728 resource_metadata

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:
@ -2575,6 +2649,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 +2662,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 +2772,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 +2821,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

@ -155,6 +155,40 @@ class SemanticToolFilterHook(CustomLogger):
return await self.filter.filter_tools(query=user_query, available_tools=expanded_tools)
def _selected_tool_names(self, filtered_tools: list[dict[str, Any]]) -> list[str]:
"""Names of the semantically selected tools, as produced by the MCP expansion."""
names = (self.filter._extract_tool_info(tool)[0] for tool in filtered_tools)
return [name for name in names if name]
@staticmethod
def _narrow_mcp_references(tools: list[Any], selected_tool_names: list[str]) -> list[Any]:
"""
Restrict each litellm_proxy MCP reference to the semantically selected tools.
The reference block is preserved rather than replaced with expanded tools, so the
MCP gateway still performs the expansion. That keeps the per-endpoint tool shape
and tool auto-execution intact. Expansion already applied any caller-supplied
allowed_tools, so this selection can only narrow a block further.
Whether an undecidable selection exposes every tool or none is owned by
SemanticMCPToolFilter.filter_tools, which returns the full set when nothing
matches; the same policy therefore governs references and plain tools. Passing an
empty selection through is safe rather than a hidden allow-all: the gateway reads
the union of every reference's allowed_tools and treats an empty union as unset.
"""
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
return [
(
{**tool, "allowed_tools": selected_tool_names}
if isinstance(tool, dict) and LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway([tool])
else tool
)
for tool in tools
]
def _is_mcp_tool(self, tool: object) -> bool:
"""
Check whether *tool* is registered in the MCP semantic router.
@ -261,36 +295,30 @@ class SemanticToolFilterHook(CustomLogger):
if self._should_expand_mcp_tools(tools):
verbose_proxy_logger.debug("Detected litellm_proxy MCP references, expanding before semantic filtering")
if not self.filter.enabled:
verbose_proxy_logger.debug("Semantic filter disabled, leaving MCP references untouched")
return None
try:
native_tools_before_expand = [t for t in tools if not (isinstance(t, dict) and t.get("type") == "mcp")]
expanded_tools = await self._expand_mcp_tools(tools, user_api_key_dict)
if not expanded_tools:
if native_tools_before_expand:
data["tools"] = native_tools_before_expand
verbose_proxy_logger.warning(
f"No MCP tools expanded, preserving {len(native_tools_before_expand)} native tools"
)
return data
verbose_proxy_logger.warning("No tools expanded from MCP references")
return None
if not self.filter.enabled:
data["tools"] = native_tools_before_expand + expanded_tools
verbose_proxy_logger.debug("Semantic filter disabled, forwarding expanded MCP tools unfiltered")
return data
filtered_expanded_tools = await self._filter_expanded_tools(data=data, expanded_tools=expanded_tools)
combined_tools = native_tools_before_expand + filtered_expanded_tools
data["tools"] = combined_tools
selected_tool_names = self._selected_tool_names(filtered_expanded_tools)
narrowed_tools = self._narrow_mcp_references(tools, selected_tool_names)
data["tools"] = narrowed_tools
self._emit_filter_metadata_safe(
data=data,
mcp_tools=expanded_tools,
filtered_mcp_tools=filtered_expanded_tools,
native_tools=native_tools_before_expand,
filtered_tools=combined_tools,
filtered_tools=narrowed_tools,
)
verbose_proxy_logger.info(
f"Expanded MCP references to {len(expanded_tools)} tools "

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

@ -2092,12 +2092,8 @@ async def cli_poll_key(
key_id: The CLI login session ID
team_id: Optional team ID to assign to the JWT. If provided, must be one of user's teams.
"""
from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
get_team_object,
get_user_object,
)
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
from litellm.proxy.proxy_server import user_api_key_cache
try:
flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=user_api_key_cache)
@ -2167,43 +2163,11 @@ async def cli_poll_key(
models=session_data.get("models", []),
)
try:
user_db_obj = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
)
except ValueError as e:
verbose_proxy_logger.debug(f"CLI poll: user lookup failed, proceeding without user budget: {e}")
user_db_obj = None
user_budget = user_db_obj.max_budget if user_db_obj is not None else None
team_budget: Optional[float] = None
team_budget_resolved = False
if team_id is not None:
try:
team_obj = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
team_budget = team_obj.max_budget
team_budget_resolved = True
except Exception:
pass
session_max_budget = (
litellm.max_ui_session_budget
if user_budget is None and (team_id is None or (team_budget_resolved and team_budget is None))
else None
)
jwt_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
user_info=user_info,
team_id=team_id,
team_alias=team_alias,
max_budget=session_max_budget,
max_budget=None,
)
# Delete cache entry (single-use)

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,
@ -56,11 +57,13 @@ class TierClassification(BaseModel):
_CLASSIFICATION_PROMPT_TEMPLATE = """Classify the complexity of the following user request into exactly one tier.
Judge the intellectual difficulty of answering correctly, not how short the request is.
Tiers:
- SIMPLE: factual lookups, greetings, short direct questions with no reasoning or code involved.
- MEDIUM: everyday requests needing some explanation or minor code/technical content.
- COMPLEX: requests involving non-trivial code, architecture, or multi-step technical work.
- REASONING: requests explicitly requiring step-by-step reasoning, analysis, or weighing tradeoffs.
- SIMPLE: greetings, chitchat, or factual lookups with a short known answer. Do not use SIMPLE for unsolved problems, proofs, deep theory, multi-step analysis, or non-trivial code, even if the request is only one sentence.
- MEDIUM: everyday requests that need some explanation, light reasoning, or minor code/technical content.
- COMPLEX: non-trivial code, architecture, multi-step technical work, or specialized domain depth.
- REASONING: open-ended analysis, proofs, famous hard problems, step-by-step reasoning, tradeoffs, or anything where a correct answer requires careful thought rather than a quick lookup.
{system_context}Request:
{prompt}"""
@ -171,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
@ -666,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.
@ -908,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,
@ -1002,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,
@ -1016,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

@ -8,14 +8,40 @@ from typing import List, Optional, cast
from litellm import verbose_logger
from litellm.caching.dual_cache import DualCache
from litellm.constants import DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT
from litellm.integrations.custom_logger import CustomLogger, Span
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import CallTypes, StandardLoggingPayload
from litellm.utils import is_prompt_caching_valid_prompt
from litellm.utils import get_prompt_cache_min_tokens, is_prompt_caching_valid_prompt
from ..prompt_caching_cache import PromptCachingCache
def _get_min_token_count_for_deployments(healthy_deployments: list[dict]) -> int:
"""
Returns the lowest minimum cacheable prefix across a model group.
This gate only decides whether the cache lookup is worth doing. It cannot cause a wrong pin,
because a deployment is only pinned when the cache already holds an entry for the prefix, and
entries are written by `async_log_success_event` against the deployment's real model. A model
that will not cache a prefix never records one, so there is nothing to pin it to.
That makes the lowest minimum in the group the correct threshold rather than the highest.
`model` here is the model-group alias the operator chose, not a model name, so the threshold
has to come from the deployments themselves, and a group may mix models whose minimums differ.
Taking the highest would skip the lookup for a prefix a lower-minimum member genuinely cached,
losing a cache hit it had earned. The lowest can only cost a lookup that finds nothing.
"""
return min(
(
get_prompt_cache_min_tokens(model=deployment["litellm_params"]["model"])
for deployment in healthy_deployments
if deployment.get("litellm_params", {}).get("model")
),
default=DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT,
)
class PromptCachingDeploymentCheck(CustomLogger):
def __init__(self, cache: DualCache):
self.cache = cache
@ -31,7 +57,8 @@ class PromptCachingDeploymentCheck(CustomLogger):
if messages is not None and is_prompt_caching_valid_prompt(
messages=messages,
model=model,
): # prompt > 1024 tokens
min_token_count=_get_min_token_count_for_deployments(healthy_deployments),
):
prompt_cache = PromptCachingCache(
cache=self.cache,
)

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,6 +128,7 @@ class SupportedGuardrailIntegrations(Enum):
RUBRIK = "rubrik"
VIGIL_GUARD = "vigil_guard"
REPELLOAI = "repelloai"
SINGULR = "singulr"
HEADROOM = "headroom"
COMPRESR = "compresr"
@ -932,6 +936,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

@ -299,6 +299,8 @@ class UsageMetadata(TypedDict, total=False):
candidatesTokenCount: int
responseTokenCount: int
cachedContentTokenCount: int
toolUsePromptTokenCount: int
toolUsePromptTokensDetails: List[PromptTokensDetails]
promptTokensDetails: List[PromptTokensDetails]
cacheTokensDetails: List[PromptTokensDetails]
thoughtsTokenCount: int

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

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

@ -197,6 +197,9 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
cache_read_input_token_cost_above_272k_tokens: Optional[float]
cache_read_input_token_cost_above_272k_tokens_priority: Optional[float]
cache_read_input_token_cost_above_512k_tokens: Optional[float]
# Smallest prefix this model will actually cache, whatever caching mechanism its provider uses.
# Absent means the provider-agnostic default applies; see MINIMUM_PROMPT_CACHE_TOKEN_COUNT.
prompt_cache_min_tokens: Optional[int]
input_cost_per_character: Optional[float] # only for vertex ai models
input_cost_per_audio_token: Optional[float]
input_cost_per_token_above_128k_tokens: Optional[float] # only for vertex ai models
@ -263,6 +266,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
"audio_transcription",
"responses",
"ocr",
"realtime",
]
]
tpm: Optional[int]
@ -325,6 +329,7 @@ class CallTypes(str, Enum):
cancel_batch = "cancel_batch"
pass_through = "pass_through_endpoint"
anthropic_messages = "anthropic_messages"
aanthropic_messages = "aanthropic_messages"
get_assistants = "get_assistants"
aget_assistants = "aget_assistants"
create_assistants = "create_assistants"
@ -398,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
#########################################################
@ -493,6 +503,7 @@ CallTypesLiteral = Literal[
"pass_through_endpoint",
"allm_passthrough_route",
"anthropic_messages",
"aanthropic_messages",
"aretrieve_batch",
"retrieve_batch",
"generate_content",
@ -1474,6 +1485,9 @@ class PromptTokensDetailsWrapper(
web_search_requests: Optional[int] = None
"""Number of web search requests made by the tool call. Used for Anthropic to calculate web search cost."""
tool_use_tokens: Optional[int] = None
"""Prompt tokens consumed by server-side tool use (e.g. Gemini grounding via googleSearch)."""
character_count: Optional[int] = None
"""Character count sent to the model. Used for Vertex AI multimodal embeddings."""
@ -1504,6 +1518,8 @@ class PromptTokensDetailsWrapper(
del self.audio_length_seconds
if self.web_search_requests is None:
del self.web_search_requests
if self.tool_use_tokens is None:
del self.tool_use_tokens
if self.cache_creation_tokens is None:
del self.cache_creation_tokens
if self.cache_creation_token_details is None:

View file

@ -73,7 +73,8 @@ from litellm.constants import (
JITTER,
MAX_RETRY_DELAY,
MAX_TOKEN_TRIMMING_ATTEMPTS,
MINIMUM_PROMPT_CACHE_TOKEN_COUNT,
DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT,
MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE,
OPENAI_EMBEDDING_PARAMS,
TOOL_CHOICE_OBJECT_TOKEN_COUNT,
)
@ -3197,6 +3198,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",
@ -5402,6 +5409,7 @@ def _get_model_info_helper(
"cache_creation_input_token_cost_above_200k_tokens", None
),
cache_read_input_token_cost=_model_info.get("cache_read_input_token_cost", None),
prompt_cache_min_tokens=_model_info.get("prompt_cache_min_tokens", None),
cache_read_input_token_cost_above_200k_tokens=_model_info.get(
"cache_read_input_token_cost_above_200k_tokens", None
),
@ -9039,16 +9047,46 @@ def should_use_cohere_v1_client(api_base: Optional[str], present_version_params:
return api_base.endswith("/v1/rerank") or (uses_v1_params and not api_base.endswith("/v2/rerank"))
def get_prompt_cache_min_tokens(model: str) -> int:
"""
Returns the smallest prefix `model` will actually cache.
Resolution order is an explicitly configured `MINIMUM_PROMPT_CACHE_TOKEN_COUNT`, then the
model's `prompt_cache_min_tokens` in the cost map, then the provider-agnostic default. The
cost map is the source of truth because the real minimum is per-model and per-platform:
Anthropic's ranges from 512 to 4096 and moves in both directions across releases, and the
same model can differ by platform.
Never raises. An unresolvable model falls back to the default rather than propagating, so a
caller cannot mistake "no entry for this model" for "this prompt is not cacheable".
"""
if MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE is not None:
return MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE
try:
min_tokens = get_model_info(model=model).get("prompt_cache_min_tokens")
except Exception:
return DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT
if min_tokens is None:
return DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT
return min_tokens
def is_prompt_caching_valid_prompt(
model: str,
messages: Optional[List[AllMessageValues]],
tools: Optional[List[ChatCompletionToolParam]] = None,
custom_llm_provider: Optional[str] = None,
min_token_count: int | None = None,
) -> bool:
"""
Returns true if the prompt is valid for prompt caching.
OpenAI + Anthropic providers have a minimum token count of 1024 for prompt caching.
The minimum cacheable prefix is per-model, so it is resolved from `model` unless the caller
passes `min_token_count`. Callers that only hold a model-group alias (the router's deployment
checks) must resolve the threshold themselves and pass it, because an alias resolves to
nothing here and would silently fall back to the default.
OpenAI's minimum is a flat 1024 across models, which the default already covers.
"""
try:
if messages is None and tools is None:
@ -9061,7 +9099,9 @@ def is_prompt_caching_valid_prompt(
model=model,
use_default_image_token_count=True,
)
return token_count >= MINIMUM_PROMPT_CACHE_TOKEN_COUNT
if min_token_count is None:
min_token_count = get_prompt_cache_min_tokens(model=model)
return token_count >= min_token_count
except Exception as e:
verbose_logger.error(f"Error in is_prompt_caching_valid_prompt: {e}")
return False

File diff suppressed because it is too large Load diff

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,11 +13,12 @@ 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)
- `gateway/` - proxy configuration only (`litellm-config.yml`); no tests
- `claude_code/` - the Claude Code compatibility matrix: drives the real `claude` CLI (and HTTP probes) against a proxy for each feature x provider cell, reporting tagged-union outcomes via the `compat_result` fixture; ships its own driver/builder/publisher and does not use the shared transport harness
- `claude_code/` - the Claude Code compatibility matrix: drives the real `claude` CLI (and HTTP probes) against a proxy for each feature x provider cell, reporting tagged-union outcomes via the `compat_result` fixture; ships its own driver/builder/publisher plus `_*_unit_tests/` trees, and does not use the shared transport harness
## Lay the pattern down in a class
@ -51,9 +52,9 @@ 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). `tests/e2e/` is for live proxy suites only; do not put unit tests here. 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
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
## Typing
@ -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

@ -180,3 +180,43 @@ def matches_id_shape(shape: IdShape, id_str: str) -> bool:
if shape == "model_encoded":
return is_model_encoded_id(id_str)
return not is_managed_id(id_str) and not is_model_encoded_id(id_str)
def coverage_cells_for_lifecycle(cap: Capability) -> tuple[str, ...]:
"""Registry cell ids that the parametrized lifecycle test covers for one capability.
OpenAI has per-scenario cells plus granular create/retrieve/cancel/list/file
cells. Other providers have one basic cell each. File-upload cells for the
batch-backing path are included when the lifecycle uploads for that provider.
"""
match cap.provider:
case "openai":
cells = (
f"llm.batches.openai_{cap.scenario}.basic.nonstream.works",
"llm.batches.openai.create.nonstream.works",
"llm.batches.openai.retrieve.nonstream.works",
"llm.batches.openai.file_lifecycle.nonstream.works",
"llm.files.openai.upload.nonstream.works",
)
if cap.can_cancel:
cells = (*cells, "llm.batches.openai.cancel.nonstream.works")
if cap.can_list:
cells = (*cells, "llm.batches.openai.list.nonstream.works")
return cells
case "azure":
return (
"llm.batches.azure_openai.basic.nonstream.works",
"llm.files.azure_openai.upload.nonstream.works",
)
case "vertex_ai":
return (
"llm.batches.vertex.basic.nonstream.works",
"llm.files.vertex.upload.nonstream.works",
)
case "bedrock":
return (
"llm.batches.bedrock.basic.nonstream.works",
"llm.files.bedrock.upload.nonstream.works",
)
case _:
return ()

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

@ -37,6 +37,7 @@ from capabilities import (
CAPABILITIES,
FILE_ID_SHAPE,
Capability,
coverage_cells_for_lifecycle,
matches_id_shape,
raw_id_matches_provider,
)
@ -168,7 +169,17 @@ def assert_batch_object(batch: BatchObject) -> None:
), "batch.created_at missing"
@pytest.mark.parametrize("cap", CAPABILITIES, ids=[c.id for c in CAPABILITIES])
@pytest.mark.parametrize(
"cap",
[
pytest.param(
cap,
id=cap.id,
marks=pytest.mark.covers(*coverage_cells_for_lifecycle(cap)),
)
for cap in CAPABILITIES
],
)
def test_batch_lifecycle(
cap: Capability,
client: BatchClient,
@ -266,6 +277,7 @@ def test_batch_lifecycle(
assert match.object == "batch"
@pytest.mark.covers("llm.batches.openai.key_model_access_denied.nonstream.works")
def test_batch_key_model_access_denied(
client: BatchClient, resources: ResourceManager, batch_deployments: None
) -> None:
@ -301,6 +313,10 @@ def test_batch_key_model_access_denied(
), f"restricted key created a batch for a disallowed model (status {denied_create.status_code})"
@pytest.mark.covers(
"llm.files.openai.upload.nonstream.works",
"llm.files.openai.delete.nonstream.works",
)
def test_file_upload_and_delete_outputs(
client: BatchClient, resources: ResourceManager, batch_deployments: None
) -> None:

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

@ -0,0 +1,78 @@
"""Shared plumbing for the GPT-5.6 (Sol / Terra / Luna) provider columns.
OpenAI shipped GPT-5.6 as a three-tier family on 2026-07-09 Sol
(flagship), Terra (balanced), Luna (fast) and Claude Code can drive
all three through a LiteLLM proxy that translates the Anthropic
Messages API to each provider's native shape. Four provider columns
cover "OpenAI plus the big three clouds":
openai OpenAI API (openai/gpt-5.6-*)
azure_openai Azure OpenAI (azure/gpt-5.6-*)
bedrock_mantle AWS Bedrock, Mantle (bedrock_mantle/openai.gpt-5.6-*,
Responses API)
vertex_ai_gpt GCP Vertex AI not_applicable Vertex does
not offer the closed-weight
GPT-5.6 family; Model Garden
carries only the open-weight
gpt-oss MaaS models
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
a static `not_applicable` and never touch the network.
"""
from __future__ import annotations
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 "
"(Sol / Terra / Luna); Model Garden carries only the open-weight "
"gpt-oss MaaS models. Convert this column's cells to live tests if "
"Google adds the GPT-5.6 models."
)
def skip_unless_mantle_cells_enabled() -> None:
"""Skip the calling test unless `COMPAT_MANTLE_CELLS` opts the
Bedrock Mantle cells in.
A skipped cell is recorded as `not_tested` in the published matrix
(see the skip handling in `tests/e2e/claude_code/conftest.py`),
which is the honest state while the AWS account has no Mantle
access to the GPT-5.6 models yet.
"""
if os.environ.get(MANTLE_CELLS_ENV, "").strip().lower() in {"1", "true", "yes"}:
return
pytest.skip(
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

@ -0,0 +1,44 @@
"""basic_messaging_non_streaming x Azure OpenAI (GPT-5.6).
Drive the real `claude` CLI in headless mode against a running LiteLLM
proxy that routes Anthropic Messages requests to Azure OpenAI
deployments of the GPT-5.6 family (Sol, Terra, Luna), and report the
outcome via `compat_result`.
Azure OpenAI serves the same chat-completions wire shape as
openai.com behind per-resource deployments; LiteLLM's `azure/gpt-*`
route handles the deployment addressing while reusing the OpenAI
translation, so this cell catches Azure-specific regressions
(auth headers, api-version pinning, deployment routing) that the
`openai` column cannot.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/basic_messaging_non_streaming/test_azure_openai.py
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ ^^^^^^^^^^^^
feature_id provider
Every GPT cell exercises the three GPT-5.6 tiers; the cell only goes
green if all three pass.
"""
from __future__ import annotations
from claude_code._basic_messaging import run_basic_messaging_cell
AZURE_OPENAI_MODELS = [
"gpt-5-6-sol-azure-openai",
"gpt-5-6-terra-azure-openai",
"gpt-5-6-luna-azure-openai",
]
def test_basic_messaging_non_streaming_azure_openai(compat_result):
"""Drive the `claude` CLI against the LiteLLM proxy and assert a
non-empty reply from each GPT-5.6 tier."""
run_basic_messaging_cell(
compat_result=compat_result,
models=AZURE_OPENAI_MODELS,
prompt="Reply with the single word 'pong' and nothing else.",
)

View file

@ -0,0 +1,46 @@
"""basic_messaging_non_streaming x AWS Bedrock Mantle (GPT-5.6).
Drive the real `claude` CLI in headless mode against a running LiteLLM
proxy that routes Anthropic Messages requests to OpenAI's GPT-5.6
family (Sol, Terra, Luna) hosted on AWS Bedrock, and report the
outcome via `compat_result`.
Bedrock exposes the GPT-5.6 models through the Mantle endpoint, which
speaks the OpenAI Responses API rather than Converse/Invoke; LiteLLM's
`bedrock_mantle/openai.gpt-*` route signs the request with SigV4 and
translates Anthropic Messages to Responses, so this cell exercises a
translation path no other column covers.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_mantle.py
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ ^^^^^^^^^^^^^^
feature_id provider
Every GPT cell exercises the three GPT-5.6 tiers; the cell only goes
green if all three pass. Mantle cells are opt-in via
COMPAT_MANTLE_CELLS=1 (see `claude_code._gpt_cells`).
"""
from __future__ import annotations
from claude_code._basic_messaging import run_basic_messaging_cell
from claude_code._gpt_cells import skip_unless_mantle_cells_enabled
BEDROCK_MANTLE_MODELS = [
"gpt-5-6-sol-bedrock-mantle",
"gpt-5-6-terra-bedrock-mantle",
"gpt-5-6-luna-bedrock-mantle",
]
def test_basic_messaging_non_streaming_bedrock_mantle(compat_result):
"""Drive the `claude` CLI against the LiteLLM proxy and assert a
non-empty reply from each GPT-5.6 tier."""
skip_unless_mantle_cells_enabled()
run_basic_messaging_cell(
compat_result=compat_result,
models=BEDROCK_MANTLE_MODELS,
prompt="Reply with the single word 'pong' and nothing else.",
)

View file

@ -0,0 +1,43 @@
"""basic_messaging_non_streaming x OpenAI (GPT-5.6).
Drive the real `claude` CLI in headless mode against a running LiteLLM
proxy that routes Anthropic Messages requests to OpenAI's GPT-5.6
family (Sol, Terra, Luna), and report the outcome via `compat_result`.
Claude Code only speaks the Anthropic Messages API; LiteLLM's
`openai/gpt-*` route translates the request to OpenAI chat completions
and maps the response back, so this cell exercises the full
cross-provider translation layer in both directions.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/basic_messaging_non_streaming/test_openai.py
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ ^^^^^^
feature_id provider
Every GPT cell exercises the three GPT-5.6 tiers; the cell only goes
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",
"gpt-5-6-terra-openai",
"gpt-5-6-luna-openai",
]
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,
prompt="Reply with the single word 'pong' and nothing else.",
)

View file

@ -0,0 +1,29 @@
"""basic_messaging_non_streaming x Vertex AI (GPT-5.6) — not applicable.
GCP is the only one of the big-three clouds without OpenAI's
closed-weight GPT-5.6 family (Sol / Terra / Luna); Vertex AI Model
Garden carries only the open-weight gpt-oss MaaS models. The cell
reports `not_applicable` so the published matrix documents the gap
explicitly instead of leaving a `not_tested` hole.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai_gpt.py
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ ^^^^^^^^^^^^^
feature_id provider
"""
from __future__ import annotations
from claude_code._gpt_cells import VERTEX_AI_GPT_NOT_APPLICABLE_REASON
def test_basic_messaging_non_streaming_vertex_ai_gpt(compat_result):
"""Record the static not_applicable outcome for this cell."""
compat_result.set(
{
"status": "not_applicable",
"reason": VERTEX_AI_GPT_NOT_APPLICABLE_REASON,
}
)

View file

@ -0,0 +1,44 @@
"""basic_messaging_streaming x Azure OpenAI (GPT-5.6).
Drive the real `claude` CLI in headless `--output-format stream-json`
mode against a running LiteLLM proxy that routes Anthropic Messages
requests to Azure OpenAI deployments of the GPT-5.6 family (Sol,
Terra, Luna), and report the outcome via `compat_result`.
Azure OpenAI streams the same chat-completions SSE shape as
openai.com; LiteLLM re-emits it as Anthropic stream events, and the
`verify_streaming=True` assertion (via `--include-partial-messages`)
proves the events arrived incrementally rather than as one buffered
response.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/basic_messaging_streaming/test_azure_openai.py
^^^^^^^^^^^^^^^^^^^^^^^^^ ^^^^^^^^^^^^
feature_id provider
Every GPT cell exercises the three GPT-5.6 tiers; the cell only goes
green if all three pass.
"""
from __future__ import annotations
from claude_code._basic_messaging import run_basic_messaging_cell
AZURE_OPENAI_MODELS = [
"gpt-5-6-sol-azure-openai",
"gpt-5-6-terra-azure-openai",
"gpt-5-6-luna-azure-openai",
]
def test_basic_messaging_streaming_azure_openai(compat_result):
"""Drive the `claude` CLI against the LiteLLM proxy and assert a
non-empty streamed reply from each GPT-5.6 tier."""
run_basic_messaging_cell(
compat_result=compat_result,
models=AZURE_OPENAI_MODELS,
prompt="Count from 1 to 5, one number per line.",
verify_streaming=True,
)

View file

@ -0,0 +1,47 @@
"""basic_messaging_streaming x AWS Bedrock Mantle (GPT-5.6).
Drive the real `claude` CLI in headless `--output-format stream-json`
mode against a running LiteLLM proxy that routes Anthropic Messages
requests to OpenAI's GPT-5.6 family (Sol, Terra, Luna) on AWS
Bedrock's Mantle endpoint, and report the outcome via `compat_result`.
Mantle streams OpenAI Responses API events over SigV4-signed SSE;
LiteLLM re-emits them as Anthropic stream events, and the
`verify_streaming=True` assertion (via `--include-partial-messages`)
proves the events arrived incrementally rather than as one buffered
response.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_mantle.py
^^^^^^^^^^^^^^^^^^^^^^^^^ ^^^^^^^^^^^^^^
feature_id provider
Every GPT cell exercises the three GPT-5.6 tiers; the cell only goes
green if all three pass. Mantle cells are opt-in via
COMPAT_MANTLE_CELLS=1 (see `claude_code._gpt_cells`).
"""
from __future__ import annotations
from claude_code._basic_messaging import run_basic_messaging_cell
from claude_code._gpt_cells import skip_unless_mantle_cells_enabled
BEDROCK_MANTLE_MODELS = [
"gpt-5-6-sol-bedrock-mantle",
"gpt-5-6-terra-bedrock-mantle",
"gpt-5-6-luna-bedrock-mantle",
]
def test_basic_messaging_streaming_bedrock_mantle(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_mantle_cells_enabled()
run_basic_messaging_cell(
compat_result=compat_result,
models=BEDROCK_MANTLE_MODELS,
prompt="Count from 1 to 5, one number per line.",
verify_streaming=True,
)

View file

@ -0,0 +1,46 @@
"""basic_messaging_streaming x OpenAI (GPT-5.6).
Drive the real `claude` CLI in headless `--output-format stream-json`
mode against a running LiteLLM proxy that routes Anthropic Messages
requests to OpenAI's GPT-5.6 family (Sol, Terra, Luna), and report the
outcome via `compat_result`.
LiteLLM translates OpenAI's chat-completions SSE chunks into Anthropic
`message_start` / `content_block_delta` / `message_stop` events on the
fly; the `verify_streaming=True` assertion (via
`--include-partial-messages`) proves the proxy re-emitted incremental
events instead of buffering the upstream stream into one response.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/basic_messaging_streaming/test_openai.py
^^^^^^^^^^^^^^^^^^^^^^^^^ ^^^^^^
feature_id provider
Every GPT cell exercises the three GPT-5.6 tiers; the cell only goes
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",
"gpt-5-6-terra-openai",
"gpt-5-6-luna-openai",
]
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,
prompt="Count from 1 to 5, one number per line.",
verify_streaming=True,
)

View file

@ -0,0 +1,29 @@
"""basic_messaging_streaming x Vertex AI (GPT-5.6) — not applicable.
GCP is the only one of the big-three clouds without OpenAI's
closed-weight GPT-5.6 family (Sol / Terra / Luna); Vertex AI Model
Garden carries only the open-weight gpt-oss MaaS models. The cell
reports `not_applicable` so the published matrix documents the gap
explicitly instead of leaving a `not_tested` hole.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai_gpt.py
^^^^^^^^^^^^^^^^^^^^^^^^^ ^^^^^^^^^^^^^
feature_id provider
"""
from __future__ import annotations
from claude_code._gpt_cells import VERTEX_AI_GPT_NOT_APPLICABLE_REASON
def test_basic_messaging_streaming_vertex_ai_gpt(compat_result):
"""Record the static not_applicable outcome for this cell."""
compat_result.set(
{
"status": "not_applicable",
"reason": VERTEX_AI_GPT_NOT_APPLICABLE_REASON,
}
)

View file

@ -53,6 +53,7 @@ VERTEX_AI_MODELS = [
]
@pytest.mark.skip(reason="stage red: Vertex returns not supported for token counting for Claude aliases")
@pytest.mark.covers("llm.messages.vertex.count_tokens.nonstream.works")
def test_count_tokens_vertex_ai(compat_result):
"""Probe `/v1/messages/count_tokens` for each Vertex AI tier and

View file

@ -153,6 +153,7 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str:
return preamble + "".join(pad_lines) + closing
@pytest.mark.skip(reason="stage red: 1M long_context not green on stage Anthropic path yet (200k sonnet / model alias)")
@pytest.mark.covers("llm.messages.anthropic.long_context_1m.nonstream.works")
def test_long_context_1m_anthropic(compat_result):
"""Drive the `claude` CLI with a ~210k-token prompt and the

View file

@ -153,6 +153,7 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str:
return preamble + "".join(pad_lines) + closing
@pytest.mark.skip(reason="stage red: 1M long_context not green on stage Azure Foundry deployments yet")
@pytest.mark.covers("llm.messages.azure_foundry.long_context_1m.nonstream.works")
def test_long_context_1m_azure(compat_result):
"""Drive the `claude` CLI (Azure (Microsoft Foundry)) with a ~210k-token prompt and the

View file

@ -153,6 +153,7 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str:
return preamble + "".join(pad_lines) + closing
@pytest.mark.skip(reason="stage red: 1M long_context not green on stage Bedrock Converse deployments yet")
@pytest.mark.covers("llm.messages.bedrock_converse.long_context_1m.nonstream.works")
def test_long_context_1m_bedrock_converse(compat_result):
"""Drive the `claude` CLI (Bedrock (Converse)) with a ~210k-token prompt and the

View file

@ -153,6 +153,7 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str:
return preamble + "".join(pad_lines) + closing
@pytest.mark.skip(reason="stage red: 1M long_context not green on stage Bedrock Invoke deployments yet")
@pytest.mark.covers("llm.messages.bedrock_invoke.long_context_1m.nonstream.works")
def test_long_context_1m_bedrock_invoke(compat_result):
"""Drive the `claude` CLI (Bedrock (Invoke)) with a ~210k-token prompt and the

View file

@ -153,6 +153,7 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str:
return preamble + "".join(pad_lines) + closing
@pytest.mark.skip(reason="stage red: 1M long_context not green on stage Vertex deployments yet")
@pytest.mark.covers("llm.messages.vertex.long_context_1m.nonstream.works")
def test_long_context_1m_vertex_ai(compat_result):
"""Drive the `claude` CLI (Vertex AI) with a ~210k-token prompt and the

View file

@ -12,13 +12,24 @@
schema_version: "1"
# Provider column order in the rendered matrix.
# Provider column order in the rendered matrix. The first five are
# the v0 Claude columns; the GPT-5.6 (Sol / Terra / Luna) columns
# added 2026-07 follow them. `vertex_ai_gpt` is a static
# not_applicable column: GCP does not offer the closed-weight GPT-5.6
# family (Model Garden carries only the open-weight gpt-oss MaaS
# models), and the column documents that gap explicitly. GPT columns
# currently back the two basic_messaging rows plus tool_use and
# tool_use_streaming; other rows render not_tested for them.
providers:
- anthropic
- bedrock_invoke
- bedrock_converse
- vertex_ai
- azure
- openai
- azure_openai
- bedrock_mantle
- vertex_ai_gpt
# Feature row order.
features:

View file

@ -40,6 +40,8 @@ of bug the row exists to surface.
from __future__ import annotations
import pytest
from claude_code._passthrough import foundry_extra_env, run_passthrough_cell
AZURE_MODELS = [
@ -49,6 +51,7 @@ AZURE_MODELS = [
]
@pytest.mark.skip(reason="stage red: /azure passthrough drops client headers (e.g. anthropic-version); product gap")
def test_passthrough_azure(compat_result):
"""Drive the `claude` CLI through `{proxy}/azure` and assert a reply."""
run_passthrough_cell(

View file

@ -88,6 +88,10 @@ def _build_minimal_pdf(marker: str) -> bytes:
return bytes(out)
@pytest.mark.skip(
reason="product bug LIT-4523: Bedrock Converse requires a text block with document; "
"re-enable when document-only content is handled"
)
@pytest.mark.covers("llm.messages.bedrock_converse.pdf_input.nonstream.works")
def test_pdf_input_bedrock_converse(compat_result, tmp_path):
base_url, api_key = require_proxy(compat_result)

View file

@ -24,6 +24,9 @@ edits:
LITELLM_COMPAT_RATE_VERTEX_AI (req/s, default 5.0)
LITELLM_COMPAT_RATE_BEDROCK_CONVERSE (req/s, default 5.0)
LITELLM_COMPAT_RATE_BEDROCK_INVOKE (req/s, default 5.0)
LITELLM_COMPAT_RATE_OPENAI (req/s, default 5.0)
LITELLM_COMPAT_RATE_AZURE_OPENAI (req/s, default 5.0)
LITELLM_COMPAT_RATE_BEDROCK_MANTLE (req/s, default 5.0)
LITELLM_COMPAT_RATE_BURST (per-bucket burst override;
default = rate)
LITELLM_COMPAT_RATE_STATE_DIR (state file directory;
@ -36,7 +39,10 @@ network.
The provider id is inferred from the model id by `infer_provider`,
mirroring the matrix's column layout (`anthropic`, `azure`,
`vertex_ai`, `bedrock_converse`, `bedrock_invoke`).
`vertex_ai`, `bedrock_converse`, `bedrock_invoke`, `openai`,
`azure_openai`, `bedrock_mantle`). The `vertex_ai_gpt` matrix column
has no bucket: its cells are static not_applicable stubs that never
reach the network.
"""
from __future__ import annotations
@ -62,6 +68,9 @@ PROVIDER_AZURE = "azure"
PROVIDER_VERTEX_AI = "vertex_ai"
PROVIDER_BEDROCK_CONVERSE = "bedrock_converse"
PROVIDER_BEDROCK_INVOKE = "bedrock_invoke"
PROVIDER_OPENAI = "openai"
PROVIDER_AZURE_OPENAI = "azure_openai"
PROVIDER_BEDROCK_MANTLE = "bedrock_mantle"
ALL_PROVIDERS = (
PROVIDER_ANTHROPIC,
@ -69,6 +78,9 @@ ALL_PROVIDERS = (
PROVIDER_VERTEX_AI,
PROVIDER_BEDROCK_CONVERSE,
PROVIDER_BEDROCK_INVOKE,
PROVIDER_OPENAI,
PROVIDER_AZURE_OPENAI,
PROVIDER_BEDROCK_MANTLE,
)
DEFAULT_RATE = 5.0 # req/s per provider, conservative starting point
@ -83,13 +95,21 @@ def infer_provider(model: str) -> str:
The matrix column layout is fixed; aliases registered in the proxy
encode the provider via a suffix (`-bedrock-converse`,
`-bedrock-invoke`, `-azure`, `-vertex`) or its absence (Anthropic).
Order matters: the bedrock suffixes both contain `bedrock`, so we
test the more-specific ones first.
`-bedrock-invoke`, `-azure`, `-vertex`, `-openai`, `-azure-openai`,
`-bedrock-mantle`) or its absence (Anthropic). Order matters:
`-azure-openai` also ends with `-openai`, and the bedrock suffixes
all contain `bedrock`, so the more-specific suffixes are tested
first.
"""
if not model:
raise ValueError("model must be a non-empty string")
lower = model.lower()
if lower.endswith("-azure-openai"):
return PROVIDER_AZURE_OPENAI
if lower.endswith("-openai"):
return PROVIDER_OPENAI
if lower.endswith("-bedrock-mantle"):
return PROVIDER_BEDROCK_MANTLE
if lower.endswith("-bedrock-converse"):
return PROVIDER_BEDROCK_CONVERSE
if lower.endswith("-bedrock-invoke"):

View file

@ -21,8 +21,16 @@
# LITELLM_COMPAT_RATE_VERTEX_AI
# LITELLM_COMPAT_RATE_BEDROCK_CONVERSE
# LITELLM_COMPAT_RATE_BEDROCK_INVOKE
# LITELLM_COMPAT_RATE_OPENAI
# LITELLM_COMPAT_RATE_AZURE_OPENAI
# LITELLM_COMPAT_RATE_BEDROCK_MANTLE
# LITELLM_COMPAT_RATE_BURST override per-bucket burst
#
# Optional env (GPT-5.6 columns):
# COMPAT_MANTLE_CELLS=1 opt the Bedrock Mantle GPT-5.6
# cells in; without it they skip
# and publish as not_tested
#
# Optional env (parallelism):
# COMPAT_XDIST_WORKERS passed to `pytest -n` (default: auto)
#
@ -57,7 +65,7 @@ results_path="${COMPAT_RESULTS_PATH:-compat-results.json}"
summary_path="${COMPAT_RATE_LIMIT_SUMMARY_PATH:-compat-rate-limit-summary.json}"
echo "[run_compat] rates:"
for provider in ANTHROPIC AZURE VERTEX_AI BEDROCK_CONVERSE BEDROCK_INVOKE; do
for provider in ANTHROPIC AZURE VERTEX_AI BEDROCK_CONVERSE BEDROCK_INVOKE OPENAI AZURE_OPENAI BEDROCK_MANTLE; do
var="LITELLM_COMPAT_RATE_${provider}"
echo " ${provider}=${!var:-default(5/s)}"
done

View file

@ -14,6 +14,14 @@
# - claude-{tier}-bedrock-converse → Bedrock Converse API
# - claude-{tier}-vertex → GCP Vertex AI
# - claude-{tier}-azure → Microsoft Foundry (Anthropic deployments)
# - gpt-5-6-{tier}-openai → OpenAI API
# - gpt-5-6-{tier}-azure-openai → Azure OpenAI deployments
# - gpt-5-6-{tier}-bedrock-mantle → Bedrock Mantle (Responses API)
#
# GPT-5.6 tiers are sol / terra / luna. There are no GPT aliases for
# GCP: Vertex AI does not offer the closed-weight GPT-5.6 family, so
# the matrix's `vertex_ai_gpt` column reports not_applicable without
# ever reaching the proxy.
model_list:
# ---- Anthropic ----
@ -122,6 +130,54 @@ model_list:
extra_headers:
anthropic-beta: "context-1m-2025-08-07"
# ---- OpenAI (GPT-5.6) ----
- model_name: gpt-5-6-sol-openai
litellm_params:
model: openai/gpt-5.6-sol
api_key: os.environ/OPENAI_API_KEY
- model_name: gpt-5-6-terra-openai
litellm_params:
model: openai/gpt-5.6-terra
api_key: os.environ/OPENAI_API_KEY
- model_name: gpt-5-6-luna-openai
litellm_params:
model: openai/gpt-5.6-luna
api_key: os.environ/OPENAI_API_KEY
# ---- Azure OpenAI (GPT-5.6) ----
- model_name: gpt-5-6-sol-azure-openai
litellm_params:
model: azure/gpt-5.6-sol
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
- model_name: gpt-5-6-terra-azure-openai
litellm_params:
model: azure/gpt-5.6-terra
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
- model_name: gpt-5-6-luna-azure-openai
litellm_params:
model: azure/gpt-5.6-luna
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
# ---- Bedrock Mantle (GPT-5.6, Responses API) ----
# Sol is only served from us-east-1 / us-east-2 as of 2026-07;
# Terra and Luna additionally have us-west-2. One region keeps the
# column comparable across tiers.
- model_name: gpt-5-6-sol-bedrock-mantle
litellm_params:
model: bedrock_mantle/openai.gpt-5.6-sol
aws_region_name: us-east-1
- model_name: gpt-5-6-terra-bedrock-mantle
litellm_params:
model: bedrock_mantle/openai.gpt-5.6-terra
aws_region_name: us-east-1
- model_name: gpt-5-6-luna-bedrock-mantle
litellm_params:
model: bedrock_mantle/openai.gpt-5.6-luna
aws_region_name: us-east-1
general_settings:
# Claude Code sends provider-specific headers (e.g. anthropic-beta) we
# want to forward verbatim to the upstream so the wire-shape under

View file

@ -54,6 +54,10 @@ def _has_thinking_block(events: Sequence[Mapping[str, Any]]) -> bool:
return False
@pytest.mark.skip(
reason="product bug LIT-4524: Bedrock Converse streaming Content block is not a text block; "
"re-enable when empty/mismatched content_block_delta is fixed"
)
@pytest.mark.covers("llm.messages.bedrock_converse.thinking.nonstream.works")
def test_thinking_bedrock_converse(compat_result):
"""Drive the `claude` CLI against the LiteLLM proxy with thinking

View file

@ -59,6 +59,7 @@ AZURE_MODELS = [
]
@pytest.mark.skip(reason="stage red: Azure Foundry tool_search_server not supported in workspace for probed models")
@pytest.mark.covers("llm.messages.azure_foundry.tool_search.nonstream.works")
def test_tool_search_azure(compat_result):
"""Probe `/v1/messages` with a `tool_search_tool_regex_20251119`

View file

@ -59,6 +59,10 @@ BEDROCK_INVOKE_MODELS = [
]
@pytest.mark.skip(
reason="product bug LIT-4522: Bedrock Invoke /v1/messages does not normalize "
"tool_search_tool_regex_20251119; re-enable when messages path matches chat path"
)
@pytest.mark.covers("llm.messages.bedrock_invoke.tool_search.nonstream.works")
def test_tool_search_bedrock_invoke(compat_result):
"""Probe `/v1/messages` with a `tool_search_tool_regex_20251119`

View file

@ -59,6 +59,7 @@ VERTEX_AI_MODELS = [
]
@pytest.mark.skip(reason="stage red: Vertex rejects tool_search when deployment extra_headers inject context-1m beta; product/config")
@pytest.mark.covers("llm.messages.vertex.tool_search.nonstream.works")
def test_tool_search_vertex_ai(compat_result):
"""Probe `/v1/messages` with a `tool_search_tool_regex_20251119`

View file

@ -0,0 +1,109 @@
"""tool_use x Azure OpenAI (GPT-5.6).
Drive the real `claude` CLI against a running LiteLLM proxy that
routes Anthropic Messages requests to Azure OpenAI deployments of the
GPT-5.6 family (Sol, Terra, Luna), ask the model to invoke a built-in
tool (`Bash`), and assert that a `tool_use` content block came back
over the wire.
Azure OpenAI serves the same function-calling wire shape as
openai.com behind per-resource deployments; LiteLLM's `azure/gpt-*`
route reuses the OpenAI tool translation on top of Azure's deployment
addressing and auth.
Bash is restricted to the exact command `echo pong` plus
`--permission-mode dontAsk`; see `tool_use/test_anthropic.py` for the
security rationale.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/tool_use/test_azure_openai.py
^^^^^^^^ ^^^^^^^^^^^^
feature_id provider
"""
from __future__ import annotations
from typing import Any, Mapping, Sequence
import pytest
from claude_code._env import require_proxy
from claude_code.cli_driver import (
ClaudeCLIError,
failure_diagnostic,
run_claude_models_parallel,
)
AZURE_OPENAI_MODELS = [
"gpt-5-6-sol-azure-openai",
"gpt-5-6-terra-azure-openai",
"gpt-5-6-luna-azure-openai",
]
TOOL_USE_PROMPT = (
"Use the Bash tool to run the command `echo pong` and report what it printed."
)
TOOL_USE_ARGS = [
"--allowed-tools",
"Bash(echo pong)",
"--permission-mode",
"dontAsk",
]
def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool:
for event in events:
if event.get("type") != "assistant":
continue
message = event.get("message") or {}
content = message.get("content")
if not isinstance(content, list):
continue
for block in content:
if isinstance(block, dict) and block.get("type") == "tool_use":
return True
return False
def test_tool_use_azure_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."""
proxy = require_proxy(compat_result)
outcomes = run_claude_models_parallel(
models=AZURE_OPENAI_MODELS,
prompt=TOOL_USE_PROMPT,
base_url=proxy.base_url,
api_key=proxy.api_key,
extra_args=TOOL_USE_ARGS,
)
failures = []
for model in AZURE_OPENAI_MODELS:
outcome = outcomes[model]
if isinstance(outcome, ClaudeCLIError):
error = f"[{model}] {outcome}"
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
if outcome.exit_code != 0:
error = f"[{model}] claude CLI failed: {failure_diagnostic(outcome)}"
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
if not _has_tool_use_event(outcome.events):
error = (
f"[{model}] no tool_use content block observed in stream-json events"
)
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
compat_result.add({"status": "pass"})
if failures:
pytest.fail("; ".join(failures), pytrace=False)

View file

@ -0,0 +1,115 @@
"""tool_use x AWS Bedrock Mantle (GPT-5.6).
Drive the real `claude` CLI against a running LiteLLM proxy that
routes Anthropic Messages requests to OpenAI's GPT-5.6 family (Sol,
Terra, Luna) on AWS Bedrock's Mantle endpoint, ask the model to invoke
a built-in tool (`Bash`), and assert that a `tool_use` content block
came back over the wire.
Mantle speaks the OpenAI Responses API, whose tool declarations and
`function_call` outputs differ from both Anthropic Messages and
chat completions; LiteLLM's `bedrock_mantle/openai.gpt-*` route
translates Anthropic `tools` into Responses tool declarations and maps
the emitted function calls back to `tool_use` blocks.
Bash is restricted to the exact command `echo pong` plus
`--permission-mode dontAsk`; see `tool_use/test_anthropic.py` for the
security rationale.
Mantle cells are opt-in via COMPAT_MANTLE_CELLS=1 (see
`claude_code._gpt_cells`).
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/tool_use/test_bedrock_mantle.py
^^^^^^^^ ^^^^^^^^^^^^^^
feature_id provider
"""
from __future__ import annotations
from typing import Any, Mapping, Sequence
import pytest
from claude_code._env import require_proxy
from claude_code._gpt_cells import skip_unless_mantle_cells_enabled
from claude_code.cli_driver import (
ClaudeCLIError,
failure_diagnostic,
run_claude_models_parallel,
)
BEDROCK_MANTLE_MODELS = [
"gpt-5-6-sol-bedrock-mantle",
"gpt-5-6-terra-bedrock-mantle",
"gpt-5-6-luna-bedrock-mantle",
]
TOOL_USE_PROMPT = (
"Use the Bash tool to run the command `echo pong` and report what it printed."
)
TOOL_USE_ARGS = [
"--allowed-tools",
"Bash(echo pong)",
"--permission-mode",
"dontAsk",
]
def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool:
for event in events:
if event.get("type") != "assistant":
continue
message = event.get("message") or {}
content = message.get("content")
if not isinstance(content, list):
continue
for block in content:
if isinstance(block, dict) and block.get("type") == "tool_use":
return True
return False
def test_tool_use_bedrock_mantle(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_mantle_cells_enabled()
proxy = require_proxy(compat_result)
outcomes = run_claude_models_parallel(
models=BEDROCK_MANTLE_MODELS,
prompt=TOOL_USE_PROMPT,
base_url=proxy.base_url,
api_key=proxy.api_key,
extra_args=TOOL_USE_ARGS,
)
failures = []
for model in BEDROCK_MANTLE_MODELS:
outcome = outcomes[model]
if isinstance(outcome, ClaudeCLIError):
error = f"[{model}] {outcome}"
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
if outcome.exit_code != 0:
error = f"[{model}] claude CLI failed: {failure_diagnostic(outcome)}"
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
if not _has_tool_use_event(outcome.events):
error = (
f"[{model}] no tool_use content block observed in stream-json events"
)
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
compat_result.add({"status": "pass"})
if failures:
pytest.fail("; ".join(failures), pytrace=False)

View file

@ -0,0 +1,110 @@
"""tool_use x OpenAI (GPT-5.6).
Drive the real `claude` CLI against a running LiteLLM proxy that
routes Anthropic Messages requests to OpenAI's GPT-5.6 family (Sol,
Terra, Luna), ask the model to invoke a built-in tool (`Bash`), and
assert that a `tool_use` content block came back over the wire.
Claude Code declares its tools in Anthropic `tools` format; LiteLLM's
`openai/gpt-*` route translates them to OpenAI function calling and
maps the returned `tool_calls` back to Anthropic `tool_use` blocks, so
this cell exercises the tool-schema translation in both directions.
Bash is restricted to the exact command `echo pong` plus
`--permission-mode dontAsk`; see `tool_use/test_anthropic.py` for the
security rationale.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/tool_use/test_openai.py
^^^^^^^^ ^^^^^^
feature_id provider
"""
from __future__ import annotations
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,
run_claude_models_parallel,
)
OPENAI_MODELS = [
"gpt-5-6-sol-openai",
"gpt-5-6-terra-openai",
"gpt-5-6-luna-openai",
]
TOOL_USE_PROMPT = (
"Use the Bash tool to run the command `echo pong` and report what it printed."
)
TOOL_USE_ARGS = [
"--allowed-tools",
"Bash(echo pong)",
"--permission-mode",
"dontAsk",
]
def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool:
for event in events:
if event.get("type") != "assistant":
continue
message = event.get("message") or {}
content = message.get("content")
if not isinstance(content, list):
continue
for block in content:
if isinstance(block, dict) and block.get("type") == "tool_use":
return True
return False
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(
models=OPENAI_MODELS,
prompt=TOOL_USE_PROMPT,
base_url=proxy.base_url,
api_key=proxy.api_key,
extra_args=TOOL_USE_ARGS,
)
failures = []
for model in OPENAI_MODELS:
outcome = outcomes[model]
if isinstance(outcome, ClaudeCLIError):
error = f"[{model}] {outcome}"
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
if outcome.exit_code != 0:
error = f"[{model}] claude CLI failed: {failure_diagnostic(outcome)}"
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
if not _has_tool_use_event(outcome.events):
error = (
f"[{model}] no tool_use content block observed in stream-json events"
)
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
compat_result.add({"status": "pass"})
if failures:
pytest.fail("; ".join(failures), pytrace=False)

View file

@ -0,0 +1,33 @@
"""tool_use x Vertex AI (GPT-5.6) — not applicable.
GCP is the only one of the big-three clouds without OpenAI's
closed-weight GPT-5.6 family (Sol / Terra / Luna); Vertex AI Model
Garden carries only the open-weight gpt-oss MaaS models. The cell
reports `not_applicable` so the published matrix documents the gap
explicitly instead of leaving a `not_tested` hole.
This stub never drives the `claude` CLI, so it grants no tools and is
exempt from the Bash allow-rule pin enforced by
`_pr_gate_unit_tests/test_bash_tool_restrictions.py`.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/tool_use/test_vertex_ai_gpt.py
^^^^^^^^ ^^^^^^^^^^^^^
feature_id provider
"""
from __future__ import annotations
from claude_code._gpt_cells import VERTEX_AI_GPT_NOT_APPLICABLE_REASON
def test_tool_use_vertex_ai_gpt(compat_result):
"""Record the static not_applicable outcome for this cell."""
compat_result.set(
{
"status": "not_applicable",
"reason": VERTEX_AI_GPT_NOT_APPLICABLE_REASON,
}
)

View file

@ -0,0 +1,138 @@
"""tool_use_streaming x Azure OpenAI (GPT-5.6).
Drive the real `claude` CLI in headless `--output-format stream-json`
mode against a running LiteLLM proxy that routes Anthropic Messages
requests to Azure OpenAI deployments of the GPT-5.6 family (Sol,
Terra, Luna), ask the model to invoke a built-in tool (`Bash`), and
assert that the upstream (a) emitted a `tool_use` content block and
(b) streamed the tool input incrementally as `input_json_delta`
events.
Azure OpenAI streams tool arguments in the same chat-completions
fragment shape as openai.com; LiteLLM must re-emit them as Anthropic
`input_json_delta` deltas rather than buffering the full input into
one complete block.
Bash is restricted to the exact command `echo pong` plus
`--permission-mode dontAsk`; see `tool_use/test_anthropic.py` for the
security rationale.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/tool_use_streaming/test_azure_openai.py
^^^^^^^^^^^^^^^^^^ ^^^^^^^^^^^^
feature_id provider
"""
from __future__ import annotations
from typing import Any, Mapping, Sequence
import pytest
from claude_code._env import require_proxy
from claude_code.cli_driver import (
ClaudeCLIError,
failure_diagnostic,
run_claude_models_parallel,
)
AZURE_OPENAI_MODELS = [
"gpt-5-6-sol-azure-openai",
"gpt-5-6-terra-azure-openai",
"gpt-5-6-luna-azure-openai",
]
TOOL_USE_PROMPT = (
"Use the Bash tool to run the command `echo pong` and report what it printed."
)
TOOL_USE_ARGS = [
"--allowed-tools",
"Bash(echo pong)",
"--permission-mode",
"dontAsk",
"--include-partial-messages",
]
def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool:
for event in events:
if event.get("type") != "assistant":
continue
message = event.get("message") or {}
content = message.get("content")
if not isinstance(content, list):
continue
for block in content:
if isinstance(block, dict) and block.get("type") == "tool_use":
return True
return False
def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int:
"""Count `input_json_delta` records among the `stream_event`
entries. Zero means the proxy collapsed the streamed tool input
into a single complete block instead of forwarding the incremental
deltas the upstream emitted."""
inner_events = (
event.get("event") for event in events if event.get("type") == "stream_event"
)
return sum(
1
for inner in inner_events
if isinstance(inner, Mapping)
and inner.get("type") == "content_block_delta"
and isinstance(inner.get("delta"), Mapping)
and inner["delta"].get("type") == "input_json_delta"
)
def test_tool_use_streaming_azure_openai(compat_result):
proxy = require_proxy(compat_result)
outcomes = run_claude_models_parallel(
models=AZURE_OPENAI_MODELS,
prompt=TOOL_USE_PROMPT,
base_url=proxy.base_url,
api_key=proxy.api_key,
extra_args=TOOL_USE_ARGS,
)
failures = []
for model in AZURE_OPENAI_MODELS:
outcome = outcomes[model]
if isinstance(outcome, ClaudeCLIError):
error = f"[{model}] {outcome}"
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
if outcome.exit_code != 0:
error = f"[{model}] claude CLI failed: {failure_diagnostic(outcome)}"
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
if not _has_tool_use_event(outcome.events):
error = (
f"[{model}] no tool_use content block observed in stream-json events"
)
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
if _count_input_json_deltas(outcome.events) == 0:
error = (
f"[{model}] no input_json_delta stream events observed; proxy "
f"likely buffered the tool input into a complete block or "
f"stripped fine-grained tool streaming"
)
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
compat_result.add({"status": "pass"})
if failures:
pytest.fail("; ".join(failures), pytrace=False)

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