mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_ocr_rust_default
This commit is contained in:
commit
cd63b40255
156 changed files with 9789 additions and 1621 deletions
3
.gitignore
vendored
3
.gitignore
vendored
|
|
@ -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*
|
||||
|
|
|
|||
BIN
dist/litellm-1.79.1.tar.gz
vendored
BIN
dist/litellm-1.79.1.tar.gz
vendored
Binary file not shown.
94
litellm-rust/Cargo.lock
generated
94
litellm-rust/Cargo.lock
generated
|
|
@ -28,12 +28,6 @@ version = "1.1.2"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0"
|
||||
|
||||
[[package]]
|
||||
name = "autocfg"
|
||||
version = "1.5.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53"
|
||||
|
||||
[[package]]
|
||||
name = "axum"
|
||||
version = "0.7.9"
|
||||
|
|
@ -392,21 +386,6 @@ dependencies = [
|
|||
"percent-encoding",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "futures"
|
||||
version = "0.3.32"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8b147ee9d1f6d097cef9ce628cd2ee62288d963e16fb287bd9286455b241382d"
|
||||
dependencies = [
|
||||
"futures-channel",
|
||||
"futures-core",
|
||||
"futures-executor",
|
||||
"futures-io",
|
||||
"futures-sink",
|
||||
"futures-task",
|
||||
"futures-util",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "futures-channel"
|
||||
version = "0.3.32"
|
||||
|
|
@ -423,17 +402,6 @@ version = "0.3.32"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d"
|
||||
|
||||
[[package]]
|
||||
name = "futures-executor"
|
||||
version = "0.3.32"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "baf29c38818342a3b26b5b923639e7b1f4a61fc5e76102d4b1981c6dc7a7579d"
|
||||
dependencies = [
|
||||
"futures-core",
|
||||
"futures-task",
|
||||
"futures-util",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "futures-io"
|
||||
version = "0.3.32"
|
||||
|
|
@ -469,7 +437,6 @@ version = "0.3.32"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6"
|
||||
dependencies = [
|
||||
"futures-channel",
|
||||
"futures-core",
|
||||
"futures-io",
|
||||
"futures-macro",
|
||||
|
|
@ -928,15 +895,6 @@ dependencies = [
|
|||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "indoc"
|
||||
version = "2.0.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "79cf5c93f93228cf8efb3ba362535fb11199ac548a09ce117c9b1adc3030d706"
|
||||
dependencies = [
|
||||
"rustversion",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ipnet"
|
||||
version = "2.12.0"
|
||||
|
|
@ -1089,15 +1047,6 @@ version = "2.8.2"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "88904434abc2901f197fe8cc55f0445e7ded921dba5911dad2e2b39b48e663c4"
|
||||
|
||||
[[package]]
|
||||
name = "memoffset"
|
||||
version = "0.9.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "488016bfae457b036d996092f6cb448677611ce4449e970ceaf42695203f218a"
|
||||
dependencies = [
|
||||
"autocfg",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "mime"
|
||||
version = "0.3.17"
|
||||
|
|
@ -1225,29 +1174,26 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "pyo3"
|
||||
version = "0.23.5"
|
||||
version = "0.29.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7778bffd85cf38175ac1f545509665d0b9b92a198ca7941f131f85f7a4f9a872"
|
||||
checksum = "cd274650b21d4bfc26a0a47587962c1edb425f69287324355cd040c3ea66071c"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"indoc",
|
||||
"libc",
|
||||
"memoffset",
|
||||
"once_cell",
|
||||
"portable-atomic",
|
||||
"pyo3-build-config",
|
||||
"pyo3-ffi",
|
||||
"pyo3-macros",
|
||||
"unindent",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pyo3-async-runtimes"
|
||||
version = "0.23.0"
|
||||
version = "0.29.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "977dc837525cfd22919ba6a831413854beb7c99a256c03bf8624ad707e45810e"
|
||||
checksum = "b3ef68daa7316a3fac65e5e18b2203f010346de1c1c53456811a2624673ab046"
|
||||
dependencies = [
|
||||
"futures",
|
||||
"futures-channel",
|
||||
"futures-util",
|
||||
"once_cell",
|
||||
"pin-project-lite",
|
||||
"pyo3",
|
||||
|
|
@ -1256,19 +1202,18 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "pyo3-build-config"
|
||||
version = "0.23.5"
|
||||
version = "0.29.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "94f6cbe86ef3bf18998d9df6e0f3fc1050a8c5efa409bf712e661a4366e010fb"
|
||||
checksum = "c5e2a7d2f0d013342f295c048ad19237add5154a55b1c5a254c0ec93d4109078"
|
||||
dependencies = [
|
||||
"once_cell",
|
||||
"target-lexicon",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pyo3-ffi"
|
||||
version = "0.23.5"
|
||||
version = "0.29.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e9f1b4c431c0bb1c8fb0a338709859eed0d030ff6daa34368d3b152a63dfdd8d"
|
||||
checksum = "ca85c467da1bbc8d866eea5deff9cf29ea5f7785054a17da36e65bda9c05845b"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"pyo3-build-config",
|
||||
|
|
@ -1276,9 +1221,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "pyo3-macros"
|
||||
version = "0.23.5"
|
||||
version = "0.29.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fbc2201328f63c4710f68abdf653c89d8dbc2858b88c5d88b0ff38a75288a9da"
|
||||
checksum = "9ac53762fd065daa3194dd09337a38bd793a188100fd1a9304c4ab312d901771"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"pyo3-macros-backend",
|
||||
|
|
@ -1288,13 +1233,12 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "pyo3-macros-backend"
|
||||
version = "0.23.5"
|
||||
version = "0.29.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fca6726ad0f3da9c9de093d6f116a93c1a38e417ed73bf138472cf4064f72028"
|
||||
checksum = "4ca3a1557399783172dc5bf39cfca835157732532cba56b71d2292161e53b362"
|
||||
dependencies = [
|
||||
"heck",
|
||||
"proc-macro2",
|
||||
"pyo3-build-config",
|
||||
"quote",
|
||||
"syn",
|
||||
]
|
||||
|
|
@ -1974,9 +1918,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "target-lexicon"
|
||||
version = "0.12.16"
|
||||
version = "0.13.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "61c41af27dd6d1e27b1b16b489db798443478cef1f06a660c96db617ba5de3b1"
|
||||
checksum = "adb6935a6f5c20170eeceb1a3835a49e12e19d792f6dd344ccc76a985ca5a6ca"
|
||||
|
||||
[[package]]
|
||||
name = "thiserror"
|
||||
|
|
@ -2250,12 +2194,6 @@ version = "1.0.24"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75"
|
||||
|
||||
[[package]]
|
||||
name = "unindent"
|
||||
version = "0.2.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7264e107f553ccae879d21fbea1d6724ac785e8c3bfc762137959b5802826ef3"
|
||||
|
||||
[[package]]
|
||||
name = "untrusted"
|
||||
version = "0.9.0"
|
||||
|
|
|
|||
|
|
@ -16,8 +16,8 @@ rust-version = "1.86"
|
|||
litellm-core = { path = "crates/core" }
|
||||
litellm-ai-gateway = { path = "crates/ai-gateway", default-features = false }
|
||||
axum = "0.7"
|
||||
pyo3 = "0.23.5"
|
||||
pyo3-async-runtimes = { version = "0.23.0", features = ["tokio-runtime"] }
|
||||
pyo3 = "0.29.0"
|
||||
pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
|
||||
rand = "0.8"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "multipart", "rustls-tls", "http2", "stream"] }
|
||||
serde = { version = "1.0", features = ["derive"] }
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -174,7 +174,7 @@ fn aocr(
|
|||
.await
|
||||
.map_err(|err| Python::with_gil(|py| core_error_to_pyerr(py, err)))?;
|
||||
|
||||
Python::with_gil(|py| json_to_py(py, value))
|
||||
Python::attach(|py| json_to_py(py, value))
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -237,7 +237,9 @@ async def anthropic_messages(
|
|||
AnthropicCacheControlHook,
|
||||
)
|
||||
|
||||
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(messages, system, kwargs)
|
||||
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(
|
||||
messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools
|
||||
)
|
||||
|
||||
original_stream = stream or kwargs.get("_websearch_interception_converted_stream", False)
|
||||
|
||||
|
|
@ -426,7 +428,9 @@ def anthropic_messages_handler(
|
|||
AnthropicCacheControlHook,
|
||||
)
|
||||
|
||||
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(messages, system, kwargs)
|
||||
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(
|
||||
messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools
|
||||
)
|
||||
|
||||
metadata = validate_anthropic_api_metadata(metadata)
|
||||
|
||||
|
|
|
|||
|
|
@ -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 ----
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -3451,7 +3451,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2.2e-05,
|
||||
"output_cost_per_token": 2.64e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -3470,7 +3470,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.00022,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -3489,7 +3489,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -4687,7 +4687,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -4707,7 +4707,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4739,7 +4739,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4771,7 +4771,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4832,7 +4832,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.0002,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -4850,7 +4850,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -7922,7 +7922,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2.2e-05,
|
||||
"output_cost_per_token": 2.64e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -7941,7 +7941,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.00022,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -7960,7 +7960,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -16272,7 +16272,7 @@
|
|||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/glm-5p2": {
|
||||
"cache_read_input_token_cost": 2.6e-07,
|
||||
"cache_read_input_token_cost": 1.4e-07,
|
||||
"input_cost_per_token": 1.4e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
|
|
@ -16686,7 +16686,7 @@
|
|||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/glm-5p2": {
|
||||
"cache_read_input_token_cost": 2.6e-07,
|
||||
"cache_read_input_token_cost": 1.4e-07,
|
||||
"input_cost_per_token": 1.4e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
|
|
@ -22094,7 +22094,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22113,7 +22113,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22207,7 +22207,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22225,7 +22225,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22243,7 +22243,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -24438,7 +24438,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24470,7 +24470,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24502,7 +24502,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24535,7 +24535,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 2.4e-05,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
|
|
@ -24570,7 +24570,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
|
|
@ -24603,7 +24603,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24635,7 +24635,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -43573,7 +43573,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -43606,7 +43606,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
|
|||
|
|
@ -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)}")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
190
litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py
Normal file
190
litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -151,6 +151,80 @@ async def _record_streaming_client_disconnect_if_needed(
|
|||
return True
|
||||
|
||||
|
||||
def _deferred_stream_logging_is_armed(request_data: dict) -> bool:
|
||||
logging_obj = request_data.get("litellm_logging_obj")
|
||||
if logging_obj is None:
|
||||
return False
|
||||
return (
|
||||
getattr(logging_obj, "_on_deferred_stream_complete", None) is not None
|
||||
and getattr(logging_obj, "_deferred_stream_complete_args", None) is not None
|
||||
)
|
||||
|
||||
|
||||
async def _bill_partial_streamed_spend_on_disconnect(request_data: dict, response: object) -> bool:
|
||||
"""
|
||||
A client disconnect throws GeneratorExit/CancelledError into the streaming
|
||||
generator, so neither the success nor the failure logging callback fires
|
||||
and the chunks already streamed (plus any sub-call cost folded into the
|
||||
logging object) would never reach spend tracking. Assemble the partial
|
||||
response from the wrapper's collected chunks and dispatch success logging
|
||||
for it; dispatch_success_handlers dedups against a natural end-of-stream
|
||||
dispatch via has_dispatched_final_stream_success.
|
||||
|
||||
Awaited directly by the shielded cleanup rather than scheduled with
|
||||
create_task: the client is already gone so the extra latency is harmless,
|
||||
and an unrooted task could be garbage-collected before it bills.
|
||||
|
||||
Returns True when a disconnect-time success event owns the request's
|
||||
max_parallel_requests slot release (one was dispatched here, or one had
|
||||
already been dispatched for this stream), so the caller can skip the
|
||||
explicit slot release and avoid a double release. Returns False when no
|
||||
success event fired (logging disabled, nothing streamed, or assembly
|
||||
failed) and the caller must release the slot itself.
|
||||
"""
|
||||
if litellm.disable_streaming_logging is True:
|
||||
return False
|
||||
logging_obj = request_data.get("litellm_logging_obj")
|
||||
if not isinstance(logging_obj, LiteLLMLoggingObj):
|
||||
return False
|
||||
if logging_obj.model_call_details.get("has_dispatched_final_stream_success"):
|
||||
# A natural end-of-stream success event already fired and released the
|
||||
# slot; do not bill again, and let the caller skip the slot release.
|
||||
return True
|
||||
chunks: object = getattr(response, "chunks", None)
|
||||
if not isinstance(chunks, list) or not chunks:
|
||||
return False
|
||||
verbose_proxy_logger.debug(
|
||||
"Billing partial streamed spend for %s chunks after client disconnect, litellm_call_id=%s",
|
||||
len(chunks),
|
||||
request_data.get("litellm_call_id"),
|
||||
)
|
||||
messages: object = getattr(response, "messages", None)
|
||||
try:
|
||||
partial_response = litellm.stream_chunk_builder(
|
||||
chunks=chunks,
|
||||
messages=messages if isinstance(messages, list) else None,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # partial billing is best-effort; never break stream teardown
|
||||
verbose_proxy_logger.debug("Failed to assemble partial streamed response for disconnect billing: %s", e)
|
||||
return False
|
||||
if partial_response is None:
|
||||
return False
|
||||
try:
|
||||
await logging_obj.dispatch_success_handlers(
|
||||
partial_response,
|
||||
cache_hit=False,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # partial billing is best-effort; never break stream teardown
|
||||
verbose_proxy_logger.debug("Failed to dispatch disconnect billing event: %s", e)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
async def _cancel_pending_gather_tasks(tasks: list["asyncio.Task[Any]"]) -> None:
|
||||
pending_tasks = [task for task in tasks if not task.done()]
|
||||
for task in pending_tasks:
|
||||
|
|
@ -851,9 +925,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# If conversion fails, use original spend
|
||||
pass
|
||||
|
||||
model_name = ProxyBaseLLMRequestProcessing._get_deployment_model_name(litellm_logging_obj)
|
||||
|
||||
headers = {
|
||||
"x-litellm-call-id": call_id,
|
||||
"x-litellm-model-id": model_id,
|
||||
"x-litellm-model-name": model_name,
|
||||
"x-litellm-cache-key": cache_key,
|
||||
"x-litellm-model-api-base": (
|
||||
api_base.split("?")[0] if api_base else None
|
||||
|
|
@ -1322,6 +1399,27 @@ class ProxyBaseLLMRequestProcessing:
|
|||
model_id = model_info.get("id", "") or ""
|
||||
return model_id
|
||||
|
||||
@staticmethod
|
||||
def _get_deployment_model_name(
|
||||
litellm_logging_obj: LiteLLMLoggingObj | None,
|
||||
) -> str | None:
|
||||
"""Extract the underlying deployment model string (e.g. ``azure/gpt-4o``).
|
||||
|
||||
The router rewrites the response ``model`` field to the model-group alias
|
||||
the client requested, so neither the response body nor the existing
|
||||
headers expose the concrete deployment model. The router records it under
|
||||
``litellm_params`` metadata as ``deployment``, so read it back from there.
|
||||
"""
|
||||
litellm_params = getattr(litellm_logging_obj, "litellm_params", None)
|
||||
if not isinstance(litellm_params, dict):
|
||||
return None
|
||||
for key in ("litellm_metadata", "metadata"):
|
||||
metadata = litellm_params.get(key, {}) or {}
|
||||
deployment = metadata.get("deployment")
|
||||
if deployment:
|
||||
return deployment
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _response_cost_from_logging_obj(
|
||||
*,
|
||||
|
|
@ -2575,6 +2673,8 @@ class ProxyBaseLLMRequestProcessing:
|
|||
response: Any,
|
||||
stream_completed: bool = False,
|
||||
client_disconnected: bool = False,
|
||||
user_api_key_dict: UserAPIKeyAuth | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
) -> None:
|
||||
with anyio.CancelScope(shield=True):
|
||||
should_record_client_disconnect = client_disconnected or (not stream_completed)
|
||||
|
|
@ -2586,7 +2686,28 @@ class ProxyBaseLLMRequestProcessing:
|
|||
client_disconnected,
|
||||
)
|
||||
if recorded_client_disconnect:
|
||||
deferred_stream_logging_armed = _deferred_stream_logging_is_armed(request_data)
|
||||
ProxyLogging._fire_deferred_stream_logging(request_data)
|
||||
# A disconnect-time success event (the deferred-guardrail flush
|
||||
# above, or the partial-spend billing below) releases the
|
||||
# request's max_parallel_requests slot through the limiter's
|
||||
# own success callback. Release the slot explicitly only when
|
||||
# no such event fires, so exactly one release happens; two
|
||||
# concurrent releases would race and double-decrement under the
|
||||
# limiter's in-memory fallback.
|
||||
success_event_owns_slot_release = deferred_stream_logging_armed
|
||||
if not deferred_stream_logging_armed:
|
||||
success_event_owns_slot_release = await _bill_partial_streamed_spend_on_disconnect(
|
||||
request_data, response
|
||||
)
|
||||
if (
|
||||
not success_event_owns_slot_release
|
||||
and proxy_logging_obj is not None
|
||||
and user_api_key_dict is not None
|
||||
):
|
||||
await proxy_logging_obj._arelease_max_parallel_requests_on_disconnect(
|
||||
user_api_key_dict, request_data
|
||||
)
|
||||
|
||||
if hasattr(response, "aclose"):
|
||||
try:
|
||||
|
|
@ -2675,12 +2796,13 @@ class ProxyBaseLLMRequestProcessing:
|
|||
except (asyncio.CancelledError, GeneratorExit):
|
||||
# Client disconnected mid-stream. CancelledError / GeneratorExit
|
||||
# are BaseException and bypass the success/failure logging
|
||||
# callbacks that release the pre-call max_parallel_requests +1;
|
||||
# release it here. This is the outermost generator Starlette closes
|
||||
# on disconnect, so the nested iterator hook (which only sees
|
||||
# GeneratorExit on GC) cannot own the refund.
|
||||
# callbacks that release the pre-call max_parallel_requests +1.
|
||||
# Flag the disconnect; the shielded cleanup in `finally` owns the
|
||||
# slot release so it can coordinate with disconnect-time success
|
||||
# billing and release exactly once. This is the outermost generator
|
||||
# Starlette closes on disconnect, so the nested iterator hook (which
|
||||
# only sees GeneratorExit on GC) cannot own the refund.
|
||||
if not stream_completed:
|
||||
proxy_logging_obj._release_max_parallel_requests_on_disconnect(user_api_key_dict)
|
||||
client_disconnected = True
|
||||
if not delivered_chunk:
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
|
|
@ -2723,6 +2845,8 @@ class ProxyBaseLLMRequestProcessing:
|
|||
response=response,
|
||||
stream_completed=stream_completed,
|
||||
client_disconnected=client_disconnected,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
44
litellm/proxy/guardrails/guardrail_hooks/singulr/__init__.py
Normal file
44
litellm/proxy/guardrails/guardrail_hooks/singulr/__init__.py
Normal 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,
|
||||
}
|
||||
216
litellm/proxy/guardrails/guardrail_hooks/singulr/singulr.py
Normal file
216
litellm/proxy/guardrails/guardrail_hooks/singulr/singulr.py
Normal 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
|
||||
|
|
@ -0,0 +1,71 @@
|
|||
from typing import TYPE_CHECKING
|
||||
|
||||
import litellm
|
||||
from litellm.types.guardrails import SupportedGuardrailIntegrations
|
||||
|
||||
from .straiker import StraikerGuardrail
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
_OPTIONAL_INIT_FIELDS = (
|
||||
"timeout",
|
||||
"max_retries",
|
||||
"initial_backoff",
|
||||
"max_backoff",
|
||||
"unreachable_fallback",
|
||||
"fail_on_error",
|
||||
"max_payload_bytes",
|
||||
"custom_headers",
|
||||
"metadata",
|
||||
"verbose",
|
||||
)
|
||||
|
||||
|
||||
def _get_config_value(litellm_params: "LitellmParams", optional_params: object, attribute_name: str) -> object:
|
||||
if optional_params is not None:
|
||||
if isinstance(optional_params, dict):
|
||||
value = optional_params.get(attribute_name)
|
||||
else:
|
||||
value = getattr(optional_params, attribute_name, None)
|
||||
if value is not None:
|
||||
return value
|
||||
return getattr(litellm_params, attribute_name, None)
|
||||
|
||||
|
||||
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
|
||||
optional_params = getattr(litellm_params, "optional_params", None)
|
||||
api_key = litellm_params.api_key
|
||||
if not api_key:
|
||||
raise ValueError("api_key is required for straiker")
|
||||
|
||||
api_base = litellm_params.api_base or "https://api.prod.straiker.ai"
|
||||
default_app = getattr(litellm_params, "default_app", None) or getattr(litellm_params, "source", None)
|
||||
source = default_app if isinstance(default_app, str) and default_app else "LiteLLM Gateway"
|
||||
kwargs: dict[str, object] = {
|
||||
field: value
|
||||
for field in _OPTIONAL_INIT_FIELDS
|
||||
for value in [_get_config_value(litellm_params, optional_params, field)]
|
||||
if value is not None
|
||||
}
|
||||
_callback = StraikerGuardrail(
|
||||
api_key=api_key,
|
||||
api_base=api_base if isinstance(api_base, str) else "https://api.prod.straiker.ai",
|
||||
source=source,
|
||||
guardrail_name=guardrail.get("guardrail_name", "straiker"),
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=litellm_params.default_on,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
litellm.logging_callback_manager.add_litellm_callback(_callback)
|
||||
return _callback
|
||||
|
||||
|
||||
guardrail_initializer_registry = {
|
||||
SupportedGuardrailIntegrations.STRAIKER.value: initialize_guardrail,
|
||||
}
|
||||
|
||||
guardrail_class_registry = {
|
||||
SupportedGuardrailIntegrations.STRAIKER.value: StraikerGuardrail,
|
||||
}
|
||||
541
litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py
Normal file
541
litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py
Normal file
|
|
@ -0,0 +1,541 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import random
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Literal, NoReturn
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._version import version as litellm_version
|
||||
from litellm.exceptions import (
|
||||
BadRequestError,
|
||||
GuardrailRaisedException,
|
||||
ModifyResponseException,
|
||||
Timeout,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
get_session_id_from_request_data,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
|
||||
STRAIKER_WEBHOOK_SCHEMA_VERSION,
|
||||
StraikerGuardrailConfigModel,
|
||||
StraikerWebhookApplication,
|
||||
StraikerWebhookContent,
|
||||
StraikerWebhookContext,
|
||||
StraikerWebhookEvent,
|
||||
StraikerWebhookIdentity,
|
||||
StraikerWebhookRequest,
|
||||
StraikerWebhookResponse,
|
||||
StraikerWebhookStream,
|
||||
StraikerWebhookUsage,
|
||||
)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, Usage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
||||
GUARDRAIL_NAME = "straiker"
|
||||
DEFAULT_BLOCK_MESSAGE = "Content violates policy"
|
||||
DEFAULT_API_BASE = "https://api.prod.straiker.ai"
|
||||
DEFAULT_MAX_PAYLOAD_BYTES = 524288
|
||||
WEBHOOK_PATH = "/api/v1/detect/webhook"
|
||||
RETRY_STATUS = frozenset({408, 429, 500, 502, 503, 504})
|
||||
UNREACHABLE_STATUS = frozenset({502, 503, 504})
|
||||
_APPLICATION_METADATA_KEYS = frozenset({"agent_id", "app_name"})
|
||||
_OPAQUE_METADATA_SCALAR_TYPES = (str, int, float, bool)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _WebhookFailure:
|
||||
message: str
|
||||
is_unreachable: bool
|
||||
|
||||
|
||||
def _as_dict(value: object) -> dict:
|
||||
return value if isinstance(value, dict) else {}
|
||||
|
||||
|
||||
def _merged_metadata(request_data: dict) -> dict:
|
||||
return {
|
||||
**_as_dict(request_data.get("metadata")),
|
||||
**_as_dict(request_data.get("litellm_metadata")),
|
||||
}
|
||||
|
||||
|
||||
def _as_optional_str(value: object) -> str | None:
|
||||
return value if isinstance(value, str) and value else None
|
||||
|
||||
|
||||
def _build_webhook_metadata(request_data: dict, default_metadata: dict[str, str]) -> dict[str, object] | None:
|
||||
out: dict[str, object] = {}
|
||||
for key, value in _as_dict(request_data.get("metadata")).items():
|
||||
if key in _APPLICATION_METADATA_KEYS or key.startswith("user_api"):
|
||||
continue
|
||||
if key == "session_id":
|
||||
continue
|
||||
if isinstance(value, _OPAQUE_METADATA_SCALAR_TYPES):
|
||||
out[key] = value
|
||||
out.update(default_metadata)
|
||||
return out or None
|
||||
|
||||
|
||||
def _extract_identity(request_data: dict) -> StraikerWebhookIdentity:
|
||||
meta = _merged_metadata(request_data)
|
||||
return StraikerWebhookIdentity(
|
||||
litellm_key=_as_optional_str(meta.get("user_api_key_alias"))
|
||||
or _as_optional_str(meta.get("user_api_key_hash"))
|
||||
or _as_optional_str(meta.get("user_api_key_token")),
|
||||
litellm_team=_as_optional_str(meta.get("user_api_key_team_alias"))
|
||||
or _as_optional_str(meta.get("user_api_key_team_id")),
|
||||
litellm_user_id=_as_optional_str(meta.get("user_api_key_user_id")),
|
||||
litellm_user_email=_as_optional_str(meta.get("user_api_key_user_email")),
|
||||
litellm_org_id=_as_optional_str(meta.get("user_api_key_org_id")),
|
||||
end_user_id=_as_optional_str(meta.get("user_api_key_end_user_id")),
|
||||
)
|
||||
|
||||
|
||||
def _resolve_provider(request_data: dict, model: str | None) -> str | None:
|
||||
litellm_params = _as_dict(request_data.get("litellm_params"))
|
||||
custom_llm_provider = request_data.get("custom_llm_provider") or litellm_params.get("custom_llm_provider")
|
||||
if custom_llm_provider:
|
||||
return custom_llm_provider
|
||||
if not model:
|
||||
return None
|
||||
try:
|
||||
_, provider, _, _ = get_llm_provider(
|
||||
model=model,
|
||||
api_base=request_data.get("api_base") or litellm_params.get("api_base"),
|
||||
api_key=request_data.get("api_key") or litellm_params.get("api_key"),
|
||||
)
|
||||
except BadRequestError:
|
||||
return None
|
||||
return provider or None
|
||||
|
||||
|
||||
def _resolve_destination(request_data: dict) -> str | None:
|
||||
litellm_params = _as_dict(request_data.get("litellm_params"))
|
||||
api_base = request_data.get("api_base") or litellm_params.get("api_base")
|
||||
if not isinstance(api_base, str):
|
||||
return None
|
||||
try:
|
||||
return urlsplit(api_base).hostname
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def _resolve_call_surface(logging_obj: LiteLLMLoggingObj | None, request_data: dict) -> str:
|
||||
call_type = (
|
||||
(getattr(logging_obj, "call_type", None) if logging_obj is not None else None)
|
||||
or request_data.get("call_type")
|
||||
or request_data.get("litellm_call_type")
|
||||
)
|
||||
return call_type if isinstance(call_type, str) and call_type else "unknown"
|
||||
|
||||
|
||||
def _response_finish_reason(response: Any) -> str | None:
|
||||
choices = getattr(response, "choices", None)
|
||||
if not isinstance(choices, list):
|
||||
return None
|
||||
for choice in choices:
|
||||
reason = getattr(choice, "finish_reason", None)
|
||||
if isinstance(reason, str) and reason:
|
||||
return reason
|
||||
return None
|
||||
|
||||
|
||||
def _build_usage(response: object) -> StraikerWebhookUsage | None:
|
||||
usage = getattr(response, "usage", None)
|
||||
if not isinstance(usage, Usage):
|
||||
return None
|
||||
input_tokens = usage.prompt_tokens
|
||||
output_tokens = usage.completion_tokens
|
||||
if input_tokens is None and output_tokens is None:
|
||||
return None
|
||||
return StraikerWebhookUsage(input_tokens=input_tokens, output_tokens=output_tokens)
|
||||
|
||||
|
||||
def _is_streamed_request(request_data: dict) -> bool:
|
||||
if request_data.get("stream") is True:
|
||||
return True
|
||||
body = _as_dict(_as_dict(request_data.get("proxy_server_request")).get("body"))
|
||||
return body.get("stream") is True
|
||||
|
||||
|
||||
class StraikerGuardrail(CustomGuardrail):
|
||||
@staticmethod
|
||||
def get_config_model() -> type[GuardrailConfigModel]:
|
||||
return StraikerGuardrailConfigModel
|
||||
|
||||
@classmethod
|
||||
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
|
||||
return [
|
||||
GuardrailEventHooks.pre_call,
|
||||
GuardrailEventHooks.post_call,
|
||||
]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str,
|
||||
api_base: str = DEFAULT_API_BASE,
|
||||
source: str = "LiteLLM Gateway",
|
||||
timeout: float = 5.0,
|
||||
max_retries: int = 2,
|
||||
initial_backoff: float = 0.1,
|
||||
max_backoff: float = 2.0,
|
||||
unreachable_fallback: Literal["fail_open", "fail_closed"] = "fail_closed",
|
||||
fail_on_error: bool = True,
|
||||
max_payload_bytes: int = DEFAULT_MAX_PAYLOAD_BYTES,
|
||||
custom_headers: dict[str, str] | None = None,
|
||||
metadata: dict[str, str] | None = None,
|
||||
verbose: bool = False,
|
||||
async_handler: httpx.AsyncClient | None = None,
|
||||
**kwargs: object,
|
||||
) -> None:
|
||||
if not api_key:
|
||||
raise ValueError("api_key must be non-empty")
|
||||
if unreachable_fallback not in ("fail_open", "fail_closed"):
|
||||
raise ValueError(f"unreachable_fallback must be 'fail_open' or 'fail_closed'; got {unreachable_fallback!r}")
|
||||
|
||||
self.api_key = api_key
|
||||
self.api_base = api_base.rstrip("/")
|
||||
self.source = source
|
||||
self.timeout = float(timeout)
|
||||
self.max_retries = max(0, int(max_retries))
|
||||
self.initial_backoff = max(0.0, float(initial_backoff))
|
||||
self.max_backoff = max(self.initial_backoff, float(max_backoff))
|
||||
self.unreachable_fallback = unreachable_fallback
|
||||
self.fail_on_error = fail_on_error
|
||||
self.max_payload_bytes = int(max_payload_bytes)
|
||||
self.custom_headers = dict(custom_headers) if custom_headers else {}
|
||||
self.default_metadata = dict(metadata) if metadata else {}
|
||||
self.verbose = bool(verbose)
|
||||
|
||||
self.streaming_end_of_stream_only = True
|
||||
self.streaming_buffer_until_moderated = True
|
||||
|
||||
self.async_handler = async_handler or get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.GuardrailCallback,
|
||||
)
|
||||
|
||||
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def _webhook_url(self) -> str:
|
||||
return f"{self.api_base}{WEBHOOK_PATH}"
|
||||
|
||||
def _headers(self) -> dict[str, str]:
|
||||
reserved = {"authorization", "content-type", "x-straiker-webhook-format"}
|
||||
extra = {k: v for k, v in self.custom_headers.items() if k.lower() not in reserved}
|
||||
return {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"X-Straiker-Webhook-Format": "litellm",
|
||||
**extra,
|
||||
}
|
||||
|
||||
def _build_application(self, request_data: dict) -> StraikerWebhookApplication:
|
||||
meta = _merged_metadata(request_data)
|
||||
agent_id = _as_optional_str(meta.get("agent_id"))
|
||||
return StraikerWebhookApplication(
|
||||
source=agent_id or self.source,
|
||||
name=_as_optional_str(meta.get("app_name")),
|
||||
)
|
||||
|
||||
def _build_context(
|
||||
self,
|
||||
request_data: dict,
|
||||
model: str | None,
|
||||
logging_obj: LiteLLMLoggingObj | None,
|
||||
) -> StraikerWebhookContext:
|
||||
return StraikerWebhookContext(
|
||||
call_surface=_resolve_call_surface(logging_obj, request_data),
|
||||
model=model,
|
||||
model_provider=_resolve_provider(request_data, model),
|
||||
destination=_resolve_destination(request_data),
|
||||
session_id=get_session_id_from_request_data(request_data),
|
||||
litellm_call_id=getattr(logging_obj, "litellm_call_id", None) if logging_obj else None,
|
||||
litellm_trace_id=getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None,
|
||||
litellm_version=litellm_version,
|
||||
)
|
||||
|
||||
def _build_envelope(
|
||||
self,
|
||||
*,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: LiteLLMLoggingObj | None,
|
||||
) -> StraikerWebhookRequest:
|
||||
model = inputs.get("model") or request_data.get("model")
|
||||
call_id = getattr(logging_obj, "litellm_call_id", None) if logging_obj else None
|
||||
event_id = f"{call_id or 'litellm'}:{input_type}"
|
||||
|
||||
content = StraikerWebhookContent(
|
||||
texts=list(inputs.get("texts") or []),
|
||||
images=list(inputs.get("images") or []),
|
||||
structured_messages=inputs.get("structured_messages"),
|
||||
tools=inputs.get("tools"),
|
||||
tool_calls=inputs.get("tool_calls"),
|
||||
)
|
||||
|
||||
if input_type == "request":
|
||||
event = StraikerWebhookEvent(type="pre_call", id=event_id)
|
||||
return StraikerWebhookRequest(
|
||||
event=event,
|
||||
request=content,
|
||||
context=self._build_context(request_data, model, logging_obj),
|
||||
identity=_extract_identity(request_data),
|
||||
application=self._build_application(request_data),
|
||||
metadata=_build_webhook_metadata(request_data, self.default_metadata),
|
||||
)
|
||||
|
||||
response_obj = request_data.get("response")
|
||||
content.finish_reason = _response_finish_reason(response_obj)
|
||||
original_messages = request_data.get("messages")
|
||||
request_content = StraikerWebhookContent(
|
||||
structured_messages=original_messages if isinstance(original_messages, list) else None,
|
||||
)
|
||||
phase: Literal["none", "assembled"] = "assembled" if _is_streamed_request(request_data) else "none"
|
||||
event = StraikerWebhookEvent(type="post_call", id=event_id, stream=StraikerWebhookStream(phase=phase))
|
||||
return StraikerWebhookRequest(
|
||||
event=event,
|
||||
request=request_content,
|
||||
response=content,
|
||||
context=self._build_context(request_data, model, logging_obj),
|
||||
identity=_extract_identity(request_data),
|
||||
application=self._build_application(request_data),
|
||||
usage=_build_usage(response_obj),
|
||||
metadata=_build_webhook_metadata(request_data, self.default_metadata),
|
||||
)
|
||||
|
||||
async def _post_webhook(self, payload: dict) -> tuple[StraikerWebhookResponse | None, _WebhookFailure | None]:
|
||||
try:
|
||||
body = json.dumps(payload).encode("utf-8")
|
||||
except (TypeError, ValueError, OverflowError) as error:
|
||||
return None, _WebhookFailure(f"request serialization failed: {error}", is_unreachable=False)
|
||||
body_bytes = len(body)
|
||||
if body_bytes > self.max_payload_bytes:
|
||||
return None, _WebhookFailure(
|
||||
f"payload {body_bytes}B exceeds max_payload_bytes {self.max_payload_bytes}",
|
||||
is_unreachable=False,
|
||||
)
|
||||
|
||||
url = self._webhook_url()
|
||||
headers = self._headers()
|
||||
attempts = self.max_retries + 1
|
||||
last_failure: _WebhookFailure | None = None
|
||||
|
||||
if self.verbose:
|
||||
verbose_proxy_logger.info(
|
||||
json.dumps(
|
||||
{
|
||||
"event": "straiker.webhook_request",
|
||||
"url": url,
|
||||
"bytes": body_bytes,
|
||||
"payload": payload,
|
||||
},
|
||||
default=str,
|
||||
)
|
||||
)
|
||||
|
||||
for attempt in range(attempts):
|
||||
try:
|
||||
resp = await self.async_handler.post(url, content=body, headers=headers, timeout=self.timeout)
|
||||
if resp.status_code == 200:
|
||||
try:
|
||||
body = resp.json()
|
||||
parsed = StraikerWebhookResponse.model_validate(body)
|
||||
except (ValidationError, json.JSONDecodeError) as ve:
|
||||
return None, _WebhookFailure(f"invalid response schema: {ve}", is_unreachable=False)
|
||||
if self.verbose:
|
||||
verbose_proxy_logger.info(
|
||||
json.dumps(
|
||||
{
|
||||
"event": "straiker.webhook_response",
|
||||
"status_code": resp.status_code,
|
||||
"body": body,
|
||||
},
|
||||
default=str,
|
||||
)
|
||||
)
|
||||
return parsed, None
|
||||
last_failure = _WebhookFailure(
|
||||
f"HTTP {resp.status_code}: {resp.text[:200]}",
|
||||
is_unreachable=resp.status_code in UNREACHABLE_STATUS,
|
||||
)
|
||||
if resp.status_code not in RETRY_STATUS:
|
||||
return None, last_failure
|
||||
except (httpx.RequestError, asyncio.TimeoutError, Timeout) as e:
|
||||
last_failure = _WebhookFailure(f"{type(e).__name__}: {e}", is_unreachable=True)
|
||||
except (json.JSONDecodeError, TypeError, ValueError) as e:
|
||||
return None, _WebhookFailure(f"{type(e).__name__}: {e}", is_unreachable=False)
|
||||
|
||||
if attempt < attempts - 1:
|
||||
backoff = min(self.initial_backoff * (2**attempt), self.max_backoff)
|
||||
await asyncio.sleep(random.uniform(0, backoff))
|
||||
|
||||
return None, last_failure or _WebhookFailure("unknown error", is_unreachable=True)
|
||||
|
||||
def _record(
|
||||
self,
|
||||
*,
|
||||
request_data: dict,
|
||||
logging_obj: LiteLLMLoggingObj | None,
|
||||
parsed: StraikerWebhookResponse,
|
||||
) -> None:
|
||||
if not self.verbose:
|
||||
return
|
||||
response_obj = request_data.get("response")
|
||||
hidden = getattr(response_obj, "_hidden_params", None)
|
||||
if isinstance(hidden, dict):
|
||||
straiker_hidden = hidden.setdefault("straiker", {})
|
||||
if isinstance(straiker_hidden, dict):
|
||||
straiker_hidden.update({"action": parsed.action, "turn_id": parsed.turn_id})
|
||||
|
||||
def _fail(
|
||||
self,
|
||||
*,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
error: str,
|
||||
is_unreachable: bool,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
fail_open = (is_unreachable and self.unreachable_fallback == "fail_open") or not self.fail_on_error
|
||||
verbose_proxy_logger.error(
|
||||
json.dumps(
|
||||
{
|
||||
"event": "straiker.error",
|
||||
"input_type": input_type,
|
||||
"error": error,
|
||||
"fail_open": fail_open,
|
||||
},
|
||||
default=str,
|
||||
)
|
||||
)
|
||||
if fail_open:
|
||||
return inputs
|
||||
self._block(
|
||||
request_data=request_data,
|
||||
input_type=input_type,
|
||||
message=f"Straiker detection unavailable: {error}",
|
||||
)
|
||||
|
||||
def _block(
|
||||
self,
|
||||
*,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
message: str,
|
||||
) -> NoReturn:
|
||||
if input_type == "request":
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name or GUARDRAIL_NAME,
|
||||
message=message,
|
||||
should_wrap_with_default_message=False,
|
||||
)
|
||||
raise ModifyResponseException(
|
||||
message=message,
|
||||
model=request_data.get("model", "unknown") or "unknown",
|
||||
request_data=request_data,
|
||||
guardrail_name=self.guardrail_name or GUARDRAIL_NAME,
|
||||
original_response=request_data.get("response"),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _intervened_inputs(
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
parsed: StraikerWebhookResponse,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
return_inputs: GenericGuardrailAPIInputs = {}
|
||||
return_inputs.update(inputs)
|
||||
if parsed.texts is not None:
|
||||
return_inputs["texts"] = parsed.texts
|
||||
return return_inputs
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
try:
|
||||
envelope = self._build_envelope(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type=input_type,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
payload = envelope.model_dump(mode="json", exclude_none=True)
|
||||
except (ValidationError, TypeError, ValueError) as error:
|
||||
return self._fail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type=input_type,
|
||||
error=str(error),
|
||||
is_unreachable=False,
|
||||
)
|
||||
|
||||
parsed, failure = await self._post_webhook(payload)
|
||||
if failure is not None:
|
||||
return self._fail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type=input_type,
|
||||
error=failure.message,
|
||||
is_unreachable=failure.is_unreachable,
|
||||
)
|
||||
|
||||
if parsed is None:
|
||||
return self._fail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type=input_type,
|
||||
error="empty response from Straiker",
|
||||
is_unreachable=False,
|
||||
)
|
||||
self._record(request_data=request_data, logging_obj=logging_obj, parsed=parsed)
|
||||
|
||||
if parsed.schema_version is not None and parsed.schema_version != STRAIKER_WEBHOOK_SCHEMA_VERSION:
|
||||
verbose_proxy_logger.warning(
|
||||
json.dumps(
|
||||
{
|
||||
"event": "straiker.schema_drift",
|
||||
"expected": STRAIKER_WEBHOOK_SCHEMA_VERSION,
|
||||
"received": parsed.schema_version,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
if parsed.action == "BLOCKED":
|
||||
self._block(
|
||||
request_data=request_data,
|
||||
input_type=input_type,
|
||||
message=parsed.blocked_reason or DEFAULT_BLOCK_MESSAGE,
|
||||
)
|
||||
if parsed.action == "GUARDRAIL_INTERVENED":
|
||||
is_streamed_response = input_type == "response" and _is_streamed_request(request_data)
|
||||
if parsed.texts is None or is_streamed_response:
|
||||
self._block(
|
||||
request_data=request_data,
|
||||
input_type=input_type,
|
||||
message=parsed.blocked_reason or DEFAULT_BLOCK_MESSAGE,
|
||||
)
|
||||
return self._intervened_inputs(inputs, parsed)
|
||||
return inputs
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from litellm.types.utils import ModelResponse
|
|||
|
||||
from .config import (
|
||||
DEFAULT_CODE_KEYWORDS,
|
||||
DEFAULT_ESCALATION_KEYWORDS,
|
||||
DEFAULT_REASONING_KEYWORDS,
|
||||
DEFAULT_SIMPLE_KEYWORDS,
|
||||
DEFAULT_TECHNICAL_KEYWORDS,
|
||||
|
|
@ -173,6 +174,11 @@ class ComplexityRouter(CustomLogger):
|
|||
self.config.custom_technical_keywords,
|
||||
)
|
||||
self.simple_keywords = self.config.simple_keywords or DEFAULT_SIMPLE_KEYWORDS
|
||||
self.escalation_keywords = (
|
||||
self.config.escalation_keywords
|
||||
if self.config.escalation_keywords is not None
|
||||
else DEFAULT_ESCALATION_KEYWORDS
|
||||
)
|
||||
|
||||
# Lazily built on first semantic request and cached for reuse (route
|
||||
# embeddings are static, only the prompt is embedded per request). The lock
|
||||
|
|
@ -668,6 +674,53 @@ class ComplexityRouter(CustomLogger):
|
|||
}
|
||||
return best_model
|
||||
|
||||
def _escalation_triggered(self, user_message: str) -> bool:
|
||||
"""Whether the prompt asks to escalate to a stronger model.
|
||||
|
||||
Matching is a case-sensitive substring test so the default "LITELLM ESCALATE"
|
||||
only fires on the deliberate, shouted form and not on incidental lowercase
|
||||
mentions of the word (e.g. "how do I escalate this ticket").
|
||||
"""
|
||||
if not self.escalation_keywords:
|
||||
return False
|
||||
return any(keyword in user_message for keyword in self.escalation_keywords)
|
||||
|
||||
def _tier_for_model(self, model: str) -> ComplexityTier | None:
|
||||
"""Return the most-severe configured tier whose pool contains this model."""
|
||||
pools = self._tier_pools()
|
||||
matched = tuple(ComplexityTier(tier_name) for tier_name, models in pools.items() if model in models)
|
||||
if not matched:
|
||||
return None
|
||||
return max(matched, key=TIER_SEVERITY_ORDER.index)
|
||||
|
||||
def _escalate_tier(self, tier: ComplexityTier) -> ComplexityTier:
|
||||
"""Bump a tier one step up to the next-higher configured tier.
|
||||
|
||||
Returns the input tier unchanged when it is already the highest configured
|
||||
tier, so escalation can never route below the model the user would otherwise
|
||||
have received.
|
||||
"""
|
||||
configured = frozenset(self.config.tiers)
|
||||
current_index = TIER_SEVERITY_ORDER.index(tier)
|
||||
higher_tiers = tuple(
|
||||
candidate for candidate in TIER_SEVERITY_ORDER[current_index + 1 :] if candidate.value in configured
|
||||
)
|
||||
return higher_tiers[0] if higher_tiers else tier
|
||||
|
||||
def _escalated_pin(self, pinned_model: str) -> str | None:
|
||||
"""Bump a session's pinned model to the next-higher configured tier.
|
||||
|
||||
Returns None when the pin no longer maps to any configured tier, signalling
|
||||
a full reclassification instead.
|
||||
"""
|
||||
pinned_tier = self._tier_for_model(pinned_model)
|
||||
if pinned_tier is None:
|
||||
return None
|
||||
escalated_tier = self._escalate_tier(pinned_tier)
|
||||
if escalated_tier == pinned_tier:
|
||||
return pinned_model
|
||||
return self.get_model_for_tier(escalated_tier)
|
||||
|
||||
def _lexical_tier_override(self, user_message: str) -> ComplexityTier | None:
|
||||
"""When keyword_tier_rules match literally, the most-severe matched tier wins.
|
||||
|
||||
|
|
@ -910,29 +963,41 @@ class ComplexityRouter(CustomLogger):
|
|||
if cache_key is not None:
|
||||
pinned_model = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
|
||||
if isinstance(pinned_model, str):
|
||||
# Refresh the TTL on every hit so an active session doesn't lose its
|
||||
# pin mid-conversation just because it outlives the original write.
|
||||
await self.litellm_router_instance.cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=pinned_model,
|
||||
ttl=self.config.session_affinity_ttl_seconds,
|
||||
)
|
||||
if self.config.adaptive:
|
||||
from litellm.router_strategy.adaptive_router.config import (
|
||||
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
|
||||
routed_model: str | None = pinned_model
|
||||
if self.escalation_keywords:
|
||||
resolved_messages = self._resolve_messages(messages, request_kwargs)
|
||||
user_message = (
|
||||
self._extract_user_message_and_system_prompt(resolved_messages)[0]
|
||||
if resolved_messages
|
||||
else None
|
||||
)
|
||||
if user_message is not None and self._escalation_triggered(user_message):
|
||||
routed_model = self._escalated_pin(pinned_model)
|
||||
if routed_model is not None:
|
||||
# Refresh the TTL on every hit so an active session doesn't lose its
|
||||
# pin mid-conversation just because it outlives the original write.
|
||||
await self.litellm_router_instance.cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=routed_model,
|
||||
ttl=self.config.session_affinity_ttl_seconds,
|
||||
)
|
||||
if self.config.adaptive:
|
||||
from litellm.router_strategy.adaptive_router.config import (
|
||||
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
|
||||
)
|
||||
|
||||
kwargs_metadata = request_kwargs.setdefault("metadata", {})
|
||||
if isinstance(kwargs_metadata, dict):
|
||||
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = pinned_model
|
||||
verbose_router_logger.info(
|
||||
f"ComplexityRouter: routing decision cause=session_affinity_pin, routed_model={pinned_model}"
|
||||
)
|
||||
has_original_messages = messages is not None and len(messages) > 0
|
||||
return PreRoutingHookResponse(
|
||||
model=pinned_model,
|
||||
messages=messages if has_original_messages else None,
|
||||
)
|
||||
kwargs_metadata = request_kwargs.setdefault("metadata", {})
|
||||
if isinstance(kwargs_metadata, dict):
|
||||
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = routed_model
|
||||
cause = "session_affinity_escalation" if routed_model != pinned_model else "session_affinity_pin"
|
||||
verbose_router_logger.info(
|
||||
f"ComplexityRouter: routing decision cause={cause}, routed_model={routed_model}"
|
||||
)
|
||||
has_original_messages = messages is not None and len(messages) > 0
|
||||
return PreRoutingHookResponse(
|
||||
model=routed_model,
|
||||
messages=messages if has_original_messages else None,
|
||||
)
|
||||
|
||||
response = await self._classify_and_route(
|
||||
model=model,
|
||||
|
|
@ -1004,13 +1069,17 @@ class ComplexityRouter(CustomLogger):
|
|||
messages=messages if has_original_messages else None,
|
||||
)
|
||||
|
||||
escalate = self._escalation_triggered(user_message)
|
||||
|
||||
override_tier = await self._resolve_keyword_tier_override(user_message, request_kwargs)
|
||||
if override_tier is not None:
|
||||
routed_model = await self._pick_model_for_tier(override_tier, messages, resolved_messages, request_kwargs)
|
||||
cause = "semantic_keyword_match" if self.config.semantic_keyword_matching else "literal_keyword_match"
|
||||
routed_tier = self._escalate_tier(override_tier) if escalate else override_tier
|
||||
routed_model = await self._pick_model_for_tier(routed_tier, messages, resolved_messages, request_kwargs)
|
||||
base_cause = "semantic_keyword_match" if self.config.semantic_keyword_matching else "literal_keyword_match"
|
||||
cause = f"{base_cause}+escalation" if escalate else base_cause
|
||||
verbose_router_logger.info(
|
||||
f"ComplexityRouter: routing decision cause={cause}, "
|
||||
f"tier={override_tier.value}, routed_model={routed_model}"
|
||||
f"tier={routed_tier.value}, routed_model={routed_model}"
|
||||
)
|
||||
return PreRoutingHookResponse(
|
||||
model=routed_model,
|
||||
|
|
@ -1018,6 +1087,9 @@ class ComplexityRouter(CustomLogger):
|
|||
)
|
||||
|
||||
tier, score, signals = await self.aclassify(user_message, system_prompt, request_kwargs)
|
||||
if escalate:
|
||||
tier = self._escalate_tier(tier)
|
||||
signals = [*signals, "escalation"]
|
||||
if self.config.adaptive:
|
||||
routed_model = self._soft_floor_pick(tier, user_message, request_kwargs)
|
||||
adaptive = self._ensure_adaptive_router()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -53,6 +53,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import (
|
|||
from litellm.types.proxy.guardrails.guardrail_hooks.cisco_ai_defense import (
|
||||
CiscoAIDefenseGuardrailConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.singulr import (
|
||||
SingulrGuardrailConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.headroom import (
|
||||
HeadroomGuardrailConfigModel,
|
||||
)
|
||||
|
|
@ -125,8 +128,10 @@ class SupportedGuardrailIntegrations(Enum):
|
|||
RUBRIK = "rubrik"
|
||||
VIGIL_GUARD = "vigil_guard"
|
||||
REPELLOAI = "repelloai"
|
||||
SINGULR = "singulr"
|
||||
HEADROOM = "headroom"
|
||||
COMPRESR = "compresr"
|
||||
STRAIKER = "straiker"
|
||||
|
||||
|
||||
class Role(Enum):
|
||||
|
|
@ -932,6 +937,7 @@ class LitellmParams(
|
|||
HiddenlayerGuardrailConfigModel,
|
||||
QostodianNexusConfigModel,
|
||||
VigilGuardGuardrailConfigModel,
|
||||
SingulrGuardrailConfigModel,
|
||||
):
|
||||
guardrail: str = Field(description="The type of guardrail integration to use")
|
||||
mode: Union[str, List[str], Mode] = Field(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
63
litellm/types/proxy/guardrails/guardrail_hooks/singulr.py
Normal file
63
litellm/types/proxy/guardrails/guardrail_hooks/singulr.py
Normal 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"
|
||||
169
litellm/types/proxy/guardrails/guardrail_hooks/straiker.py
Normal file
169
litellm/types/proxy/guardrails/guardrail_hooks/straiker.py
Normal file
|
|
@ -0,0 +1,169 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall
|
||||
|
||||
from .base import GuardrailConfigModel
|
||||
|
||||
StraikerWebhookEventType = Literal["pre_call", "post_call"]
|
||||
StraikerWebhookStreamPhase = Literal["none", "assembled"]
|
||||
StraikerWebhookAction = Literal["NONE", "BLOCKED", "GUARDRAIL_INTERVENED"]
|
||||
|
||||
STRAIKER_WEBHOOK_SCHEMA_VERSION = "1"
|
||||
|
||||
|
||||
class StraikerWebhookStream(BaseModel):
|
||||
phase: StraikerWebhookStreamPhase = "none"
|
||||
index: int | None = None
|
||||
|
||||
|
||||
class StraikerWebhookEvent(BaseModel):
|
||||
type: StraikerWebhookEventType
|
||||
id: str
|
||||
stream: StraikerWebhookStream = Field(default_factory=StraikerWebhookStream)
|
||||
|
||||
|
||||
class StraikerWebhookContent(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
texts: list[str] = Field(default_factory=list)
|
||||
images: list[str] = Field(default_factory=list)
|
||||
structured_messages: list[AllMessageValues] | None = None
|
||||
tools: list[dict[str, object]] | None = None
|
||||
tool_calls: list[ChatCompletionToolCallChunk] | list[ChatCompletionMessageToolCall] | None = None
|
||||
finish_reason: str | None = None
|
||||
|
||||
|
||||
class StraikerWebhookUsage(BaseModel):
|
||||
input_tokens: int | None = None
|
||||
output_tokens: int | None = None
|
||||
|
||||
|
||||
class StraikerWebhookContext(BaseModel):
|
||||
call_surface: str
|
||||
model: str | None = None
|
||||
model_provider: str | None = None
|
||||
destination: str | None = None
|
||||
session_id: str | None = None
|
||||
litellm_call_id: str | None = None
|
||||
litellm_trace_id: str | None = None
|
||||
litellm_version: str | None = None
|
||||
|
||||
|
||||
class StraikerWebhookIdentity(BaseModel):
|
||||
litellm_key: str | None = None
|
||||
litellm_team: str | None = None
|
||||
litellm_user_id: str | None = None
|
||||
litellm_user_email: str | None = None
|
||||
litellm_org_id: str | None = None
|
||||
end_user_id: str | None = None
|
||||
|
||||
|
||||
class StraikerWebhookApplication(BaseModel):
|
||||
source: str
|
||||
name: str | None = None
|
||||
|
||||
|
||||
class StraikerWebhookRequest(BaseModel):
|
||||
schema_version: str = STRAIKER_WEBHOOK_SCHEMA_VERSION
|
||||
event: StraikerWebhookEvent
|
||||
request: StraikerWebhookContent
|
||||
response: StraikerWebhookContent | None = None
|
||||
context: StraikerWebhookContext
|
||||
identity: StraikerWebhookIdentity
|
||||
application: StraikerWebhookApplication
|
||||
usage: StraikerWebhookUsage | None = None
|
||||
metadata: dict[str, object] | None = None
|
||||
|
||||
|
||||
class StraikerWebhookResponse(BaseModel):
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
action: StraikerWebhookAction = "NONE"
|
||||
blocked_reason: str | None = None
|
||||
texts: list[str] | None = None
|
||||
schema_version: str | None = None
|
||||
turn_id: str | None = Field(default=None, alias="turnId")
|
||||
|
||||
|
||||
class StraikerGuardrailConfigModelOptionalParams(BaseModel):
|
||||
timeout: float | None = Field(
|
||||
default=5.0,
|
||||
gt=0.0,
|
||||
description="Per-attempt HTTP timeout in seconds.",
|
||||
)
|
||||
max_retries: int | None = Field(
|
||||
default=2,
|
||||
ge=0,
|
||||
description="Retries on transient HTTP (408/429/5xx) and network errors.",
|
||||
)
|
||||
initial_backoff: float | None = Field(
|
||||
default=0.1,
|
||||
ge=0.0,
|
||||
description="Initial retry backoff in seconds.",
|
||||
)
|
||||
max_backoff: float | None = Field(
|
||||
default=2.0,
|
||||
ge=0.0,
|
||||
description="Maximum retry backoff in seconds.",
|
||||
)
|
||||
unreachable_fallback: Literal["fail_open", "fail_closed"] | None = Field(
|
||||
default="fail_closed",
|
||||
description="Behavior when Straiker is unreachable after retries.",
|
||||
)
|
||||
fail_on_error: bool | None = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"Behavior on any guardrail error, not just unreachability. True (default) blocks "
|
||||
"the request on error; False logs and allows the request to proceed."
|
||||
),
|
||||
)
|
||||
max_payload_bytes: int | None = Field(
|
||||
default=524288,
|
||||
gt=0,
|
||||
description="Maximum serialized webhook payload size sent to Straiker.",
|
||||
)
|
||||
custom_headers: dict[str, str] | None = Field(
|
||||
default=None,
|
||||
description="Additional HTTP headers sent to Straiker, excluding Authorization and the webhook-format header.",
|
||||
)
|
||||
metadata: dict[str, str] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Default metadata key/values added to the webhook metadata bag on every request. "
|
||||
"On key conflict with request-derived metadata, these configured values win."
|
||||
),
|
||||
)
|
||||
verbose: bool | None = Field(
|
||||
default=False,
|
||||
description="Log webhook request/response payloads and record action/turn_id in response hidden params.",
|
||||
)
|
||||
|
||||
|
||||
class StraikerGuardrailConfigModel(GuardrailConfigModel[StraikerGuardrailConfigModelOptionalParams]):
|
||||
api_key: str = Field(
|
||||
min_length=1,
|
||||
description="Straiker DefendAI environment API key (Bearer token). Env: STRAIKER_API_KEY.",
|
||||
json_schema_extra={"secret": True},
|
||||
)
|
||||
|
||||
api_base: str | None = Field(
|
||||
default="https://api.prod.straiker.ai",
|
||||
description="Straiker API base URL. Use the regional variant for non-US tenants.",
|
||||
)
|
||||
|
||||
default_app: str | None = Field(
|
||||
default="LiteLLM Gateway",
|
||||
description=(
|
||||
"Default application registered in the Straiker Defend Console. "
|
||||
"Overridden per-request by metadata.agent_id when present."
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "Straiker"
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -266,6 +266,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
"audio_transcription",
|
||||
"responses",
|
||||
"ocr",
|
||||
"realtime",
|
||||
]
|
||||
]
|
||||
tpm: Optional[int]
|
||||
|
|
@ -402,6 +403,11 @@ class CallTypes(str, Enum):
|
|||
vector_store_search = "vector_store_search"
|
||||
avector_store_search = "avector_store_search"
|
||||
|
||||
ingest = "ingest"
|
||||
aingest = "aingest"
|
||||
query = "query"
|
||||
aquery = "aquery"
|
||||
|
||||
#########################################################
|
||||
# Container Call Types
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -3196,6 +3196,12 @@ def get_optional_params_embeddings(
|
|||
non_default_params=non_default_params, optional_params={}, kwargs=kwargs
|
||||
)
|
||||
elif custom_llm_provider == "vertex_ai" or custom_llm_provider == "gemini":
|
||||
# OpenAI SDKs (and litellm's own client) send encoding_format="float"
|
||||
# by default; float lists are exactly what the vertex API returns, so
|
||||
# the param is a no-op — don't reject the provider default. Other
|
||||
# values (e.g. "base64") stay on the unsupported-param path below.
|
||||
if non_default_params.get("encoding_format") == "float":
|
||||
non_default_params.pop("encoding_format")
|
||||
supported_params = get_supported_openai_params(
|
||||
model=model,
|
||||
custom_llm_provider="vertex_ai",
|
||||
|
|
|
|||
|
|
@ -3451,7 +3451,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2.2e-05,
|
||||
"output_cost_per_token": 2.64e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -3470,7 +3470,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.00022,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -3489,7 +3489,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -4687,7 +4687,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -4707,7 +4707,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4739,7 +4739,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4771,7 +4771,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4832,7 +4832,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.0002,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -4850,7 +4850,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -7922,7 +7922,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2.2e-05,
|
||||
"output_cost_per_token": 2.64e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -7941,7 +7941,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.00022,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -7960,7 +7960,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -16272,7 +16272,7 @@
|
|||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/glm-5p2": {
|
||||
"cache_read_input_token_cost": 2.6e-07,
|
||||
"cache_read_input_token_cost": 1.4e-07,
|
||||
"input_cost_per_token": 1.4e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
|
|
@ -16686,7 +16686,7 @@
|
|||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/glm-5p2": {
|
||||
"cache_read_input_token_cost": 2.6e-07,
|
||||
"cache_read_input_token_cost": 1.4e-07,
|
||||
"input_cost_per_token": 1.4e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
|
|
@ -22169,7 +22169,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22188,7 +22188,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22282,7 +22282,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22300,7 +22300,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22318,7 +22318,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -24513,7 +24513,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24545,7 +24545,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24577,7 +24577,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24610,7 +24610,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 2.4e-05,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
|
|
@ -24645,7 +24645,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
|
|
@ -24678,7 +24678,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24710,7 +24710,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -43694,7 +43694,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -43727,7 +43727,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
|
|||
|
|
@ -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
28
router_plugins.json
Normal 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"]
|
||||
}
|
||||
]
|
||||
|
|
@ -13,6 +13,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
|
|||
- `realtime/` - realtime websocket sessions, including the pipecat audio path
|
||||
- `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `ratelimit/` (rpm/tpm blocks, window reset, pacing headers on live traffic), `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window), and `spend_tracking/` (spend logging and cost attribution on `/spend/*`)
|
||||
- `management/` - key/team/user/organization management routes: create/update/delete persistence via the info routes, team membership, and llm-only-key route denials; also the dashboard UI behavior on top of them, driven through the proxy-served UI at /ui with playwright (optional dep behind importorskip)
|
||||
- `mcp/` - the MCP server surface over api_key auth: an admin registers an upstream MCP server through the management API and grants keys access via `object_permission.mcp_servers`, then the suite asserts tool listing and calling honor that permission (a key without the grant sees none of the server's tools and is refused a `tools/call` with a 403)
|
||||
- `logging/` - logging-integration delivery (datadog and friends)
|
||||
- `security/` - secret handling and log-leak protection
|
||||
- `router/` - routing and reliability behavior (fallbacks, cooldowns)
|
||||
|
|
@ -51,7 +52,7 @@ The shape is layered so tests stay declarative
|
|||
|
||||
Each suite provides its own `client` fixture (see `llm_translation/passthrough_client.py`), a frozen dataclass that holds the shared `Gateway` and adds suite-specific routes. Cleanup runs through that same `Gateway`, so whatever keys or customers your test creates get torn down by the `resources` fixture
|
||||
|
||||
Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The skip-vs-fail split is deliberate: a test marked `e2e` skips when no proxy answers its liveness probe, but once a request reaches the proxy any wrong behavior is a hard failure, never a skip
|
||||
Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The harness hard-fails and never skips: a test marked `e2e` fails when no proxy answers its liveness probe, and once a request reaches the proxy any wrong behavior is likewise a hard failure, so a missing proxy turns the run red instead of being mistaken for a pass
|
||||
|
||||
Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure coverage of the harness itself carries no marker and runs regardless. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache
|
||||
|
||||
|
|
@ -131,7 +132,7 @@ Quota Management - behavior features (entity- or config-driven caps and their ac
|
|||
quota_management.<behavior>.<variant>.<assertion>
|
||||
behavior : ratelimit | budget | spend_tracking
|
||||
variant : <ratelimit> rpm | tpm | priority_generous | priority_strict
|
||||
<budget> key | internal_user | end_user | organization | team_member | tag
|
||||
<budget> key | internal_user | end_user | organization | team | team_member | tag
|
||||
| model_max | soft | key_multi_window | team_multi_window
|
||||
| fallback | spend_counter
|
||||
<spend_tracking> chat_completions | stream | embeddings | cache_hit | key_rollup
|
||||
|
|
@ -139,7 +140,8 @@ quota_management.<behavior>.<variant>.<assertion>
|
|||
| spend_calculate | pagination
|
||||
assertion : blocks_over_limit | resets_after_window | headers_report_remaining | picks_under_tpm
|
||||
| blocks_then_resets | resets_windows_independently | alerts_without_blocking
|
||||
| isolates_per_model | routes_to_fallback | reseed_matches_db | logs_cost | zero_cost
|
||||
| isolates_per_model | isolates_per_member | enforced_across_keys | routes_to_fallback
|
||||
| reseed_matches_db | logs_cost | zero_cost
|
||||
| matches_sum_of_logs | loses_no_spend | attributes_spend | writes_own_rows
|
||||
| writes_failure_row | returns_cost | keeps_total
|
||||
e.g. quota_management.ratelimit.rpm.blocks_over_limit exercised_on=[chat_completions, messages]
|
||||
|
|
@ -173,7 +175,7 @@ other.<area>.<case>.<assertion>
|
|||
```
|
||||
|
||||
## Hard Rules
|
||||
- no monkeypatching, mock tests or unit tests of any kind. if a contributor asks you to write an end to end test, do NOT stage a unit test with it. if you find a product gap, call it out in the PR description
|
||||
- no monkeypatching or mock tests, and never substitute a unit test for e2e feature coverage: a product feature is proven end to end against a live proxy, not with a unit test. if a contributor asks you to write an end to end test, do NOT stage a unit test of the feature with it; if you find a product gap, call it out in the PR description. tests that cover the harness itself are the exception and are allowed (for example `coverage_registry/test_collector.py`, which unit-tests the coverage collector): they carry no `e2e` marker, exercise harness plumbing rather than a product feature, and run whether or not a proxy is up
|
||||
|
||||
- use model management endpoints to create new models for a test. this could be in a conftest / inline for each test. ask the user what they want.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(...)`.
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
@ -16,12 +16,12 @@ cover "OpenAI plus the big three clouds":
|
|||
carries only the open-weight
|
||||
gpt-oss MaaS models
|
||||
|
||||
The openai and azure_openai columns run unconditionally, like every
|
||||
other live column: the environments that run the suite carry
|
||||
`OPENAI_API_KEY` and `AZURE_API_BASE` + `AZURE_API_KEY` pointing at a
|
||||
resource with gpt-5.6 deployments. The bedrock_mantle column is
|
||||
opt-in via `COMPAT_MANTLE_CELLS=1` because the AWS account is still
|
||||
waiting on the Bedrock Mantle allowlist for the `openai.gpt-5.6-*`
|
||||
The azure_openai column runs unconditionally when Azure gpt-5.6
|
||||
deployments exist. The openai column is opt-in via
|
||||
`COMPAT_OPENAI_GPT_CELLS=1` because under the full stage suite those
|
||||
cells routinely burn minutes on Claude CLI timeouts. The bedrock_mantle
|
||||
column is opt-in via `COMPAT_MANTLE_CELLS=1` because the AWS account is
|
||||
still waiting on the Bedrock Mantle allowlist for the `openai.gpt-5.6-*`
|
||||
models; until the flag is set each Mantle cell skips and its matrix
|
||||
cell publishes as `not_tested` instead of a credential-shaped red.
|
||||
The `vertex_ai_gpt` column needs no flag either way: its cells report
|
||||
|
|
@ -35,6 +35,7 @@ import os
|
|||
import pytest
|
||||
|
||||
MANTLE_CELLS_ENV = "COMPAT_MANTLE_CELLS"
|
||||
OPENAI_GPT_CELLS_ENV = "COMPAT_OPENAI_GPT_CELLS"
|
||||
|
||||
VERTEX_AI_GPT_NOT_APPLICABLE_REASON = (
|
||||
"GCP Vertex AI does not offer OpenAI's closed-weight GPT-5.6 family "
|
||||
|
|
@ -59,3 +60,19 @@ def skip_unless_mantle_cells_enabled() -> None:
|
|||
f"Bedrock Mantle GPT-5.6 cells are opt-in; set {MANTLE_CELLS_ENV}=1 "
|
||||
"once the AWS account is allowlisted for the openai.gpt-5.6-* models"
|
||||
)
|
||||
|
||||
|
||||
def skip_unless_openai_gpt_cells_enabled() -> None:
|
||||
"""Skip OpenAI GPT-5.6 columns unless `COMPAT_OPENAI_GPT_CELLS` opts them in.
|
||||
|
||||
Under the full stage suite these cells routinely hit 120s Claude CLI
|
||||
timeouts and rate-limit-shaped retries across Sol/Terra/Luna, burning
|
||||
~8+ minutes per cell without a stable green. Opt in when exercising
|
||||
the OpenAI GPT translation path in isolation.
|
||||
"""
|
||||
if os.environ.get(OPENAI_GPT_CELLS_ENV, "").strip().lower() in {"1", "true", "yes"}:
|
||||
return
|
||||
pytest.skip(
|
||||
f"OpenAI GPT-5.6 cells are opt-in; set {OPENAI_GPT_CELLS_ENV}=1 "
|
||||
"to run them (stage suite timeouts under concurrent load)"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ green if all three pass.
|
|||
from __future__ import annotations
|
||||
|
||||
from claude_code._basic_messaging import run_basic_messaging_cell
|
||||
from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled
|
||||
|
||||
OPENAI_MODELS = [
|
||||
"gpt-5-6-sol-openai",
|
||||
|
|
@ -34,6 +35,7 @@ OPENAI_MODELS = [
|
|||
def test_basic_messaging_non_streaming_openai(compat_result):
|
||||
"""Drive the `claude` CLI against the LiteLLM proxy and assert a
|
||||
non-empty reply from each GPT-5.6 tier."""
|
||||
skip_unless_openai_gpt_cells_enabled()
|
||||
run_basic_messaging_cell(
|
||||
compat_result=compat_result,
|
||||
models=OPENAI_MODELS,
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ green if all three pass.
|
|||
from __future__ import annotations
|
||||
|
||||
from claude_code._basic_messaging import run_basic_messaging_cell
|
||||
from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled
|
||||
|
||||
OPENAI_MODELS = [
|
||||
"gpt-5-6-sol-openai",
|
||||
|
|
@ -36,6 +37,7 @@ OPENAI_MODELS = [
|
|||
def test_basic_messaging_streaming_openai(compat_result):
|
||||
"""Drive the `claude` CLI against the LiteLLM proxy and assert a
|
||||
non-empty streamed reply from each GPT-5.6 tier."""
|
||||
skip_unless_openai_gpt_cells_enabled()
|
||||
run_basic_messaging_cell(
|
||||
compat_result=compat_result,
|
||||
models=OPENAI_MODELS,
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ from typing import Any, Mapping, Sequence
|
|||
import pytest
|
||||
|
||||
from claude_code._env import require_proxy
|
||||
from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled
|
||||
from claude_code.cli_driver import (
|
||||
ClaudeCLIError,
|
||||
failure_diagnostic,
|
||||
|
|
@ -69,6 +70,7 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool:
|
|||
def test_tool_use_openai(compat_result):
|
||||
"""Drive the `claude` CLI against the LiteLLM proxy and assert a
|
||||
tool call was emitted on the wire by each GPT-5.6 tier."""
|
||||
skip_unless_openai_gpt_cells_enabled()
|
||||
proxy = require_proxy(compat_result)
|
||||
|
||||
outcomes = run_claude_models_parallel(
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ from typing import Any, Mapping, Sequence
|
|||
import pytest
|
||||
|
||||
from claude_code._env import require_proxy
|
||||
from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled
|
||||
from claude_code.cli_driver import (
|
||||
ClaudeCLIError,
|
||||
failure_diagnostic,
|
||||
|
|
@ -87,6 +88,7 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int:
|
|||
|
||||
|
||||
def test_tool_use_streaming_openai(compat_result):
|
||||
skip_unless_openai_gpt_cells_enabled()
|
||||
proxy = require_proxy(compat_result)
|
||||
|
||||
outcomes = run_claude_models_parallel(
|
||||
|
|
|
|||
|
|
@ -15,14 +15,14 @@ shared fixtures build on it.
|
|||
|
||||
import functools
|
||||
import sys
|
||||
from collections.abc import Generator, Iterator
|
||||
from collections.abc import Iterator
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
from e2e_config import CONTROL_PLANE_BASE_URL, PROXY_BASE_URL
|
||||
from e2e_result_reporter import covers_from_item, format_e2e_result_line, result_from_pytest
|
||||
from junit_properties import attach_result_properties
|
||||
from lifecycle import GatewayProvider, ResourceManager
|
||||
|
||||
|
||||
|
|
@ -40,6 +40,17 @@ def pytest_configure(config: pytest.Config) -> None:
|
|||
)
|
||||
|
||||
|
||||
def pytest_collection_modifyitems(items: list[pytest.Item]) -> None:
|
||||
"""Attach the two custom signals (suite package and covered cell ids) to every
|
||||
test's user_properties so the standard JUnit report (`--junitxml`) records them
|
||||
as `<property>` entries, on every outcome including skips and setup errors.
|
||||
Downstream (Loki/Grafana) reads outcome and duration from the standard report
|
||||
and these properties for package rollups and coverage drill-down. See
|
||||
junit_properties.py."""
|
||||
for item in items:
|
||||
attach_result_properties(item)
|
||||
|
||||
|
||||
def _liveness_reason(label: str, base_url: str) -> str | None:
|
||||
"""None if `base_url` answers its liveness probe, else a failure reason."""
|
||||
try:
|
||||
|
|
@ -86,30 +97,6 @@ def pytest_runtest_call(item: pytest.Item) -> None:
|
|||
item.session.stash[_E2E_TEST_RAN] = True
|
||||
|
||||
|
||||
@pytest.hookimpl(wrapper=True, tryfirst=True)
|
||||
def pytest_runtest_makereport(
|
||||
item: pytest.Item, call: pytest.CallInfo[object]
|
||||
) -> Generator[None, pytest.TestReport, pytest.TestReport]:
|
||||
"""Emit one structured E2E_RESULT line per finished test for Loki/Grafana.
|
||||
|
||||
Status-history panels should aggregate by package (and optional covers), not
|
||||
scrape pytest progress basenames. See e2e_result_reporter.py.
|
||||
"""
|
||||
report = yield
|
||||
result = result_from_pytest(
|
||||
nodeid=str(report.nodeid),
|
||||
when=str(report.when),
|
||||
failed=bool(report.failed),
|
||||
skipped=bool(report.skipped),
|
||||
passed=bool(report.passed),
|
||||
duration_seconds=float(report.duration),
|
||||
covers=covers_from_item(item),
|
||||
)
|
||||
if result is not None:
|
||||
print(format_e2e_result_line(result), flush=True)
|
||||
return report
|
||||
|
||||
|
||||
def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
|
||||
"""Once the whole e2e session is done (all suites), truncate the spend logs so
|
||||
the DB doesn't accumulate test rows. Sessions where no e2e test body ran leave
|
||||
|
|
@ -132,13 +119,6 @@ def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
|
|||
if spend_dir in sys.path:
|
||||
sys.path.remove(spend_dir)
|
||||
|
||||
try:
|
||||
from bob_the_builder import remediate
|
||||
|
||||
remediate(session)
|
||||
except Exception as exc: # noqa: BLE001 - remediation is best-effort
|
||||
print(f"devin remediation best-effort failed: {exc}")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def resources(client: GatewayProvider) -> Iterator[ResourceManager]:
|
||||
|
|
|
|||
|
|
@ -7,14 +7,20 @@
|
|||
- {id: quota_management.ratelimit.priority_generous.picks_under_tpm, module: quota_management, tier: P1, behavior: ratelimit, variant: priority_generous, assertions: [picks_under_tpm], exercised_on: [chat_completions, messages], source: "dynamic_rate_limiter_v3.py:36-52", rationale: "Generous mode (<80% sat) allows priority borrowing"}
|
||||
- {id: quota_management.ratelimit.priority_strict.picks_under_tpm, module: quota_management, tier: P1, behavior: ratelimit, variant: priority_strict, assertions: [picks_under_tpm], exercised_on: [chat_completions, messages], source: "dynamic_rate_limiter_v3.py:53-71", rationale: "Strict mode (>=80% sat) enforces priority fairness"}
|
||||
- {id: quota_management.budget.key.blocks_over_limit, module: quota_management, tier: P0, behavior: budget, variant: key, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "A key's max_budget blocks further paid calls once spend crosses it"}
|
||||
- {id: quota_management.budget.team.blocks_over_limit, module: quota_management, tier: P0, behavior: budget, variant: team, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "A team's max_budget blocks every key on the team once combined spend crosses it, including keys that spent nothing themselves"}
|
||||
- {id: quota_management.budget.internal_user.blocks_over_limit, module: quota_management, tier: P1, behavior: budget, variant: internal_user, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "An internal user's max_budget governs personal keys"}
|
||||
- {id: quota_management.budget.internal_user.enforced_across_keys, module: quota_management, tier: P1, behavior: budget, variant: internal_user, assertions: [enforced_across_keys], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "An internal user's max_budget governs every personal key it owns; a second untouched key is blocked once the shared user budget is exhausted"}
|
||||
- {id: quota_management.budget.end_user.blocks_over_limit, module: quota_management, tier: P1, behavior: budget, variant: end_user, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "A customer (end-user) max_budget blocks calls attributed via user="}
|
||||
- {id: quota_management.budget.organization.blocks_over_limit, module: quota_management, tier: P1, behavior: budget, variant: organization, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "An organization's max_budget blocks keys under its teams"}
|
||||
- {id: quota_management.budget.team_member.blocks_over_limit, module: quota_management, tier: P1, behavior: budget, variant: team_member, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "A member's per-team budget blocks independently of the team budget"}
|
||||
- {id: quota_management.budget.team_member.isolates_per_member, module: quota_management, tier: P1, behavior: budget, variant: team_member, assertions: [isolates_per_member], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "One team member's exhausted per-team budget does not block a different member on the same team"}
|
||||
- {id: quota_management.budget.tag.blocks_over_limit, module: quota_management, tier: P1, behavior: budget, variant: tag, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "router_strategy/budget_limiter.py", rationale: "Proxy-level tag budgets block tagged requests at the cap"}
|
||||
- {id: quota_management.budget.model_max.isolates_per_model, module: quota_management, tier: P1, behavior: budget, variant: model_max, assertions: [isolates_per_model], exercised_on: [chat_completions], source: "proxy/hooks/model_max_budget_limiter.py", rationale: "model_max_budget caps one model without touching a sibling's budget"}
|
||||
- {id: quota_management.budget.soft.alerts_without_blocking, module: quota_management, tier: P1, behavior: budget, variant: soft, assertions: [alerts_without_blocking], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "soft_budget alerts but never blocks traffic"}
|
||||
- {id: quota_management.budget.key.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: key, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_duration zeroes key spend after the window; a blocked key serves again"}
|
||||
- {id: quota_management.budget.team.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: team, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_duration zeroes a team's spend after the window; every key on the team serves again"}
|
||||
- {id: quota_management.budget.organization.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: organization, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "An org budget resets after its window; keys under the org serve again"}
|
||||
- {id: quota_management.budget.internal_user.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: internal_user, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "An internal user's budget resets after its window; their personal and team-member keys serve again"}
|
||||
- {id: quota_management.budget.team_member.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: team_member, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "Member per-team budget reset keeps advancing window after window"}
|
||||
- {id: quota_management.budget.key_multi_window.blocks_then_resets, module: quota_management, tier: P1, behavior: budget, variant: key_multi_window, assertions: [blocks_then_resets], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_limits enforce within a short window and serve again in the next"}
|
||||
- {id: quota_management.budget.key_multi_window.resets_windows_independently, module: quota_management, tier: P2, behavior: budget, variant: key_multi_window, assertions: [resets_windows_independently], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "Each window of a multi-window budget resets on its own schedule"}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
# local setup to run e2e tests
|
||||
configs:
|
||||
mcp_upstream_server:
|
||||
file: ../mcp_tests/mcp_e2e_upstream_server.py
|
||||
litellm_config:
|
||||
content: |
|
||||
general_settings:
|
||||
|
|
@ -131,7 +133,27 @@ services:
|
|||
target: /app/config.yaml
|
||||
command: ["--config", "/app/config.yaml", "--port", "4000"]
|
||||
|
||||
# throwaway db
|
||||
# deterministic self-hosted upstream MCP server (FastMCP add/multiply over
|
||||
# streamable-http), reachable by the litellm container at mcp-upstream:8090/mcp.
|
||||
# Not a depends_on of litellm on purpose: only the mcp suite needs it, and it
|
||||
# boots long before the proxy is live, so it must not gate the other suites'
|
||||
# stack. The suite registers it through /v1/mcp/server at test time.
|
||||
mcp-upstream:
|
||||
image: ghcr.io/berriai/litellm:main-latest
|
||||
entrypoint: ["python3", "/app/mcp_upstream_server.py"]
|
||||
environment:
|
||||
MCP_HOST: 0.0.0.0
|
||||
MCP_PORT: "8090"
|
||||
configs:
|
||||
- source: mcp_upstream_server
|
||||
target: /app/mcp_upstream_server.py
|
||||
healthcheck:
|
||||
test: ["CMD", "python3", "-c", "import socket; socket.create_connection(('127.0.0.1', 8090), 2).close()"]
|
||||
interval: 3s
|
||||
timeout: 3s
|
||||
retries: 40
|
||||
|
||||
# throwaway db
|
||||
db:
|
||||
image: postgres:16
|
||||
environment:
|
||||
|
|
|
|||
|
|
@ -1,144 +0,0 @@
|
|||
"""Structured e2e result lines for Loki / Grafana status history.
|
||||
|
||||
Pytest progress lines are a bad dashboard source: they only expose file basenames,
|
||||
break under quiet modes, and force status-history rows to explode with suite growth.
|
||||
|
||||
Each finished test emits one logfmt line:
|
||||
|
||||
E2E_RESULT package=logging file=test_langfuse_e2e.py outcome=failed
|
||||
duration_ms=1234 node_id=logging/test_langfuse_e2e.py::TestX::test_y
|
||||
covers=logging.langfuse.team.success
|
||||
|
||||
Grafana package status-history queries max(fail) by package over E2E_RESULT lines.
|
||||
Drill-down uses node_id / covers in Explore, not status-history cardinality.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable, Sequence
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Literal, Protocol, runtime_checkable
|
||||
|
||||
Outcome = Literal["passed", "failed", "error", "skipped"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class E2EResult:
|
||||
package: str
|
||||
file: str
|
||||
outcome: Outcome
|
||||
duration_ms: int
|
||||
node_id: str
|
||||
covers: tuple[str, ...]
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _MarkerArgs(Protocol):
|
||||
args: Sequence[object]
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _ItemWithCovers(Protocol):
|
||||
def iter_markers(self, name: str) -> Iterable[object]: ...
|
||||
|
||||
|
||||
def package_from_nodeid(nodeid: str) -> str:
|
||||
"""Top-level suite package under tests/e2e/, or 'root' for top-level files.
|
||||
|
||||
Pytest nodeids are relative to the invocation cwd. Repo-root runs look like
|
||||
`tests/e2e/logging/...`; suite-cwd runs look like `logging/...`. Strip the
|
||||
`tests/e2e` prefix so package is the suite dir either way.
|
||||
"""
|
||||
path_part = nodeid.split("::", 1)[0].replace("\\", "/")
|
||||
parts = tuple(p for p in path_part.split("/") if p and p != ".")
|
||||
if len(parts) >= 3 and parts[0] == "tests" and parts[1] == "e2e":
|
||||
parts = parts[2:]
|
||||
if len(parts) <= 1:
|
||||
return "root"
|
||||
return parts[0]
|
||||
|
||||
|
||||
def file_from_nodeid(nodeid: str) -> str:
|
||||
path_part = nodeid.split("::", 1)[0].replace("\\", "/")
|
||||
return Path(path_part).name
|
||||
|
||||
|
||||
def covers_from_item(item: object) -> tuple[str, ...]:
|
||||
"""Read @pytest.mark.covers cell ids from a pytest Item."""
|
||||
if not isinstance(item, _ItemWithCovers):
|
||||
return ()
|
||||
return tuple(
|
||||
dict.fromkeys(
|
||||
arg
|
||||
for marker in item.iter_markers(name="covers")
|
||||
if isinstance(marker, _MarkerArgs)
|
||||
for arg in marker.args
|
||||
if isinstance(arg, str) and arg
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def outcome_from_report(when: str, failed: bool, skipped: bool, passed: bool) -> Outcome | None:
|
||||
"""Map pytest TestReport fields to a terminal outcome. None if not final."""
|
||||
if when == "setup" and skipped:
|
||||
return "skipped"
|
||||
if when == "setup" and failed:
|
||||
return "error"
|
||||
if when != "call":
|
||||
return None
|
||||
if skipped:
|
||||
return "skipped"
|
||||
if failed:
|
||||
return "failed"
|
||||
if passed:
|
||||
return "passed"
|
||||
return "failed"
|
||||
|
||||
|
||||
def _logfmt_escape(value: str) -> str:
|
||||
if value == "":
|
||||
return '""'
|
||||
needs_quote = any(ch.isspace() or ch in "\"=\\" for ch in value)
|
||||
if not needs_quote:
|
||||
return value
|
||||
escaped = value.replace("\\", "\\\\").replace('"', '\\"')
|
||||
return f'"{escaped}"'
|
||||
|
||||
|
||||
def format_e2e_result_line(result: E2EResult) -> str:
|
||||
covers = ",".join(result.covers)
|
||||
fields = (
|
||||
("package", result.package),
|
||||
("file", result.file),
|
||||
("outcome", result.outcome),
|
||||
("duration_ms", str(result.duration_ms)),
|
||||
("node_id", result.node_id),
|
||||
("covers", covers),
|
||||
)
|
||||
body = " ".join(f"{key}={_logfmt_escape(value)}" for key, value in fields)
|
||||
return f"E2E_RESULT {body}"
|
||||
|
||||
|
||||
def result_from_pytest(
|
||||
*,
|
||||
nodeid: str,
|
||||
when: str,
|
||||
failed: bool,
|
||||
skipped: bool,
|
||||
passed: bool,
|
||||
duration_seconds: float,
|
||||
covers: tuple[str, ...] = (),
|
||||
) -> E2EResult | None:
|
||||
outcome = outcome_from_report(when=when, failed=failed, skipped=skipped, passed=passed)
|
||||
if outcome is None:
|
||||
return None
|
||||
duration_ms = max(0, int(round(duration_seconds * 1000)))
|
||||
return E2EResult(
|
||||
package=package_from_nodeid(nodeid),
|
||||
file=file_from_nodeid(nodeid),
|
||||
outcome=outcome,
|
||||
duration_ms=duration_ms,
|
||||
node_id=nodeid,
|
||||
covers=covers,
|
||||
)
|
||||
|
|
@ -1,66 +0,0 @@
|
|||
# Grafana: package status history for e2e
|
||||
|
||||
Dashboard: [LiteLLM E2E](https://berriai.grafana.net/d/mup2cfn/litellm-e2e) (`mup2cfn`).
|
||||
|
||||
The old **test suite status history** panel scraped pytest progress lines and
|
||||
grouped by **file basename** (`test_foo.py`). That does not scale: multi-class
|
||||
files collapse to one bit, and full `node_id` cardinality melts status-history.
|
||||
|
||||
## Emitter
|
||||
|
||||
After each test finishes, `tests/e2e/conftest.py` prints one logfmt line:
|
||||
|
||||
```
|
||||
E2E_RESULT package=logging file=test_langfuse_e2e.py outcome=failed duration_ms=1500 node_id="logging/..." covers=cell.id
|
||||
```
|
||||
|
||||
## Panel: package status history (replace panel 11)
|
||||
|
||||
**Type:** Status history
|
||||
**Interval:** 15m (or 1h for multi-day ranges)
|
||||
**Description:** Per top-level package under `tests/e2e/`: red if any test failed or errored in the bucket.
|
||||
|
||||
```logql
|
||||
max by (package) (
|
||||
max_over_time(
|
||||
{service_name="litellm-e2e"}
|
||||
|= "E2E_RESULT"
|
||||
| logfmt
|
||||
| outcome != ""
|
||||
| label_format result=`{{ if or (eq .outcome "failed") (eq .outcome "error") }}1{{ else }}0{{ end }}`
|
||||
| unwrap result
|
||||
[$__interval]
|
||||
)
|
||||
)
|
||||
```
|
||||
|
||||
Value mappings: `0` → Pass (green), `1` → Fail (red).
|
||||
|
||||
If `service_name` is missing on older scrapes, use:
|
||||
|
||||
```logql
|
||||
{cluster="berrie-litellm-stage", pod=~"litellm-e2e-.+"}
|
||||
```
|
||||
|
||||
instead of `{service_name="litellm-e2e"}`.
|
||||
|
||||
## Panel: failed tests (logs drill-down)
|
||||
|
||||
```logql
|
||||
{service_name="litellm-e2e"} |= "E2E_RESULT" | logfmt | outcome=~"failed|error"
|
||||
```
|
||||
|
||||
Show fields: `package`, `file`, `node_id`, `covers`, `duration_ms`.
|
||||
|
||||
## Panel (optional): filter by package variable
|
||||
|
||||
Dashboard variable `package` (custom or from label_values on E2E_RESULT):
|
||||
|
||||
```logql
|
||||
{service_name="litellm-e2e"} |= "E2E_RESULT" | logfmt | package=`$package` | outcome=~"failed|error"
|
||||
```
|
||||
|
||||
## Do not
|
||||
|
||||
- Put full `node_id` as the status-history series key (cardinality).
|
||||
- Rely on `::S+ PASSED` progress regex as the primary signal once E2E_RESULT is live.
|
||||
59
tests/e2e/junit_properties.py
Normal file
59
tests/e2e/junit_properties.py
Normal file
|
|
@ -0,0 +1,59 @@
|
|||
"""Custom per-test signals for the standard JUnit reporter.
|
||||
|
||||
The e2e suite ships results to Loki/Grafana from a standard pytest JUnit report
|
||||
(`--junitxml=e2e-report.xml`), not a bespoke log line. JUnit already records
|
||||
outcome, duration, and node id for every `<testcase>`; the only signals it cannot
|
||||
derive on its own are the normalized suite package and the coverage-registry cell
|
||||
ids a test covers. Those ride along as JUnit `<property>` entries via each item's
|
||||
`user_properties`, attached in `conftest.py::pytest_collection_modifyitems`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def package_from_nodeid(nodeid: str) -> str:
|
||||
"""Top-level suite package under tests/e2e/, or 'root' for top-level files.
|
||||
|
||||
Pytest nodeids are relative to the invocation cwd. Repo-root runs look like
|
||||
`tests/e2e/logging/...`; suite-cwd runs look like `logging/...`. Strip the
|
||||
`tests/e2e` prefix so package is the suite dir either way.
|
||||
"""
|
||||
path_part = nodeid.split("::", 1)[0].replace("\\", "/")
|
||||
raw = tuple(p for p in path_part.split("/") if p and p != ".")
|
||||
parts = raw[2:] if len(raw) >= 3 and raw[0] == "tests" and raw[1] == "e2e" else raw
|
||||
if len(parts) <= 1:
|
||||
return "root"
|
||||
return parts[0]
|
||||
|
||||
|
||||
def dedupe_covers(marker_args: Iterable[tuple[object, ...]]) -> tuple[str, ...]:
|
||||
"""Flatten @pytest.mark.covers arg lists into unique, order-preserving cell
|
||||
ids, dropping anything that is not a non-empty string."""
|
||||
return tuple(dict.fromkeys(arg for args in marker_args for arg in args if isinstance(arg, str) and arg))
|
||||
|
||||
|
||||
def covers_from_item(item: pytest.Item) -> tuple[str, ...]:
|
||||
"""Read @pytest.mark.covers cell ids off a pytest Item, order-preserving."""
|
||||
return dedupe_covers(marker.args for marker in item.iter_markers(name="covers"))
|
||||
|
||||
|
||||
def result_properties(item: pytest.Item) -> tuple[tuple[str, str], ...]:
|
||||
"""The custom signals a standard reporter cannot derive: the normalized suite
|
||||
package and the comma-joined coverage-registry cell ids this test covers."""
|
||||
return (
|
||||
("package", package_from_nodeid(item.nodeid)),
|
||||
("covers", ",".join(covers_from_item(item))),
|
||||
)
|
||||
|
||||
|
||||
def attach_result_properties(item: pytest.Item) -> None:
|
||||
"""Attach result_properties to an item's user_properties, idempotently: a
|
||||
second call is a no-op, so a collection that runs the hook more than once
|
||||
never emits duplicate <property> entries."""
|
||||
if any(name == "package" for name, _ in item.user_properties):
|
||||
return
|
||||
item.user_properties.extend(result_properties(item))
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
"""LLM-translation suite's `client` fixture.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. PassthroughClient holds the shared
|
||||
Gateway, so the `resources` fixture cleans up keys this suite creates.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -49,10 +49,10 @@ kept commented out in `PROVIDERS` until they pass end-to-end here; re-enable the
|
|||
uncommenting their entry.
|
||||
|
||||
Every provider is provisioned and asserted; the suite never skips a provider. Per
|
||||
`tests/e2e/CLAUDE.md` the only sanctioned skip is the whole-suite proxy-liveness
|
||||
skip, so a provider whose credentials or upstream realtime model are missing on the
|
||||
gateway is a hard failure, not a skip. Give the gateway each provider's credentials
|
||||
to turn its tests green.
|
||||
`tests/e2e/CLAUDE.md` there is no sanctioned skip: the whole-suite proxy-liveness
|
||||
probe hard-fails when no proxy answers, and a provider whose credentials or upstream
|
||||
realtime model are missing on the gateway is likewise a hard failure, not a skip.
|
||||
Give the gateway each provider's credentials to turn its tests green.
|
||||
|
||||
## Running
|
||||
|
||||
|
|
@ -63,5 +63,5 @@ the deployments itself), then
|
|||
uv run pytest tests/e2e/llm_translation/realtime/ -v
|
||||
```
|
||||
|
||||
The whole suite skips only when no proxy answers `GET /health/liveliness` at
|
||||
The whole suite hard-fails at setup when no proxy answers `GET /health/liveliness` at
|
||||
`LITELLM_PROXY_URL` (default `http://localhost:4000`).
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
"""Realtime suite's `client` and `realtime_models` fixtures.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. RealtimeClient holds the shared Gateway,
|
||||
so the `resources` fixture cleans up keys this suite creates.
|
||||
|
||||
|
|
|
|||
|
|
@ -6,10 +6,10 @@ schema: the session lifecycle, the canonical response event sequence with a
|
|||
reconstructed transcript and usage, and a full tool-call round-trip (call ->
|
||||
tool result -> a follow-up response that uses the result).
|
||||
|
||||
One GA-speaking client validates every provider; only the model alias changes. A
|
||||
provider whose realtime alias is not configured on the proxy skips (skip on
|
||||
environment); once it is configured, a protocol failure is a hard failure. See
|
||||
REALTIME_COVERAGE_MATRIX.md.
|
||||
One GA-speaking client validates every provider; only the model alias changes.
|
||||
Every provider is provisioned at session start, so a missing realtime alias is a
|
||||
hard failure, not a skip; once configured, a protocol failure is likewise a hard
|
||||
failure. See REALTIME_COVERAGE_MATRIX.md.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
|
|
|||
|
|
@ -7,9 +7,9 @@ references the proxy resolves at call time, so adding a provider is a new type
|
|||
rather than another inline body. Start the proxy with the Rust OCR path enabled:
|
||||
|
||||
Each case creates its deployment, drives a real /v1/ocr call, and asserts a
|
||||
well-formed OCR document comes back. Per the e2e "skip on environment, fail on
|
||||
behavior" rule, a case skips when no proxy answers but fails (never skips) once a
|
||||
request reaches it: the proxy fetches each provider's referenced secrets, so a
|
||||
well-formed OCR document comes back. Per the e2e hard-fail contract, a case
|
||||
fails when no proxy answers and also fails once a request reaches it: the proxy
|
||||
fetches each provider's referenced secrets, so a
|
||||
missing credential surfaces as a live provider error rather than silent green.
|
||||
"""
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
"""Management suite fixtures: the client plus a logged-in dashboard page.
|
||||
|
||||
Lifecycle/skip/marker live in the parent conftest. The browser fixtures drive
|
||||
Lifecycle/liveness gate/marker live in the parent conftest. The browser fixtures drive
|
||||
the dashboard the proxy serves at /ui, so browser tests exercise exactly what an
|
||||
end user sees. playwright is an optional dependency loaded behind importorskip
|
||||
inside the fixture, so the API tests in this suite collect and run without it:
|
||||
|
|
|
|||
16
tests/e2e/mcp/conftest.py
Normal file
16
tests/e2e/mcp/conftest.py
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
"""MCP suite's `client` fixture.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness handling, and the
|
||||
`e2e`/`covers` markers live in the parent tests/e2e/conftest.py. McpClient holds
|
||||
the shared Gateway, so the `resources` fixture tears down whatever this suite
|
||||
creates (keys via the Gateway, MCP servers via the deferred cleanups).
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from mcp_client import McpClient, build_client
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def client() -> McpClient:
|
||||
return build_client()
|
||||
153
tests/e2e/mcp/mcp_client.py
Normal file
153
tests/e2e/mcp/mcp_client.py
Normal file
|
|
@ -0,0 +1,153 @@
|
|||
"""Client for the MCP e2e suite: admin server registration plus the api_key tool
|
||||
surface.
|
||||
|
||||
An admin registers an upstream MCP server through the management API
|
||||
(`/v1/mcp/server`, persisted in the DB) and grants a virtual key access to it via
|
||||
`object_permission.mcp_servers`. Keys then reach the server through the REST bridge
|
||||
the proxy exposes for api_key auth (`/mcp-rest/tools/list`, `/mcp-rest/tools/call`),
|
||||
which `user_api_key_auth` gates the same way the JSON-RPC `/mcp` surface does. The
|
||||
request/response bodies are co-located here because only this suite speaks MCP.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, RootModel
|
||||
|
||||
from e2e_gateway import Gateway, build_gateway
|
||||
from e2e_http import Headers, NoBody, Result, unwrap
|
||||
from models import KeyGenerateBody, ObjectPermission
|
||||
|
||||
|
||||
class ApiKeyHeaders(Headers):
|
||||
x_litellm_api_key: str = Field(serialization_alias="x-litellm-api-key")
|
||||
|
||||
|
||||
class McpServerNewBody(BaseModel):
|
||||
server_name: str
|
||||
alias: str
|
||||
url: str
|
||||
transport: str = "http"
|
||||
|
||||
|
||||
class McpServerNewResponse(BaseModel):
|
||||
server_id: str
|
||||
|
||||
|
||||
class McpServerRow(BaseModel):
|
||||
server_id: str
|
||||
alias: str | None = None
|
||||
url: str | None = None
|
||||
|
||||
|
||||
class McpServersListResponse(RootModel[list[McpServerRow]]):
|
||||
pass
|
||||
|
||||
|
||||
class McpToolMcpInfo(BaseModel):
|
||||
server_id: str | None = None
|
||||
alias: str | None = None
|
||||
|
||||
|
||||
class McpToolEntry(BaseModel):
|
||||
name: str
|
||||
description: str | None = None
|
||||
mcp_info: McpToolMcpInfo | None = None
|
||||
|
||||
|
||||
class McpToolsListResponse(BaseModel):
|
||||
tools: list[McpToolEntry] = []
|
||||
error: str | None = None
|
||||
message: str | None = None
|
||||
|
||||
def tool_names_for_server(self, server_id: str) -> frozenset[str]:
|
||||
return frozenset(
|
||||
tool.name
|
||||
for tool in self.tools
|
||||
if tool.mcp_info is not None and tool.mcp_info.server_id == server_id
|
||||
)
|
||||
|
||||
|
||||
class McpCallToolBody(BaseModel):
|
||||
name: str
|
||||
arguments: dict[str, int]
|
||||
server_id: str
|
||||
|
||||
|
||||
class McpCallContent(BaseModel):
|
||||
type: str | None = None
|
||||
text: str | None = None
|
||||
|
||||
|
||||
class McpCallToolResponse(BaseModel):
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
content: list[McpCallContent] = []
|
||||
is_error: bool | None = Field(default=None, alias="isError")
|
||||
|
||||
@property
|
||||
def first_text(self) -> str | None:
|
||||
return self.content[0].text if self.content else None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class McpClient:
|
||||
gateway: Gateway
|
||||
|
||||
def register_server(self, *, server_name: str, alias: str, url: str) -> str:
|
||||
return unwrap(
|
||||
self.gateway.transport.post(
|
||||
"/v1/mcp/server",
|
||||
headers=self.gateway.transport.master,
|
||||
json=McpServerNewBody(server_name=server_name, alias=alias, url=url),
|
||||
response_type=McpServerNewResponse,
|
||||
)
|
||||
).server_id
|
||||
|
||||
def delete_server(self, server_id: str) -> None:
|
||||
_ = self.gateway.transport.delete(
|
||||
f"/v1/mcp/server/{server_id}",
|
||||
headers=self.gateway.transport.master,
|
||||
json=NoBody(),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
||||
def registered_servers(self) -> list[McpServerRow]:
|
||||
return unwrap(
|
||||
self.gateway.transport.get(
|
||||
"/v1/mcp/server",
|
||||
headers=self.gateway.transport.master,
|
||||
params=NoBody(),
|
||||
response_type=McpServersListResponse,
|
||||
)
|
||||
).root
|
||||
|
||||
def generate_key(self, *, user_id: str, mcp_servers: list[str] | None) -> str:
|
||||
object_permission = (
|
||||
ObjectPermission(mcp_servers=mcp_servers) if mcp_servers is not None else None
|
||||
)
|
||||
return self.gateway.generate_key(
|
||||
KeyGenerateBody(models=[], user_id=user_id, object_permission=object_permission)
|
||||
)
|
||||
|
||||
def list_tools(self, key: str) -> Result[McpToolsListResponse]:
|
||||
return self.gateway.transport.get(
|
||||
"/mcp-rest/tools/list",
|
||||
headers=ApiKeyHeaders(x_litellm_api_key=key),
|
||||
params=NoBody(),
|
||||
response_type=McpToolsListResponse,
|
||||
)
|
||||
|
||||
def call_tool(
|
||||
self, key: str, *, server_id: str, name: str, arguments: dict[str, int]
|
||||
) -> Result[McpCallToolResponse]:
|
||||
return self.gateway.transport.post(
|
||||
"/mcp-rest/tools/call",
|
||||
headers=ApiKeyHeaders(x_litellm_api_key=key),
|
||||
json=McpCallToolBody(name=name, arguments=arguments, server_id=server_id),
|
||||
response_type=McpCallToolResponse,
|
||||
)
|
||||
|
||||
|
||||
def build_client() -> McpClient:
|
||||
return McpClient(gateway=build_gateway())
|
||||
103
tests/e2e/mcp/test_mcp_key_access_e2e.py
Normal file
103
tests/e2e/mcp/test_mcp_key_access_e2e.py
Normal file
|
|
@ -0,0 +1,103 @@
|
|||
"""Live e2e: a virtual key without MCP access is denied an MCP server's tools.
|
||||
|
||||
An admin registers an upstream MCP server through the management API (persisted in
|
||||
the DB, picked up without a restart) and queues its deletion. Two keys are created
|
||||
against that one server: one granted access through `object_permission.mcp_servers`
|
||||
and one with no MCP grant at all. The permitted key is the control that proves the
|
||||
upstream is alive and the tool is callable, so a failure on the denied key is an
|
||||
authorization denial rather than a dead server. The denied key must then see none
|
||||
of the server's tools on `tools/list` and must be refused with a 403 on
|
||||
`tools/call`.
|
||||
|
||||
Both the recorded state (the server is registered; the permitted key resolves its
|
||||
tools) and the enforced behavior (the unpermitted key sees nothing and is blocked)
|
||||
are asserted, so a regression that leaks tools to an ungranted key or drops the
|
||||
call-time permission check fails here.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import UnknownApiError, unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from mcp_client import McpClient
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
MCP_UPSTREAM_URL = os.environ.get("E2E_MCP_UPSTREAM_URL", "http://mcp-upstream:8090/mcp")
|
||||
MATH_TOOLS = frozenset({"add", "multiply"})
|
||||
|
||||
|
||||
def _register_math_server(client: McpClient, resources: ResourceManager) -> str:
|
||||
name = f"e2e_math_{unique_marker()}"
|
||||
server_id = client.register_server(server_name=name, alias=name, url=MCP_UPSTREAM_URL)
|
||||
resources.defer(lambda: client.delete_server(server_id))
|
||||
return server_id
|
||||
|
||||
|
||||
def _key(client: McpClient, resources: ResourceManager, *, mcp_servers: list[str] | None) -> str:
|
||||
label = "allowed" if mcp_servers else "denied"
|
||||
key = client.generate_key(user_id=f"e2e-mcp-{label}-{unique_marker()}", mcp_servers=mcp_servers)
|
||||
resources.defer(lambda: client.gateway.delete_key(key))
|
||||
return key
|
||||
|
||||
|
||||
def _assert_registered(client: McpClient, server_id: str) -> None:
|
||||
registered = {row.server_id for row in client.registered_servers()}
|
||||
assert server_id in registered, f"registered server {server_id} absent from /v1/mcp/server: {registered}"
|
||||
|
||||
|
||||
class TestMcpKeyWithoutAccessIsDenied:
|
||||
@pytest.mark.covers("mcp.list_tools.api_key.denied_without_permission")
|
||||
def test_list_tools_denied_without_permission(
|
||||
self, client: McpClient, resources: ResourceManager
|
||||
) -> None:
|
||||
server_id = _register_math_server(client, resources)
|
||||
_assert_registered(client, server_id)
|
||||
|
||||
permitted_key = _key(client, resources, mcp_servers=[server_id])
|
||||
denied_key = _key(client, resources, mcp_servers=None)
|
||||
|
||||
permitted_tools = unwrap(client.list_tools(permitted_key)).tool_names_for_server(server_id)
|
||||
assert MATH_TOOLS <= permitted_tools, (
|
||||
f"granted key did not see the server's tools (upstream dead or grant not applied): "
|
||||
f"{permitted_tools}"
|
||||
)
|
||||
|
||||
denied_tools = unwrap(client.list_tools(denied_key)).tool_names_for_server(server_id)
|
||||
assert denied_tools == frozenset(), (
|
||||
f"ungranted key saw the server's tools; tools/list leaked across the permission "
|
||||
f"boundary: {denied_tools}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("mcp.call_tool.api_key.denied_without_permission")
|
||||
def test_call_tool_denied_without_permission(
|
||||
self, client: McpClient, resources: ResourceManager
|
||||
) -> None:
|
||||
server_id = _register_math_server(client, resources)
|
||||
_assert_registered(client, server_id)
|
||||
|
||||
permitted_key = _key(client, resources, mcp_servers=[server_id])
|
||||
denied_key = _key(client, resources, mcp_servers=None)
|
||||
|
||||
permitted_tools = unwrap(client.list_tools(permitted_key)).tool_names_for_server(server_id)
|
||||
assert "add" in permitted_tools, (
|
||||
f"granted key did not discover the add tool (upstream dead or grant not applied): "
|
||||
f"{permitted_tools}"
|
||||
)
|
||||
|
||||
permitted_call = unwrap(
|
||||
client.call_tool(permitted_key, server_id=server_id, name="add", arguments={"a": 3, "b": 4})
|
||||
)
|
||||
assert permitted_call.is_error is not True, f"granted key's tool call errored: {permitted_call}"
|
||||
assert permitted_call.first_text == "7", (
|
||||
f"granted key's add(3, 4) did not return 7 (upstream not reachable): {permitted_call}"
|
||||
)
|
||||
|
||||
match client.call_tool(denied_key, server_id=server_id, name="add", arguments={"a": 3, "b": 4}):
|
||||
case UnknownApiError(status_code=403, body=body):
|
||||
assert "access_denied" in body, f"403 was not an MCP access denial: {body}"
|
||||
case other:
|
||||
pytest.fail(f"ungranted key's tool call was not refused with 403 access_denied: {other}")
|
||||
|
|
@ -39,6 +39,10 @@ class KeyMetadata(BaseModel):
|
|||
logging: list[KeyLoggingCallback] | None = None
|
||||
|
||||
|
||||
class ObjectPermission(BaseModel):
|
||||
mcp_servers: list[str] | None = None
|
||||
|
||||
|
||||
class KeyGenerateBody(BaseModel):
|
||||
models: list[str] = []
|
||||
duration: str | None = None
|
||||
|
|
@ -57,6 +61,7 @@ class KeyGenerateBody(BaseModel):
|
|||
rpm_limit: int | None = None
|
||||
allowed_routes: list[str] | None = None
|
||||
metadata: KeyMetadata | None = None
|
||||
object_permission: ObjectPermission | None = None
|
||||
|
||||
|
||||
class KeyGenerateResponse(BaseModel):
|
||||
|
|
|
|||
|
|
@ -33,12 +33,26 @@ _TEAM_READY_SLEEP_SECONDS = 0.4
|
|||
|
||||
class UserNewBody(BaseModel):
|
||||
max_budget: float
|
||||
budget_duration: str | None = None
|
||||
|
||||
|
||||
class UserNewResponse(BaseModel):
|
||||
user_id: str
|
||||
|
||||
|
||||
class UserInfoParams(BaseModel):
|
||||
user_id: str
|
||||
|
||||
|
||||
class UserInfoRow(BaseModel):
|
||||
spend: float | None = None
|
||||
max_budget: float | None = None
|
||||
|
||||
|
||||
class UserInfoResponse(BaseModel):
|
||||
user_info: UserInfoRow | None = None
|
||||
|
||||
|
||||
class UserDeleteBody(BaseModel):
|
||||
user_ids: list[str]
|
||||
|
||||
|
|
@ -51,6 +65,7 @@ class CustomerNewBody(BaseModel):
|
|||
class OrgNewBody(BaseModel):
|
||||
organization_alias: str
|
||||
max_budget: float
|
||||
budget_duration: str | None = None
|
||||
|
||||
|
||||
class OrgNewResponse(BaseModel):
|
||||
|
|
@ -61,6 +76,14 @@ class OrgDeleteBody(BaseModel):
|
|||
organization_ids: list[str]
|
||||
|
||||
|
||||
class OrgInfoParams(BaseModel):
|
||||
organization_id: str
|
||||
|
||||
|
||||
class OrgInfoResponse(BaseModel):
|
||||
budget_id: str | None = None
|
||||
|
||||
|
||||
class TeamMember(BaseModel):
|
||||
role: str
|
||||
user_id: str
|
||||
|
|
@ -69,6 +92,7 @@ class TeamMember(BaseModel):
|
|||
class TeamNewBody(BaseModel):
|
||||
team_alias: str
|
||||
max_budget: float | None = None
|
||||
budget_duration: str | None = None
|
||||
organization_id: str | None = None
|
||||
budget_limits: list[BudgetWindow] | None = None
|
||||
|
||||
|
|
@ -244,12 +268,12 @@ class BudgetClient:
|
|||
|
||||
# ---- internal user --------------------------------------------------
|
||||
|
||||
def create_user(self, *, max_budget: float) -> str:
|
||||
def create_user(self, *, max_budget: float, budget_duration: str | None = None) -> str:
|
||||
return unwrap(
|
||||
self.gateway.transport.post(
|
||||
"/user/new",
|
||||
headers=self.gateway.transport.master,
|
||||
json=UserNewBody(max_budget=max_budget),
|
||||
json=UserNewBody(max_budget=max_budget, budget_duration=budget_duration),
|
||||
response_type=UserNewResponse,
|
||||
)
|
||||
).user_id
|
||||
|
|
@ -262,6 +286,19 @@ class BudgetClient:
|
|||
response_type=NoBody,
|
||||
)
|
||||
|
||||
def user_info(self, user_id: str) -> UserInfoRow | None:
|
||||
result = self.gateway.transport.get(
|
||||
"/user/info",
|
||||
headers=self.gateway.transport.master,
|
||||
params=UserInfoParams(user_id=user_id),
|
||||
response_type=UserInfoResponse,
|
||||
)
|
||||
match result:
|
||||
case Success(data=data):
|
||||
return data.user_info
|
||||
case _:
|
||||
return None
|
||||
|
||||
# ---- customer / end-user -------------------------------------------
|
||||
|
||||
def create_customer(self, customer_id: str, *, max_budget: float) -> str:
|
||||
|
|
@ -275,16 +312,36 @@ class BudgetClient:
|
|||
|
||||
# ---- organization ---------------------------------------------------
|
||||
|
||||
def create_org(self, *, max_budget: float, alias: str) -> str:
|
||||
def create_org(self, *, max_budget: float, alias: str, budget_duration: str | None = None) -> str:
|
||||
return unwrap(
|
||||
self.gateway.transport.post(
|
||||
"/organization/new",
|
||||
headers=self.gateway.transport.master,
|
||||
json=OrgNewBody(organization_alias=alias, max_budget=max_budget),
|
||||
json=OrgNewBody(
|
||||
organization_alias=alias,
|
||||
max_budget=max_budget,
|
||||
budget_duration=budget_duration,
|
||||
),
|
||||
response_type=OrgNewResponse,
|
||||
)
|
||||
).organization_id
|
||||
|
||||
def org_budget_id(self, org_id: str) -> str | None:
|
||||
"""The id of the budget row backing an org; its budget_reset_at is read via
|
||||
budget_info (LIT-4570: /organization/new stores budget_duration without
|
||||
scheduling budget_reset_at, so the reset job's first tick schedules it)."""
|
||||
result = self.gateway.transport.get(
|
||||
"/organization/info",
|
||||
headers=self.gateway.transport.master,
|
||||
params=OrgInfoParams(organization_id=org_id),
|
||||
response_type=OrgInfoResponse,
|
||||
)
|
||||
match result:
|
||||
case Success(data=data):
|
||||
return data.budget_id
|
||||
case _:
|
||||
return None
|
||||
|
||||
def delete_org(self, org_id: str) -> None:
|
||||
_ = self.gateway.transport.delete(
|
||||
"/organization/delete",
|
||||
|
|
@ -300,6 +357,7 @@ class BudgetClient:
|
|||
*,
|
||||
alias: str,
|
||||
max_budget: float | None = None,
|
||||
budget_duration: str | None = None,
|
||||
organization_id: str | None = None,
|
||||
budget_limits: list[BudgetWindow] | None = None,
|
||||
) -> str:
|
||||
|
|
@ -310,6 +368,7 @@ class BudgetClient:
|
|||
json=TeamNewBody(
|
||||
team_alias=alias,
|
||||
max_budget=max_budget,
|
||||
budget_duration=budget_duration,
|
||||
organization_id=organization_id,
|
||||
budget_limits=budget_limits,
|
||||
),
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
"""Budgets suite's `client` fixture.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. BudgetClient holds the shared Gateway,
|
||||
so the `resources` fixture cleans up keys through it; tests register entity deletes
|
||||
via `resources.defer(...)`.
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Each entity is an E2ECase (lifecycle.E2ECase) driven by run_case: init() creates
|
|||
the budgeted entity + a key, run() drives spend until a `budget_exceeded` block,
|
||||
teardown() deletes everything init() created (always runs, even on failure/skip).
|
||||
Covers the entities with no prior live coverage - internal user, end-user,
|
||||
organization, team member. See BUDGET_TEST_COVERAGE_MATRIX.md.
|
||||
organization, team member - plus key and team. See BUDGET_TEST_COVERAGE_MATRIX.md.
|
||||
|
||||
A non-budget error fails hard (never a skip); if calls never get blocked, budget
|
||||
enforcement is broken -> fail.
|
||||
|
|
@ -18,16 +18,17 @@ import pytest
|
|||
|
||||
from budget_client import BudgetClient, is_budget_block
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from e2e_http import StreamingResponse, require_successful_call
|
||||
from lifecycle import run_case
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> None:
|
||||
"""Send paid calls until the entity's budget blocks one. Key/user/org/member
|
||||
block within a couple calls off real-time reservation counters; the end-user
|
||||
budget enforces off table spend that lands on the batch write, so it takes a
|
||||
few more. A non-budget error fails hard (never a skip)."""
|
||||
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> StreamingResponse:
|
||||
"""Send paid calls until the entity's budget blocks one; return the blocked
|
||||
response so callers can assert on its shape. Key/user/org/member block within
|
||||
a couple calls off real-time reservation counters; the end-user budget
|
||||
enforces off table spend that lands on the batch write, so it takes a few
|
||||
more. A non-budget error fails hard (never a skip)."""
|
||||
for _ in range(40):
|
||||
result = client.chat(
|
||||
key,
|
||||
|
|
@ -37,7 +38,7 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") ->
|
|||
user=user or None,
|
||||
)
|
||||
if is_budget_block(result):
|
||||
return
|
||||
return result
|
||||
require_successful_call(result)
|
||||
time.sleep(2)
|
||||
pytest.fail("budget never enforced within the call budget")
|
||||
|
|
@ -69,18 +70,85 @@ class _BudgetCase:
|
|||
|
||||
|
||||
class KeyBudgetCase(_BudgetCase):
|
||||
"""A bare key (no team_id / user_id) carrying its own max_budget, so only the
|
||||
key-level budget can be the thing that blocks. The refusal must be a 429
|
||||
budget_exceeded; any other error already fails via _assert_budget_blocks."""
|
||||
|
||||
def init(self) -> None:
|
||||
self.key = self.client.generate_key(max_budget=3e-6)
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
|
||||
def run(self) -> None:
|
||||
blocked = _assert_budget_blocks(self.client, self.key)
|
||||
assert blocked.status_code == 429, (
|
||||
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
|
||||
)
|
||||
|
||||
|
||||
class TeamBudgetCase(_BudgetCase):
|
||||
"""An admin caps a whole team: two keys under a tiny-budget team, neither with
|
||||
a key-level budget. Key A is driven until the team cap blocks it; key B's very
|
||||
first call must then be refused too, proving the cap sits on the team, not the
|
||||
key that spent. Both refusals must be 429 budget_exceeded."""
|
||||
|
||||
def init(self) -> None:
|
||||
team_id = self.client.create_team(
|
||||
alias=f"e2e-budget-team-{unique_marker()}", max_budget=3e-6
|
||||
)
|
||||
self._undo.append(lambda: self.client.delete_team(team_id))
|
||||
self.key = self.client.generate_key(team_id=team_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
self._sibling_key = self.client.generate_key(team_id=team_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self._sibling_key))
|
||||
|
||||
def run(self) -> None:
|
||||
blocked = _assert_budget_blocks(self.client, self.key)
|
||||
assert blocked.status_code == 429, (
|
||||
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
|
||||
)
|
||||
sibling = self.client.chat(
|
||||
self._sibling_key,
|
||||
"claude-haiku-4-5",
|
||||
f"spend {unique_marker()}",
|
||||
max_tokens=16,
|
||||
)
|
||||
assert is_budget_block(sibling) and sibling.status_code == 429, (
|
||||
f"a sibling key on the capped team must get the same 429 budget_exceeded, "
|
||||
f"got {sibling.status_code}: {sibling.body[:200]}"
|
||||
)
|
||||
|
||||
|
||||
class InternalUserBudgetCase(_BudgetCase):
|
||||
"""A user's max_budget follows the person, not the key. The capped user holds
|
||||
two personal keys (no team, no key budgets) plus a team-member key on an
|
||||
uncapped team; once the first personal key is refused, the other two must be
|
||||
refused as well - a second key is not a fresh allowance, and since #32005 the
|
||||
user budget draws down team keys too. All refusals must be 429 budget_exceeded."""
|
||||
|
||||
def init(self) -> None:
|
||||
user_id = self.client.create_user(max_budget=3e-6)
|
||||
self._undo.append(lambda: self.client.delete_user(user_id))
|
||||
# personal key (no team) -> the user budget governs
|
||||
self.key = self.client.generate_key(user_id=user_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
self._second_key = self.client.generate_key(user_id=user_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self._second_key))
|
||||
team_id = self.client.create_team(alias=f"e2e-budget-team-{unique_marker()}")
|
||||
self._undo.append(lambda: self.client.delete_team(team_id))
|
||||
self.client.add_team_member(team_id, user_id)
|
||||
self._team_key = self.client.generate_key(team_id=team_id, user_id=user_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self._team_key))
|
||||
|
||||
def run(self) -> None:
|
||||
blocked = _assert_budget_blocks(self.client, self.key)
|
||||
assert blocked.status_code == 429, (
|
||||
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
|
||||
)
|
||||
for label, key in (("second personal key", self._second_key), ("team-member key", self._team_key)):
|
||||
result = self.client.chat(key, "claude-haiku-4-5", f"spend {unique_marker()}", max_tokens=16)
|
||||
assert is_budget_block(result) and result.status_code == 429, (
|
||||
f"the {label} of a user over budget must get the same 429 budget_exceeded, "
|
||||
f"got {result.status_code}: {result.body[:200]}"
|
||||
)
|
||||
|
||||
|
||||
class EndUserBudgetCase(_BudgetCase):
|
||||
|
|
@ -97,34 +165,67 @@ class EndUserBudgetCase(_BudgetCase):
|
|||
|
||||
|
||||
class OrganizationBudgetCase(_BudgetCase):
|
||||
"""Org carries the tiny budget; the team under it and the key carry none, so
|
||||
the org is the only entity that can block (the historically weak link). The
|
||||
refusal must be a 429 budget_exceeded that names the org as the blocker."""
|
||||
|
||||
def init(self) -> None:
|
||||
# Org carries the tiny budget; the team under it has none, so a block here
|
||||
# is org-level enforcement (the historically weak link).
|
||||
org_id = self.client.create_org(
|
||||
self._org_id = self.client.create_org(
|
||||
max_budget=3e-6, alias=f"e2e-budget-org-{unique_marker()}"
|
||||
)
|
||||
self._undo.append(lambda: self.client.delete_org(org_id))
|
||||
self._undo.append(lambda: self.client.delete_org(self._org_id))
|
||||
team_id = self.client.create_team(
|
||||
alias=f"e2e-budget-team-{unique_marker()}", organization_id=org_id
|
||||
alias=f"e2e-budget-team-{unique_marker()}", organization_id=self._org_id
|
||||
)
|
||||
self._undo.append(lambda: self.client.delete_team(team_id))
|
||||
self.key = self.client.generate_key(team_id=team_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
|
||||
def run(self) -> None:
|
||||
blocked = _assert_budget_blocks(self.client, self.key)
|
||||
assert blocked.status_code == 429, (
|
||||
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
|
||||
)
|
||||
assert f"Organization={self._org_id}" in blocked.body, (
|
||||
f"refusal must name the org as the blocker, got: {blocked.body[:200]}"
|
||||
)
|
||||
|
||||
|
||||
class TeamMemberBudgetCase(_BudgetCase):
|
||||
"""Member A's per-team budget is tiny while the team and both members' user
|
||||
budgets are roomy (100.0), so the only cap that can trip is A's: a block
|
||||
proves member-level enforcement and must be a 429 budget_exceeded. Teammate
|
||||
B, uncapped on the same team, must keep serving after A is cut off, proving
|
||||
the member cap does not leak onto the team or its members."""
|
||||
|
||||
def init(self) -> None:
|
||||
# Member's per-team budget is tiny while the team has a large budget, so a
|
||||
# block proves member-level (not team-level) enforcement.
|
||||
team_id = self.client.create_team(
|
||||
self._team_id = self.client.create_team(
|
||||
alias=f"e2e-budget-team-{unique_marker()}", max_budget=100.0
|
||||
)
|
||||
self._undo.append(lambda: self.client.delete_team(team_id))
|
||||
user_id = self.client.create_user(max_budget=100.0)
|
||||
self._undo.append(lambda: self.client.delete_user(user_id))
|
||||
self.client.add_team_member(team_id, user_id, max_budget_in_team=3e-6)
|
||||
self.key = self.client.generate_key(team_id=team_id, user_id=user_id)
|
||||
self._undo.append(lambda: self.client.delete_team(self._team_id))
|
||||
self._member_id = self.client.create_user(max_budget=100.0)
|
||||
self._undo.append(lambda: self.client.delete_user(self._member_id))
|
||||
self.client.add_team_member(self._team_id, self._member_id, max_budget_in_team=3e-6)
|
||||
self.key = self.client.generate_key(team_id=self._team_id, user_id=self._member_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
teammate_id = self.client.create_user(max_budget=100.0)
|
||||
self._undo.append(lambda: self.client.delete_user(teammate_id))
|
||||
self.client.add_team_member(self._team_id, teammate_id)
|
||||
self._teammate_key = self.client.generate_key(team_id=self._team_id, user_id=teammate_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self._teammate_key))
|
||||
|
||||
def run(self) -> None:
|
||||
blocked = _assert_budget_blocks(self.client, self.key)
|
||||
assert blocked.status_code == 429, (
|
||||
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
|
||||
)
|
||||
teammate = self.client.chat(
|
||||
self._teammate_key,
|
||||
"claude-haiku-4-5",
|
||||
f"spend {unique_marker()}",
|
||||
max_tokens=16,
|
||||
)
|
||||
require_successful_call(teammate)
|
||||
|
||||
|
||||
def _case_id(case_cls: Type[_BudgetCase]) -> str:
|
||||
|
|
@ -138,6 +239,10 @@ def _case_id(case_cls: Type[_BudgetCase]) -> str:
|
|||
KeyBudgetCase,
|
||||
marks=pytest.mark.covers("quota_management.budget.key.blocks_over_limit"),
|
||||
),
|
||||
pytest.param(
|
||||
TeamBudgetCase,
|
||||
marks=pytest.mark.covers("quota_management.budget.team.blocks_over_limit"),
|
||||
),
|
||||
pytest.param(
|
||||
InternalUserBudgetCase,
|
||||
marks=pytest.mark.covers("quota_management.budget.internal_user.blocks_over_limit"),
|
||||
|
|
|
|||
|
|
@ -1,12 +1,4 @@
|
|||
"""Live e2e: a key budget resets (zeroes spend) after its budget_duration.
|
||||
|
||||
Short budget_duration (30s) + the fast-rescheduled reset job: a key blocked for
|
||||
exceeding its max_budget starts succeeding again once the duration elapses and the
|
||||
reset job zeroes key.spend. Closes the reset-zeroing gap in
|
||||
BUDGET_TEST_COVERAGE_MATRIX.md (reset_budget_for_litellm_keys), which the unit
|
||||
suite covers but no live test did - distinct from the per-window reset in
|
||||
test_multi_window_budget_e2e.py.
|
||||
"""
|
||||
"""Live e2e: an entity blocked over its max_budget serves again after its budget_duration window."""
|
||||
|
||||
import time
|
||||
|
||||
|
|
@ -19,42 +11,107 @@ from lifecycle import ResourceManager
|
|||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
TINY_CAP = 3e-6
|
||||
WINDOW = "30s"
|
||||
RESET_DEADLINE_SECONDS = 150
|
||||
|
||||
|
||||
def _call(client: BudgetClient, key: str):
|
||||
return client.chat(
|
||||
key, "claude-haiku-4-5", f"reset {unique_marker()}", max_tokens=16
|
||||
)
|
||||
return client.chat(key, "claude-haiku-4-5", f"reset {unique_marker()}", max_tokens=16)
|
||||
|
||||
|
||||
@pytest.mark.covers("quota_management.budget.key.resets_after_window")
|
||||
def test_key_budget_resets_after_duration(
|
||||
client: BudgetClient, resources: ResourceManager
|
||||
) -> None:
|
||||
key = client.generate_key(max_budget=3e-6, budget_duration="30s")
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
# 1. exceed the budget -> litellm returns budget_exceeded
|
||||
blocked = False
|
||||
for _ in range(20):
|
||||
def _drive_to_block(client: BudgetClient, key: str) -> None:
|
||||
"""Spend until the cap blocks a call, staying under one window so the block
|
||||
is observed before the reset job can fire; fail hard if enforcement never trips."""
|
||||
for _ in range(12):
|
||||
result = _call(client, key)
|
||||
if is_budget_block(result):
|
||||
blocked = True
|
||||
break
|
||||
return
|
||||
require_successful_call(result)
|
||||
time.sleep(2)
|
||||
assert blocked, "key budget never enforced"
|
||||
pytest.fail("budget never enforced before the window could reset")
|
||||
|
||||
# 2. once the 30s duration elapses + the reset job runs, key.spend zeroes and
|
||||
# calls flow again. The window is wall-clock-aligned, so the reset lands up to
|
||||
# a window later, then the rescheduler (~15-20s) zeroes the spend; allow
|
||||
# generous headroom over that. A stuck rescheduler is caught by the wait-loop
|
||||
# timeout, not this elapsed bound.
|
||||
start = time.monotonic()
|
||||
while time.monotonic() < start + 150:
|
||||
|
||||
def _poll_until_serves_again(client: BudgetClient, key: str) -> None:
|
||||
"""Poll past the window until the blocked key serves again; every refusal must
|
||||
stay a budget block, so a crashed reset path or provider error fails loudly."""
|
||||
deadline = time.monotonic() + RESET_DEADLINE_SECONDS
|
||||
while time.monotonic() < deadline:
|
||||
time.sleep(5)
|
||||
result = _call(client, key)
|
||||
if result.ok:
|
||||
assert time.monotonic() - start < 120, "reset too slow for a 30s budget"
|
||||
return
|
||||
assert is_budget_block(result), f"non-budget error: {result.body[:200]}"
|
||||
pytest.fail("key budget never reset within 150s")
|
||||
if not is_budget_block(result):
|
||||
pytest.fail(f"non-budget error during reset wait: HTTP {result.status_code}: {result.body[:200]}")
|
||||
pytest.fail(f"budget never reset within {RESET_DEADLINE_SECONDS}s")
|
||||
|
||||
|
||||
class TestBudgetResetDiagonal:
|
||||
@pytest.mark.covers("quota_management.budget.key.resets_after_window")
|
||||
def test_bare_key_budget_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
|
||||
key = client.generate_key(max_budget=TINY_CAP, budget_duration=WINDOW)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
_drive_to_block(client, key)
|
||||
_poll_until_serves_again(client, key)
|
||||
|
||||
@pytest.mark.covers("quota_management.budget.team.resets_after_window")
|
||||
def test_team_budget_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
|
||||
team_id = client.create_team(
|
||||
alias=f"e2e-team-reset-{unique_marker()}", max_budget=TINY_CAP, budget_duration=WINDOW
|
||||
)
|
||||
resources.defer(lambda: client.delete_team(team_id))
|
||||
key = client.generate_key(team_id=team_id)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
_drive_to_block(client, key)
|
||||
_poll_until_serves_again(client, key)
|
||||
|
||||
@pytest.mark.covers("quota_management.budget.organization.resets_after_window")
|
||||
def test_org_budget_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
|
||||
org_id = client.create_org(
|
||||
max_budget=TINY_CAP, alias=f"e2e-org-reset-{unique_marker()}", budget_duration=WINDOW
|
||||
)
|
||||
resources.defer(lambda: client.delete_org(org_id))
|
||||
team_id = client.create_team(alias=f"e2e-org-team-{unique_marker()}", organization_id=org_id)
|
||||
resources.defer(lambda: client.delete_team(team_id))
|
||||
key = client.generate_key(team_id=team_id)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
budget_id = client.org_budget_id(org_id)
|
||||
assert budget_id, "org created without a budget row"
|
||||
deadline = time.monotonic() + 30
|
||||
while not any(row.budget_reset_at for row in client.budget_info(budget_id)):
|
||||
if time.monotonic() > deadline:
|
||||
pytest.fail("org budget window never scheduled by the reset job")
|
||||
time.sleep(2)
|
||||
|
||||
_drive_to_block(client, key)
|
||||
_poll_until_serves_again(client, key)
|
||||
|
||||
@pytest.mark.covers("quota_management.budget.internal_user.resets_after_window")
|
||||
def test_personal_key_user_budget_resets_after_window(
|
||||
self, client: BudgetClient, resources: ResourceManager
|
||||
) -> None:
|
||||
user_id = client.create_user(max_budget=TINY_CAP, budget_duration=WINDOW)
|
||||
resources.defer(lambda: client.delete_user(user_id))
|
||||
key = client.generate_key(user_id=user_id)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
_drive_to_block(client, key)
|
||||
_poll_until_serves_again(client, key)
|
||||
|
||||
@pytest.mark.covers("quota_management.budget.internal_user.resets_after_window")
|
||||
def test_team_member_key_user_budget_resets_after_window(
|
||||
self, client: BudgetClient, resources: ResourceManager
|
||||
) -> None:
|
||||
user_id = client.create_user(max_budget=TINY_CAP, budget_duration=WINDOW)
|
||||
resources.defer(lambda: client.delete_user(user_id))
|
||||
team_id = client.create_team(alias=f"e2e-user-team-reset-{unique_marker()}")
|
||||
resources.defer(lambda: client.delete_team(team_id))
|
||||
client.add_team_member(team_id, user_id, max_budget_in_team=100.0)
|
||||
key = client.generate_key(team_id=team_id, user_id=user_id)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
_drive_to_block(client, key)
|
||||
_poll_until_serves_again(client, key)
|
||||
|
|
|
|||
|
|
@ -23,14 +23,16 @@ pytestmark = pytest.mark.e2e
|
|||
WINDOW_SECONDS = 30 # the tight window; calls succeed again only after it elapses
|
||||
# Prefer the OpenAI cheap model for this polling test: under the full stage suite
|
||||
# Claude chat latency + ALB target idle timeout (~60s) can surface as awselb 502
|
||||
# HTML mid-wait, which is not a budget signal. gpt-5.5 + 1 token stays well under
|
||||
# that ceiling so the wait loop measures window reset, not provider/ALB timeout.
|
||||
# HTML mid-wait, which is not a budget signal. gpt-5.5 stays well under that
|
||||
# ceiling so the wait loop measures window reset, not provider/ALB timeout.
|
||||
# max_tokens must be >1: gpt-5.5 refuses completions that hit the output limit
|
||||
# mid-message when capped at 1 token.
|
||||
MODEL = CHEAP_OPENAI_MODEL
|
||||
|
||||
|
||||
def _call(client: BudgetClient, key: str):
|
||||
return client.chat(
|
||||
key, MODEL, f"window {unique_marker()}", max_tokens=1
|
||||
key, MODEL, f"window {unique_marker()}", max_tokens=16
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,119 @@
|
|||
"""Live e2e: per-team-member budgets are enforced independently between members.
|
||||
|
||||
Two members share one team that has a large team budget. The tight member is capped
|
||||
at a tiny per-team budget and spends past it; the roomy member has plenty of room.
|
||||
Once the tight member is blocked with budget_exceeded, the roomy member still serves
|
||||
on the same team, its calls land in the spend logs under its own user id, and the
|
||||
tight member stays blocked. A shared or leaky member counter would either block the
|
||||
roomy member too or let the tight member back through once its peer spent.
|
||||
"""
|
||||
|
||||
import time
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient, is_budget_block
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import Success, require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatBody, ChatMessage
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
MODEL = "gpt-5.5"
|
||||
TEAM_BUDGET = 100.0
|
||||
TIGHT_MEMBER_BUDGET = 3e-6
|
||||
ROOMY_MEMBER_BUDGET = 100.0
|
||||
ROOMY_BURST = 3
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Pair:
|
||||
team_id: str
|
||||
tight_user_id: str
|
||||
roomy_user_id: str
|
||||
tight_key: str
|
||||
roomy_key: str
|
||||
|
||||
|
||||
@pytest.fixture(scope="class")
|
||||
def pair(client: BudgetClient) -> Iterator[_Pair]:
|
||||
"""One team with a large budget and two members on it: a tight member capped at
|
||||
a tiny per-team budget and a roomy member with headroom, each with their own key.
|
||||
Shared across the class and torn down LIFO best-effort when it finishes."""
|
||||
resources = ResourceManager(client=client.gateway)
|
||||
try:
|
||||
marker = unique_marker()
|
||||
team_id = client.create_team(alias=f"e2e-member-iso-{marker}", max_budget=TEAM_BUDGET)
|
||||
resources.defer(lambda: client.delete_team(team_id))
|
||||
tight_user = client.create_user(max_budget=TEAM_BUDGET)
|
||||
resources.defer(lambda: client.delete_user(tight_user))
|
||||
roomy_user = client.create_user(max_budget=TEAM_BUDGET)
|
||||
resources.defer(lambda: client.delete_user(roomy_user))
|
||||
client.add_team_member(team_id, tight_user, max_budget_in_team=TIGHT_MEMBER_BUDGET)
|
||||
client.add_team_member(team_id, roomy_user, max_budget_in_team=ROOMY_MEMBER_BUDGET)
|
||||
tight_key = client.generate_key(team_id=team_id, user_id=tight_user)
|
||||
resources.defer(lambda: client.delete_key(tight_key))
|
||||
roomy_key = client.generate_key(team_id=team_id, user_id=roomy_user)
|
||||
resources.defer(lambda: client.delete_key(roomy_key))
|
||||
yield _Pair(
|
||||
team_id=team_id,
|
||||
tight_user_id=tight_user,
|
||||
roomy_user_id=roomy_user,
|
||||
tight_key=tight_key,
|
||||
roomy_key=roomy_key,
|
||||
)
|
||||
finally:
|
||||
resources.teardown()
|
||||
|
||||
|
||||
def _roomy_send(client: BudgetClient, key: str) -> str:
|
||||
"""One roomy-member call that must go through; returns its request id."""
|
||||
match client.gateway.chat(
|
||||
key,
|
||||
ChatBody(
|
||||
model=MODEL,
|
||||
messages=[ChatMessage(role="user", content=f"roomy {unique_marker()}")],
|
||||
max_tokens=16,
|
||||
),
|
||||
):
|
||||
case Success(data=response):
|
||||
assert response.id is not None, "roomy member call returned no id"
|
||||
return response.id
|
||||
case other:
|
||||
pytest.fail(f"roomy member call failed while a peer was over budget: {other}")
|
||||
|
||||
|
||||
class TestTeamMemberBudgetIsolation:
|
||||
@pytest.mark.covers("quota_management.budget.team_member.isolates_per_member")
|
||||
def test_blocked_member_does_not_block_peer(self, client: BudgetClient, pair: _Pair) -> None:
|
||||
blocked = False
|
||||
for _ in range(40):
|
||||
result = client.chat(pair.tight_key, MODEL, f"tight {unique_marker()}", max_tokens=16)
|
||||
if is_budget_block(result):
|
||||
blocked = True
|
||||
break
|
||||
require_successful_call(result)
|
||||
time.sleep(2)
|
||||
assert blocked, "tight member's per-team budget never enforced"
|
||||
|
||||
sent = frozenset(_roomy_send(client, pair.roomy_key) for _ in range(ROOMY_BURST))
|
||||
|
||||
assert is_budget_block(
|
||||
client.chat(pair.tight_key, MODEL, f"tight {unique_marker()}", max_tokens=16)
|
||||
), "tight member stopped being blocked once the peer spent"
|
||||
|
||||
rows = client.gateway.poll_logs_for_key(
|
||||
pair.roomy_key, predicate=lambda rs: bool(sent & {r.request_id for r in rs})
|
||||
)
|
||||
logged = [row for row in rows if row.request_id in sent]
|
||||
assert logged, "none of the roomy member's calls reached the spend logs"
|
||||
for row in logged:
|
||||
assert row.user == pair.roomy_user_id, (
|
||||
f"roomy call {row.request_id} logged under user {row.user}, not {pair.roomy_user_id}"
|
||||
)
|
||||
assert row.team_id == pair.team_id, (
|
||||
f"roomy call {row.request_id} logged under team {row.team_id}, not {pair.team_id}"
|
||||
)
|
||||
|
|
@ -0,0 +1,79 @@
|
|||
"""Live e2e: a per-user max_budget is enforced across ALL of that user's keys.
|
||||
|
||||
An internal user's budget governs every personal key it owns, not only the one
|
||||
that happened to spend it down. One user with a tiny max_budget owns two keys:
|
||||
driving the first key to a budget_exceeded block then makes a fresh, untouched
|
||||
second key of the same user (which carries no budget of its own, so nothing but the
|
||||
shared user budget can block it) reject the same way, and the user's recorded spend
|
||||
has crossed the cap. A key-scoped-only budget would leave the second key serving.
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient, is_budget_block
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import StreamingResponse, require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
MODEL = "gpt-5.5"
|
||||
TINY_CAP = 3e-6
|
||||
RECORDED_SPEND_DEADLINE_SECONDS = 90
|
||||
SECOND_KEY_BLOCK_ATTEMPTS = 6
|
||||
|
||||
|
||||
def _call(client: BudgetClient, key: str) -> StreamingResponse:
|
||||
return client.chat(key, MODEL, f"across {unique_marker()}", max_tokens=16)
|
||||
|
||||
|
||||
def _drive_to_block(client: BudgetClient, key: str, subject: str) -> None:
|
||||
for _ in range(40):
|
||||
result = _call(client, key)
|
||||
if is_budget_block(result):
|
||||
return
|
||||
require_successful_call(result)
|
||||
time.sleep(2)
|
||||
pytest.fail(f"user budget never enforced on {subject} within the call budget")
|
||||
|
||||
|
||||
def _expect_prompt_block(client: BudgetClient, key: str, subject: str) -> None:
|
||||
"""The shared user budget is already exhausted before this key makes a single
|
||||
call, so a key with no budget of its own must be rejected promptly. The small
|
||||
bounded retry only absorbs spend-propagation lag between the two keys; it is far
|
||||
below the spend a key-scoped budget would need to accumulate to block itself, so
|
||||
a block here can only come from the shared user budget."""
|
||||
for _ in range(SECOND_KEY_BLOCK_ATTEMPTS):
|
||||
result = _call(client, key)
|
||||
if is_budget_block(result):
|
||||
return
|
||||
require_successful_call(result)
|
||||
time.sleep(2)
|
||||
pytest.fail(
|
||||
f"{subject} was not blocked by the shared user budget within {SECOND_KEY_BLOCK_ATTEMPTS} calls"
|
||||
)
|
||||
|
||||
|
||||
class TestUserBudgetAcrossKeys:
|
||||
@pytest.mark.covers("quota_management.budget.internal_user.enforced_across_keys")
|
||||
def test_user_budget_blocks_a_second_key(self, client: BudgetClient, resources: ResourceManager) -> None:
|
||||
user_id = client.create_user(max_budget=TINY_CAP)
|
||||
resources.defer(lambda: client.delete_user(user_id))
|
||||
|
||||
first_key = client.generate_key(user_id=user_id)
|
||||
resources.defer(lambda: client.delete_key(first_key))
|
||||
second_key = client.generate_key(user_id=user_id)
|
||||
resources.defer(lambda: client.delete_key(second_key))
|
||||
|
||||
_drive_to_block(client, first_key, "the first key")
|
||||
_expect_prompt_block(client, second_key, "the second key")
|
||||
|
||||
deadline = time.monotonic() + RECORDED_SPEND_DEADLINE_SECONDS
|
||||
while time.monotonic() < deadline:
|
||||
info = client.user_info(user_id)
|
||||
if info is not None and (info.spend or 0.0) >= TINY_CAP:
|
||||
return
|
||||
time.sleep(5)
|
||||
pytest.fail(f"user spend never reached the {TINY_CAP} cap in the recorded state")
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
"""Quota-management suite's `client` fixture.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. QuotaClient holds the shared Gateway,
|
||||
so the `resources` fixture cleans up keys through it.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -80,5 +80,5 @@ proxy + SpendLogs rows. Status: `covered` / `partial` / `gap`.
|
|||
`proxy_batch_write_at` (~60s) means rows land late; every read polls to a deadline.
|
||||
Fresh scoped key per test (isolation, xdist-safe, cleaned up). Assert invariants
|
||||
(`spend > 0`, `total == prompt + completion`, aggregate == sum), not literal
|
||||
$/token values, so pricing drift is not a failure. Skip on environment (no proxy /
|
||||
no provider key), fail on behavior (a real 2xx call with a wrong/missing row).
|
||||
$/token values, so pricing drift is not a failure. Hard-fail when no proxy
|
||||
answers, fail on behavior (a real 2xx call with a wrong/missing row).
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
"""Spend-tracking suite's `client` fixture and driver-model registration.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. SpendClient exposes the shared Gateway
|
||||
(GatewayProvider), so the `resources` fixture cleans up keys and customers this
|
||||
suite creates.
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
"""Router suite's `client` fixture.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. ComplexityRouterClient holds the shared
|
||||
Gateway, so the `resources` fixture cleans up keys this suite creates.
|
||||
|
||||
|
|
|
|||
|
|
@ -172,7 +172,7 @@ def anthropic_messages():
|
|||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Here is the full text of a complex legal agreement" * 400,
|
||||
"text": "Here is the full text of a complex legal agreement" * 500,
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
|
|
|
|||
40
tests/mcp_tests/mcp_e2e_upstream_server.py
Normal file
40
tests/mcp_tests/mcp_e2e_upstream_server.py
Normal file
|
|
@ -0,0 +1,40 @@
|
|||
"""Deterministic upstream MCP server for the mcp e2e suite.
|
||||
|
||||
A tiny FastMCP server exposing `add` and `multiply` over streamable-http so the
|
||||
suite has a self-hosted, offline upstream to register and exercise. DNS-rebinding
|
||||
protection is turned off because the litellm container reaches this over the
|
||||
compose network by service name (`mcp-upstream:8090`), not localhost, and the
|
||||
stack is an isolated throwaway. Bind host/port come from MCP_HOST/MCP_PORT.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
|
||||
mcp: FastMCP = FastMCP(
|
||||
"e2e-math",
|
||||
host=os.getenv("MCP_HOST", "0.0.0.0"),
|
||||
port=int(os.getenv("MCP_PORT", "8090")),
|
||||
transport_security=TransportSecuritySettings(enable_dns_rebinding_protection=False),
|
||||
)
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
def add(a: int, b: int) -> int:
|
||||
"""Add two integers"""
|
||||
return a + b
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
def multiply(a: int, b: int) -> int:
|
||||
"""Multiply two integers"""
|
||||
return a * b
|
||||
|
||||
|
||||
def main() -> None:
|
||||
mcp.run(transport="streamable-http")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -935,8 +935,8 @@ async def test_get_tools_from_mcp_servers():
|
|||
mcp_auth_header=mock_auth_header,
|
||||
mcp_servers=["server1"],
|
||||
)
|
||||
assert len(result) == 1, "Should only return tools from server1"
|
||||
assert result[0].name == "tool1", "Should return tool from server1"
|
||||
assert len(result.tools) == 1, "Should only return tools from server1"
|
||||
assert result.tools[0].name == "tool1", "Should return tool from server1"
|
||||
|
||||
# Test Case 2: Without specific MCP servers
|
||||
# Create a different mock manager for the second test case
|
||||
|
|
@ -978,9 +978,9 @@ async def test_get_tools_from_mcp_servers():
|
|||
mcp_auth_header=mock_auth_header,
|
||||
mcp_servers=None,
|
||||
)
|
||||
assert len(result) == 2, "Should return tools from all servers"
|
||||
assert len(result.tools) == 2, "Should return tools from all servers"
|
||||
assert (
|
||||
result[0].name == "tool1" and result[1].name == "tool2"
|
||||
result.tools[0].name == "tool1" and result.tools[1].name == "tool2"
|
||||
), "Should return tools from all servers"
|
||||
|
||||
#
|
||||
|
|
@ -1015,8 +1015,8 @@ async def test_get_tools_from_mcp_servers():
|
|||
mcp_auth_header=mock_auth_header,
|
||||
mcp_servers=["group-a"],
|
||||
)
|
||||
assert len(result) == 1, "Should only return tools from server3"
|
||||
assert result[0].name == "tool1", "Should return tool from server1"
|
||||
assert len(result.tools) == 1, "Should only return tools from server3"
|
||||
assert result.tools[0].name == "tool1", "Should return tool from server1"
|
||||
|
||||
except AssertionError as e:
|
||||
pytest.fail(f"Test failed: {str(e)}")
|
||||
|
|
@ -2436,11 +2436,12 @@ async def test_filter_tools_by_allowed_tools_integration():
|
|||
mock_client_constructor,
|
||||
):
|
||||
# Call _get_tools_from_mcp_servers which should apply the filtering
|
||||
filtered_tools = await _get_tools_from_mcp_servers(
|
||||
listing = await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=mock_user_auth,
|
||||
mcp_auth_header="Bearer test_token",
|
||||
mcp_servers=None, # Get from all servers
|
||||
)
|
||||
filtered_tools = listing.tools
|
||||
|
||||
# Verify that only allowed tools are returned
|
||||
assert (
|
||||
|
|
@ -2549,11 +2550,12 @@ async def test_filter_tools_by_disallowed_tools_integration():
|
|||
mock_client_constructor,
|
||||
):
|
||||
# Call _get_tools_from_mcp_servers which should apply the filtering
|
||||
filtered_tools = await _get_tools_from_mcp_servers(
|
||||
listing = await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=mock_user_auth,
|
||||
mcp_auth_header="Bearer test_token",
|
||||
mcp_servers=None, # Get from all servers
|
||||
)
|
||||
filtered_tools = listing.tools
|
||||
|
||||
# Verify that only safe tools are returned (dangerous tools filtered out)
|
||||
assert (
|
||||
|
|
@ -2650,11 +2652,12 @@ async def test_filter_tools_no_restrictions_integration():
|
|||
mock_client_constructor,
|
||||
):
|
||||
# Call _get_tools_from_mcp_servers which should apply the filtering
|
||||
filtered_tools = await _get_tools_from_mcp_servers(
|
||||
listing = await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=mock_user_auth,
|
||||
mcp_auth_header="Bearer test_token",
|
||||
mcp_servers=None, # Get from all servers
|
||||
)
|
||||
filtered_tools = listing.tools
|
||||
|
||||
# Should return all tools when no restrictions
|
||||
assert (
|
||||
|
|
|
|||
|
|
@ -1820,7 +1820,7 @@ def test_init_auto_router_deployment_success(mock_auto_router, model_list):
|
|||
|
||||
# Verify the auto-router was added to the router's auto_routers dict
|
||||
assert "test-auto-router" in router.auto_routers
|
||||
assert router.auto_routers["test-auto-router"] == mock_auto_router_instance
|
||||
assert router.auto_routers["test-auto-router"][0].strategy == mock_auto_router_instance
|
||||
|
||||
|
||||
@patch("litellm.router_strategy.auto_router.auto_router.AutoRouter")
|
||||
|
|
@ -1833,7 +1833,11 @@ def test_init_auto_router_deployment_duplicate_model_name(mock_auto_router, mode
|
|||
mock_auto_router.return_value = mock_auto_router_instance
|
||||
|
||||
# Add an existing auto-router
|
||||
router.auto_routers["test-auto-router"] = mock_auto_router_instance
|
||||
from litellm.types.router import TaggedPreRoutingStrategy
|
||||
|
||||
router.auto_routers["test-auto-router"] = [
|
||||
TaggedPreRoutingStrategy(tags=(), strategy=mock_auto_router_instance)
|
||||
]
|
||||
|
||||
# Try to add another auto-router with the same name
|
||||
litellm_params = LiteLLM_Params(
|
||||
|
|
@ -1849,7 +1853,7 @@ def test_init_auto_router_deployment_duplicate_model_name(mock_auto_router, mode
|
|||
)
|
||||
|
||||
with pytest.raises(
|
||||
ValueError, match="Auto-router deployment test-auto-router already exists"
|
||||
ValueError, match="Auto-router deployment test-auto-router with tags .* already exists"
|
||||
):
|
||||
router.init_auto_router_deployment(deployment)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,9 @@ import copy
|
|||
import datetime
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
import unittest
|
||||
from typing import List, Optional, Tuple
|
||||
from unittest.mock import ANY, MagicMock, Mock, patch
|
||||
|
|
@ -1533,3 +1535,242 @@ class TestApplyToAnthropicMessagesRequest:
|
|||
sys_blocks = sum(1 for b in (result_sys or []) if isinstance(b, dict) and b.get("cache_control") is not None)
|
||||
total_blocks = sys_blocks + sum(AnthropicCacheControlHook._count_cache_control_blocks(m) for m in result_msgs)
|
||||
assert total_blocks <= 4
|
||||
|
||||
|
||||
class TestEnableAnthropicPromptCaching:
|
||||
"""Auto-injected default breakpoints via litellm.enable_anthropic_prompt_caching."""
|
||||
|
||||
MESSAGES: List[AllMessageValues] = [
|
||||
{"role": "system", "content": "a long system prompt"},
|
||||
{"role": "user", "content": "first turn"},
|
||||
{"role": "assistant", "content": "a reply"},
|
||||
{"role": "user", "content": "latest turn"},
|
||||
]
|
||||
|
||||
def _points(self, model="claude-sonnet-4-5", provider="anthropic", messages=None, system=None, tools=None):
|
||||
return AnthropicCacheControlHook.get_default_injection_points(
|
||||
messages=copy.deepcopy(self.MESSAGES) if messages is None else messages,
|
||||
system=system,
|
||||
model=model,
|
||||
custom_llm_provider=provider,
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
def test_disabled_by_default(self):
|
||||
assert litellm.enable_anthropic_prompt_caching is False
|
||||
assert self._points() == []
|
||||
|
||||
def test_injects_system_and_trailing_turn(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert self._points() == [
|
||||
{"location": "message", "role": "system", "index": None, "control": {"type": "ephemeral"}},
|
||||
{"location": "message", "role": None, "index": -1, "control": {"type": "ephemeral"}},
|
||||
]
|
||||
|
||||
def test_bedrock_claude_is_injected(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
points = self._points(model="us.anthropic.claude-sonnet-4-5-20250929-v1:0", provider="bedrock")
|
||||
assert [p["index"] for p in points] == [None, -1]
|
||||
|
||||
@pytest.mark.parametrize("model, provider", [("gpt-4o", "openai"), ("gemini-2.0-flash", "gemini")])
|
||||
def test_non_anthropic_providers_never_injected(self, monkeypatch, model, provider):
|
||||
"""These report supports_prompt_caching=True but never consume cache_control markers."""
|
||||
from litellm.utils import supports_prompt_caching
|
||||
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert supports_prompt_caching(model=model, custom_llm_provider=provider) is True
|
||||
assert self._points(model=model, provider=provider) == []
|
||||
|
||||
def test_model_without_caching_support_not_injected(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert self._points(model="anthropic.claude-3-5-sonnet-20240620-v1:0", provider="bedrock") == []
|
||||
|
||||
def test_stands_down_when_client_sent_cache_control(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
messages = [
|
||||
{"role": "system", "content": [{"type": "text", "text": "s", "cache_control": {"type": "ephemeral"}}]},
|
||||
{"role": "user", "content": "latest turn"},
|
||||
]
|
||||
assert self._points(messages=messages) == []
|
||||
|
||||
def test_stands_down_when_system_block_has_cache_control(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
system = [{"type": "text", "text": "s", "cache_control": {"type": "ephemeral"}}]
|
||||
assert self._points(messages=[{"role": "user", "content": "hi"}], system=system) == []
|
||||
|
||||
@staticmethod
|
||||
def _tools(count: int, cached: bool) -> List[dict]:
|
||||
tool: dict = {"type": "function", "function": {"name": "t", "description": "d", "parameters": {}}}
|
||||
if cached:
|
||||
tool["cache_control"] = {"type": "ephemeral"}
|
||||
return [{**tool, "function": {**tool["function"], "name": f"t{i}"}} for i in range(count)]
|
||||
|
||||
def test_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
|
||||
"""Caching just the tool definitions is a normal client pattern, and those
|
||||
breakpoints count toward the provider's four-block limit. Three of them plus
|
||||
our two would be five, which Anthropic rejects outright."""
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert self._points(tools=self._tools(3, cached=True)) == []
|
||||
|
||||
def test_injects_when_tools_carry_no_cache_control(self, monkeypatch):
|
||||
"""Tools alone must not suppress injection; only client-marked ones do."""
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert [p["index"] for p in self._points(tools=self._tools(3, cached=False))] == [None, -1]
|
||||
|
||||
@pytest.mark.parametrize("tools", [None, []])
|
||||
def test_absent_tools_do_not_suppress_injection(self, monkeypatch, tools):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert [p["index"] for p in self._points(tools=tools)] == [None, -1]
|
||||
|
||||
def test_seed_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
|
||||
"""Same guard on the /chat/completions seeding path."""
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
params: dict = {}
|
||||
AnthropicCacheControlHook.maybe_seed_default_injection_points(
|
||||
non_default_params=params,
|
||||
messages=copy.deepcopy(self.MESSAGES),
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
tools=self._tools(3, cached=True),
|
||||
)
|
||||
assert "cache_control_injection_points" not in params
|
||||
|
||||
def test_v1_messages_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
|
||||
"""Same guard on the /v1/messages path, where tools reach the hook directly."""
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
messages = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]
|
||||
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
|
||||
copy.deepcopy(messages),
|
||||
"sys",
|
||||
{},
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
tools=self._tools(3, cached=True),
|
||||
)
|
||||
assert result_sys == "sys"
|
||||
assert result_msgs == messages
|
||||
|
||||
def test_default_ttl_is_anthropics_five_minute_cache(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert all(p["control"] == {"type": "ephemeral"} for p in self._points())
|
||||
|
||||
@pytest.mark.parametrize("ttl", ["5m", "1h"])
|
||||
def test_ttl_override_applied(self, monkeypatch, ttl):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
monkeypatch.setattr(litellm, "anthropic_prompt_caching_ttl", ttl)
|
||||
assert all(p["control"] == {"type": "ephemeral", "ttl": ttl} for p in self._points())
|
||||
|
||||
def test_seed_does_not_override_configured_points(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
configured = [{"location": "message", "role": "user", "index": 0}]
|
||||
params = {"cache_control_injection_points": configured}
|
||||
AnthropicCacheControlHook.maybe_seed_default_injection_points(
|
||||
non_default_params=params,
|
||||
messages=copy.deepcopy(self.MESSAGES),
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
assert params["cache_control_injection_points"] is configured
|
||||
|
||||
def test_seed_adds_defaults_when_enabled(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
params: dict = {}
|
||||
AnthropicCacheControlHook.maybe_seed_default_injection_points(
|
||||
non_default_params=params,
|
||||
messages=copy.deepcopy(self.MESSAGES),
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
assert [p["index"] for p in params["cache_control_injection_points"]] == [None, -1]
|
||||
|
||||
def test_seed_is_noop_when_disabled(self):
|
||||
params: dict = {}
|
||||
AnthropicCacheControlHook.maybe_seed_default_injection_points(
|
||||
non_default_params=params,
|
||||
messages=copy.deepcopy(self.MESSAGES),
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
assert params == {}
|
||||
|
||||
def test_v1_messages_applies_defaults_end_to_end(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
messages = [
|
||||
{"role": "user", "content": [{"type": "text", "text": "first"}]},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "reply"}]},
|
||||
{"role": "user", "content": [{"type": "text", "text": "latest"}]},
|
||||
]
|
||||
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
|
||||
messages,
|
||||
"a system prompt",
|
||||
{},
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
|
||||
assert result_sys == [{"type": "text", "text": "a system prompt", "cache_control": {"type": "ephemeral"}}]
|
||||
assert result_msgs[-1]["content"][-1]["cache_control"] == {"type": "ephemeral"}
|
||||
assert "cache_control" not in result_msgs[0]["content"][-1]
|
||||
|
||||
def test_v1_messages_is_noop_when_disabled(self):
|
||||
messages = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]
|
||||
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
|
||||
messages,
|
||||
"sys",
|
||||
{},
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
|
||||
assert result_sys == "sys"
|
||||
assert result_msgs == messages
|
||||
|
||||
|
||||
class TestAnthropicPromptCachingEnvVars:
|
||||
"""Both settings are read from the environment at import, so an admin can enable
|
||||
auto-caching without a config file. Each case re-imports litellm in a subprocess
|
||||
so the env is read fresh without contaminating this process's module graph.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _import_litellm_with_env(env_override: dict) -> Tuple[bool, Optional[str]]:
|
||||
env = os.environ.copy()
|
||||
env.pop("LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING", None)
|
||||
env.pop("LITELLM_ANTHROPIC_PROMPT_CACHING_TTL", None)
|
||||
env.update(env_override)
|
||||
script = textwrap.dedent(
|
||||
"""
|
||||
import json, litellm
|
||||
print(json.dumps([litellm.enable_anthropic_prompt_caching, litellm.anthropic_prompt_caching_ttl]))
|
||||
"""
|
||||
)
|
||||
result = subprocess.run(
|
||||
[sys.executable, "-c", script], capture_output=True, text=True, env=env, timeout=300
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
enabled, ttl = json.loads(result.stdout.strip().splitlines()[-1])
|
||||
return enabled, ttl
|
||||
|
||||
def test_unset_env_leaves_auto_caching_off(self):
|
||||
assert self._import_litellm_with_env({}) == (False, None)
|
||||
|
||||
@pytest.mark.parametrize("value", ["true", "True", "TRUE"])
|
||||
def test_env_enables_auto_caching_case_insensitively(self, value):
|
||||
enabled, _ = self._import_litellm_with_env({"LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING": value})
|
||||
assert enabled is True
|
||||
|
||||
@pytest.mark.parametrize("value", ["false", "0", "yes", ""])
|
||||
def test_env_only_enables_on_true(self, value):
|
||||
enabled, _ = self._import_litellm_with_env({"LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING": value})
|
||||
assert enabled is False
|
||||
|
||||
@pytest.mark.parametrize("value", ["5m", "1h"])
|
||||
def test_ttl_env_is_applied(self, value):
|
||||
_, ttl = self._import_litellm_with_env({"LITELLM_ANTHROPIC_PROMPT_CACHING_TTL": value})
|
||||
assert ttl == value
|
||||
|
||||
@pytest.mark.parametrize("value", ["10m", "1H", "3600", "ephemeral"])
|
||||
def test_unsupported_ttl_env_falls_back_to_provider_default(self, value):
|
||||
"""An unparseable TTL must fall back to Anthropic's 5m default, never reach the provider verbatim."""
|
||||
_, ttl = self._import_litellm_with_env({"LITELLM_ANTHROPIC_PROMPT_CACHING_TTL": value})
|
||||
assert ttl is None
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue