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

# Conflicts:
#	type-discipline-budget.json
This commit is contained in:
mateo-berri 2026-08-18 16:11:12 -07:00
commit c435c25da2
259 changed files with 21880 additions and 7178 deletions

View file

@ -57,7 +57,7 @@
"limit": 5681
},
"reportMissingTypeArgument": {
"limit": 15606
"limit": 15605
},
"reportMissingTypeStubs": {
"limit": 40
@ -108,7 +108,7 @@
"limit": 39154
},
"reportUnknownParameterType": {
"limit": 19945
"limit": 19944
},
"reportUnknownVariableType": {
"limit": 30772
@ -132,7 +132,7 @@
"limit": 27
},
"reportUnusedClass": {
"limit": 22
"limit": 21
},
"reportUnusedFunction": {
"limit": 139

View file

@ -41,6 +41,11 @@ OBJECT_KEYS: dict[str, JsonSchema] = {
},
"additionalProperties": False,
},
"guardrail_cost_per_unit": {
"type": "object",
"description": "USD cost per billable guardrail unit, keyed by the provider's usage counter name (e.g. Bedrock's contentPolicyUnits).",
"additionalProperties": NONNEG_NUMBER,
},
"metadata": {
"type": "object",
"description": "Free-form notes about the entry (e.g. pricing derivation).",

View file

@ -792,6 +792,8 @@ def _populate_provider_model_sets(model_cost_map: Dict) -> None:
nlp_cloud_models.add(key)
elif value.get("litellm_provider") == "aleph_alpha":
aleph_alpha_models.add(key)
elif value.get("litellm_provider") == "bedrock" and value.get("mode") == "guardrail":
pass
elif value.get("litellm_provider") == "bedrock" and not is_bedrock_pricing_only_model(key):
bedrock_models.add(key)
elif value.get("litellm_provider") == "bedrock_converse":

View file

@ -107,7 +107,7 @@ async def acreate_batch(
completion_window: Literal["24h"],
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"],
input_file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy"] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -157,7 +157,7 @@ def create_batch(
completion_window: Literal["24h"],
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"],
input_file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy"] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -339,7 +339,9 @@ def create_batch(
@client
async def aretrieve_batch(
batch_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai",
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic"
] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -385,7 +387,9 @@ def _handle_retrieve_batch_providers_without_provider_config(
litellm_params: dict,
_retrieve_batch_request: RetrieveBatchRequest,
_is_async: bool,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai",
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic"
] = "openai",
logging_obj: Any | None = None,
):
api_base: str | None = None
@ -508,7 +512,9 @@ def _handle_retrieve_batch_providers_without_provider_config(
@client
def retrieve_batch(
batch_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai",
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic"
] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -826,7 +832,7 @@ def list_batches(
async def acancel_batch(
batch_id: str,
model: str | None = None,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy"] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -872,7 +878,7 @@ async def acancel_batch(
def cancel_batch(
batch_id: str,
model: str | None = None,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock"] | str = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy"] | str = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,

View file

@ -12,8 +12,11 @@ This module is dependency-injected: callers pass the proxy ``llm_router`` and
from __future__ import annotations
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any, Final
import litellm
if TYPE_CHECKING:
from litellm.router import Router
@ -41,3 +44,28 @@ def build_router_embedding_metadata(
metadata: Final[dict[str, Any]] = dict(request_metadata or {})
metadata["semantic-cache-embedding"] = True
return metadata
def resolve_embedding_max_input_tokens(
configured_max_input_tokens: int | None,
embedding_model: str,
router: Router | None,
) -> int | None:
"""Explicit cache setting first, else the Router deployment's configured ``max_input_tokens``."""
if configured_max_input_tokens is not None:
return configured_max_input_tokens
if router is None:
return None
deployment_max_input_tokens, _ = router.get_configured_token_limits(embedding_model)
return deployment_max_input_tokens
def truncate_embedding_input(prompt: str, embedding_model: str, max_input_tokens: int | None) -> str:
"""Keep only the first ``max_input_tokens`` tokens of ``prompt`` for the embedding call."""
if max_input_tokens is None:
return prompt
tokens: Final[Sequence[int]] = litellm.encode(model=embedding_model, text=prompt)
if len(tokens) <= max_input_tokens:
return prompt
truncated: Final[str] = litellm.decode(model=embedding_model, tokens=tokens[:max_input_tokens])
return truncated

View file

@ -97,6 +97,7 @@ class Cache:
qdrant_quantization_config: str | None = None,
qdrant_semantic_cache_embedding_model: str = "text-embedding-ada-002",
qdrant_semantic_cache_vector_size: int | None = None,
semantic_cache_embedding_max_input_tokens: int | None = None,
# GCP IAM authentication parameters
gcp_service_account: str | None = None,
gcp_ssl_ca_certs: str | None = None,
@ -122,6 +123,7 @@ class Cache:
qdrant_api_key (str, optional): The api_key for the local or cloud qdrant cluster.
qdrant_collection_name (str, optional): The name for your qdrant collection. Required if type is "qdrant-semantic".
similarity_threshold (float, optional): The similarity threshold for semantic-caching, Required if type is "redis-semantic" or "qdrant-semantic".
semantic_cache_embedding_max_input_tokens (int, optional): Truncate prompts to this many tokens before embedding them for semantic caching. Defaults to the embedding deployment's configured max_input_tokens.
# Disk Cache Args
disk_cache_dir (str, optional): The directory for the disk cache. Defaults to None.
@ -192,6 +194,7 @@ class Cache:
similarity_threshold=similarity_threshold,
embedding_model=redis_semantic_cache_embedding_model,
index_name=redis_semantic_cache_index_name,
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
**kwargs,
)
elif type == LiteLLMCacheType.VALKEY_SEMANTIC:
@ -207,6 +210,7 @@ class Cache:
embedding_model=valkey_semantic_cache_embedding_model,
index_name=valkey_semantic_cache_index_name,
startup_nodes=redis_startup_nodes,
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
**kwargs,
)
elif type == LiteLLMCacheType.QDRANT_SEMANTIC:
@ -218,6 +222,7 @@ class Cache:
quantization_config=qdrant_quantization_config,
embedding_model=qdrant_semantic_cache_embedding_model,
vector_size=qdrant_semantic_cache_vector_size,
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
)
elif type == LiteLLMCacheType.LOCAL:
self.cache = InMemoryCache()

View file

@ -12,7 +12,7 @@ import ast
import asyncio
import json
import os
from typing import Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, cast
import litellm
from litellm._logging import print_verbose
@ -22,12 +22,21 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
)
from litellm.types.utils import EmbeddingResponse
from ._embedding_router import build_router_embedding_metadata, resolve_embedding_router
from ._embedding_router import (
build_router_embedding_metadata,
resolve_embedding_max_input_tokens,
resolve_embedding_router,
truncate_embedding_input,
)
from .base_cache import BaseCache
if TYPE_CHECKING:
from litellm.router import Router
class QdrantSemanticCache(BaseCache):
CACHE_KEY_FIELD_NAME = "litellm_cache_key"
embedding_max_input_tokens: int | None = None
def __init__(
self,
@ -39,6 +48,7 @@ class QdrantSemanticCache(BaseCache):
embedding_model="text-embedding-ada-002",
host_type=None,
vector_size=None,
embedding_max_input_tokens: int | None = None,
):
from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
@ -57,6 +67,7 @@ class QdrantSemanticCache(BaseCache):
raise Exception("similarity_threshold must be provided, passed None")
self.similarity_threshold = similarity_threshold
self.embedding_model = embedding_model
self.embedding_max_input_tokens = embedding_max_input_tokens
self.vector_size = vector_size if vector_size is not None else QDRANT_VECTOR_SIZE
headers = {}
@ -188,6 +199,13 @@ class QdrantSemanticCache(BaseCache):
cached_key: Final = payload.get(self.CACHE_KEY_FIELD_NAME)
return cached_key is not None and str(cached_key) == str(key)
def _embedding_input(self, prompt: str, router: "Router | None") -> str:
return truncate_embedding_input(
prompt,
self.embedding_model,
resolve_embedding_max_input_tokens(self.embedding_max_input_tokens, self.embedding_model, router),
)
def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse:
"""Embed via the proxy Router when it serves the model, else direct."""
try:
@ -197,16 +215,17 @@ class QdrantSemanticCache(BaseCache):
llm_router = None
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
embedding_input: Final = self._embedding_input(prompt, router)
if router is not None:
return router.embedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
metadata=build_router_embedding_metadata(metadata),
)
return litellm.embedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
)
@ -218,17 +237,18 @@ class QdrantSemanticCache(BaseCache):
llm_router = None
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
embedding_input: Final = self._embedding_input(prompt, router)
if router is not None:
return await router.aembedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
metadata=build_router_embedding_metadata(metadata),
)
return await litellm.aembedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
)

View file

@ -14,7 +14,7 @@ import asyncio
import json
import os
from collections.abc import Callable, Mapping
from typing import Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, cast
import litellm
from litellm._logging import print_verbose, verbose_logger
@ -23,9 +23,17 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
)
from litellm.types.utils import EmbeddingResponse
from ._embedding_router import build_router_embedding_metadata, resolve_embedding_router
from ._embedding_router import (
build_router_embedding_metadata,
resolve_embedding_max_input_tokens,
resolve_embedding_router,
truncate_embedding_input,
)
from .base_cache import BaseCache
if TYPE_CHECKING:
from litellm.router import Router
class RedisSemanticCache(BaseCache):
"""
@ -38,6 +46,7 @@ class RedisSemanticCache(BaseCache):
DEFAULT_REDIS_INDEX_NAME: str = "litellm_semantic_cache_index"
CACHE_KEY_FIELD_NAME: str = "litellm_cache_key"
embedding_max_input_tokens: int | None = None
def __init__(
self,
@ -48,6 +57,7 @@ class RedisSemanticCache(BaseCache):
similarity_threshold: float | None = None,
embedding_model: str = "text-embedding-ada-002",
index_name: str | None = None,
embedding_max_input_tokens: int | None = None,
**kwargs: object,
):
"""
@ -62,6 +72,8 @@ class RedisSemanticCache(BaseCache):
where 1.0 requires exact matches and 0.0 accepts any match
embedding_model: Model to use for generating embeddings
index_name: Name for the Redis index
embedding_max_input_tokens: Truncate prompts to this many tokens before
embedding; defaults to the Router deployment's configured max_input_tokens
ttl: Default time-to-live for cache entries in seconds
**kwargs: Additional arguments passed to the Redis client
@ -86,6 +98,7 @@ class RedisSemanticCache(BaseCache):
# While similarity: 1 = most similar, 0 = least similar
self.distance_threshold = 1 - similarity_threshold
self.embedding_model = embedding_model
self.embedding_max_input_tokens = embedding_max_input_tokens
# Set up Redis connection
if redis_url is None:
@ -307,6 +320,13 @@ class RedisSemanticCache(BaseCache):
return dict_method()
return value
def _embedding_input(self, prompt: str, router: "Router | None") -> str:
return truncate_embedding_input(
prompt,
self.embedding_model,
resolve_embedding_max_input_tokens(self.embedding_max_input_tokens, self.embedding_model, router),
)
def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> list[float]:
"""
Routes through the proxy Router when the embedding model is a Router
@ -320,12 +340,13 @@ class RedisSemanticCache(BaseCache):
llm_router = None
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
embedding_input: Final = self._embedding_input(prompt, router)
if router is not None:
embedding_response = cast(
EmbeddingResponse,
router.embedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
metadata=build_router_embedding_metadata(metadata),
),
@ -335,7 +356,7 @@ class RedisSemanticCache(BaseCache):
EmbeddingResponse,
litellm.embedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
),
)
@ -490,18 +511,19 @@ class RedisSemanticCache(BaseCache):
llm_router = None
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
embedding_input: Final = self._embedding_input(prompt, router)
try:
if router is not None:
embedding_response = await router.aembedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
metadata=build_router_embedding_metadata(metadata),
)
else:
embedding_response = await litellm.aembedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
)
return embedding_response["data"][0]["embedding"]

View file

@ -17,7 +17,6 @@ RedisSemanticCache since those are backend agnostic.
import asyncio
import hashlib
import os
import struct
from dataclasses import dataclass
from typing import Any, Final
@ -29,6 +28,7 @@ from redis.commands.search.query import Query
from litellm._logging import print_verbose
from litellm._uuid import uuid
from litellm.llms.valkey.common_utils import build_valkey_url, pack_vector
from .redis_semantic_cache import RedisSemanticCache
@ -61,6 +61,7 @@ class ValkeySemanticCache(RedisSemanticCache):
startup_nodes: list | None = None,
sync_client: Redis | None = None,
async_client: AsyncRedis | None = None,
embedding_max_input_tokens: int | None = None,
**kwargs: Any,
):
if similarity_threshold is None:
@ -78,6 +79,7 @@ class ValkeySemanticCache(RedisSemanticCache):
self.similarity_threshold = similarity_threshold
self.embedding_model = embedding_model
self.embedding_max_input_tokens = embedding_max_input_tokens
self.index_name = index_name or self.DEFAULT_VALKEY_INDEX_NAME
self.key_prefix = f"{self.index_name}:"
self._index_dim: int | None = None
@ -92,19 +94,17 @@ class ValkeySemanticCache(RedisSemanticCache):
@staticmethod
def _build_valkey_url(host: str | None, port: str | None, password: str | None, ssl: bool = False) -> str:
host = host or os.environ.get("VALKEY_HOST") or os.environ.get("REDIS_HOST")
port = port or os.environ.get("VALKEY_PORT") or os.environ.get("REDIS_PORT")
password = password or os.environ.get("VALKEY_PASSWORD") or os.environ.get("REDIS_PASSWORD")
resolved_host: Final = host or os.environ.get("VALKEY_HOST") or os.environ.get("REDIS_HOST")
resolved_port: Final = port or os.environ.get("VALKEY_PORT") or os.environ.get("REDIS_PORT")
resolved_password: Final = password or os.environ.get("VALKEY_PASSWORD") or os.environ.get("REDIS_PASSWORD")
if not host or not port:
if not resolved_host or not resolved_port:
raise ValueError(
"Missing required Valkey configuration. Provide host and port "
"(or VALKEY_HOST/VALKEY_PORT), or pass redis_url."
)
credentials: Final = f":{password}@" if password else ""
scheme: Final = "rediss" if ssl else "redis"
return f"{scheme}://{credentials}{host}:{port}"
return build_valkey_url(host=resolved_host, port=resolved_port, password=resolved_password, ssl=ssl)
@classmethod
def _scope_tag(cls, key: str) -> str:
@ -116,7 +116,7 @@ class ValkeySemanticCache(RedisSemanticCache):
@staticmethod
def _embedding_to_bytes(embedding: list[float]) -> bytes:
return struct.pack(f"<{len(embedding)}f", *embedding)
return pack_vector(embedding)
def _index_schema(self, dim: int) -> tuple[TagField, VectorField]:
return (

View file

@ -1,5 +1,6 @@
import os
import sys
from types import MappingProxyType
from typing import Final, Literal
from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_or_none
@ -1764,3 +1765,7 @@ PTU_LAPSED_ALERT_LIMIT: Final[int] = 10
# one run delete a charge another just wrote. A stale row is hours old and a concurrent
# one is seconds old, so a few minutes separates them.
PTU_PRUNE_SKEW_GRACE_SECONDS: Final[int] = 300
# Shared read-only empty mapping, for defaulting optional Mapping parameters without
# constructing a fresh mutable dict at each call site.
EMPTY_MAPPING: Final = MappingProxyType({})

View file

@ -23,12 +23,15 @@ FileCreateProvider = Literal[
"vertex_ai",
"bedrock",
"hosted_vllm",
"litellm_proxy",
"manus",
"anthropic",
]
FileRetrieveProvider = Literal["openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "manus", "anthropic"]
FileDeleteProvider = Literal["openai", "azure", "gemini", "manus", "anthropic"]
FileListProvider = Literal["openai", "azure", "manus", "anthropic"]
FileRetrieveProvider = Literal[
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic"
]
FileDeleteProvider = Literal["openai", "azure", "gemini", "litellm_proxy", "manus", "anthropic"]
FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic"]
import litellm
from litellm import get_secret_str
from litellm.files.streaming import FileContentStreamingResponse

View file

@ -1,7 +1,9 @@
from collections.abc import AsyncIterator, Iterator
from typing import Literal, NamedTuple
FileContentProvider = Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic", "manus"]
FileContentProvider = Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "manus"
]
class FileContentStreamingResult(NamedTuple):

View file

@ -64,6 +64,10 @@ from litellm.integrations.mlflow import MlflowLogger
from litellm.integrations.sqs import SQSLogger
from litellm.litellm_core_utils.core_helpers import reconstruct_model_name
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import (
cost_breakdown_with_guardrail,
guardrail_information_cost,
)
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
StandardBuiltInToolCostTracking,
)
@ -5650,12 +5654,14 @@ def get_standard_logging_object_payload(
base_model = metadata.get("deployment")
custom_pricing: Final = use_custom_pricing_for_model(litellm_params=litellm_params)
raw_response_cost: Final = kwargs.get("response_cost")
response_cost: Final[float] = raw_response_cost or 0.0
llm_response_cost: Final[float] = raw_response_cost or 0.0
guardrail_cost: Final = guardrail_information_cost(metadata.get("standard_logging_guardrail_information"))
response_cost: Final[float] = llm_response_cost + guardrail_cost
# clean up litellm hidden params
clean_hidden_params: Final = StandardLoggingPayloadSetup.get_hidden_params(hidden_params)
if clean_hidden_params["response_cost"] is None and raw_response_cost is not None:
clean_hidden_params["response_cost"] = response_cost
clean_hidden_params["response_cost"] = llm_response_cost
model_cost_information: Final = StandardLoggingPayloadSetup.get_model_cost_information(
base_model=base_model,
@ -5735,7 +5741,7 @@ def get_standard_logging_object_payload(
metadata=clean_metadata,
cache_key=clean_hidden_params["cache_key"],
response_cost=response_cost,
cost_breakdown=logging_obj.cost_breakdown,
cost_breakdown=cost_breakdown_with_guardrail(logging_obj.cost_breakdown, guardrail_cost),
total_tokens=usage_dict.get("total_tokens", 0),
prompt_tokens=usage_dict.get("prompt_tokens", 0),
completion_tokens=usage_dict.get("completion_tokens", 0),

View file

@ -0,0 +1,78 @@
import math
from collections.abc import Mapping
from typing import Final
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
import litellm
from litellm._logging import verbose_logger
from litellm.types.utils import CostBreakdown
BEDROCK_GUARDRAIL_PRICING_KEY: Final = "bedrock/guardrails"
class GuardrailPricing(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
guardrail_cost_per_unit: Mapping[str, float]
class GuardrailCostEntry(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
guardrail_cost: float | None = None
GuardrailInformationShape = tuple[GuardrailCostEntry, ...] | GuardrailCostEntry | None
_GUARDRAIL_INFORMATION_ADAPTER: Final[TypeAdapter[GuardrailInformationShape]] = TypeAdapter(GuardrailInformationShape)
def _bedrock_guardrail_pricing(aws_region_name: str | None) -> GuardrailPricing | None:
regional_key: Final = f"bedrock/{aws_region_name}/guardrails" if aws_region_name else None
for key in (regional_key, BEDROCK_GUARDRAIL_PRICING_KEY):
if key is None or key not in litellm.model_cost:
continue
try:
return GuardrailPricing.model_validate(litellm.model_cost[key])
except ValidationError as e:
verbose_logger.warning("Ignoring malformed guardrail pricing entry %s: %s", key, e)
return None
def bedrock_guardrail_cost(usage_units: Mapping[str, int], aws_region_name: str | None) -> float:
pricing: Final = _bedrock_guardrail_pricing(aws_region_name)
if pricing is None:
return 0.0
return sum(units * pricing.guardrail_cost_per_unit.get(counter, 0.0) for counter, units in usage_units.items())
def _billable_entry_cost(entry: GuardrailCostEntry) -> float:
cost: Final = entry.guardrail_cost
if cost is None or not math.isfinite(cost) or cost <= 0.0:
return 0.0
return cost
def guardrail_information_cost(guardrail_information: object) -> float:
try:
parsed: Final = _GUARDRAIL_INFORMATION_ADAPTER.validate_python(guardrail_information)
except ValidationError:
return 0.0
if parsed is None:
return 0.0
if isinstance(parsed, GuardrailCostEntry):
return _billable_entry_cost(parsed)
return sum(_billable_entry_cost(entry) for entry in parsed)
def cost_breakdown_with_guardrail(cost_breakdown: CostBreakdown | None, guardrail_cost: float) -> CostBreakdown | None:
if guardrail_cost <= 0.0:
return cost_breakdown
existing: Final[CostBreakdown] = cost_breakdown if cost_breakdown is not None else CostBreakdown()
merged: Final[CostBreakdown] = {
**existing,
"guardrail_cost": guardrail_cost,
"total_cost": existing.get("total_cost", 0.0) + guardrail_cost,
}
return merged

View file

@ -42,14 +42,18 @@ _VALID_DATA_RESIDENCIES: Final = frozenset(r.value for r in DataResidency)
# Pre-resolved service-tier cost-key suffixes (e.g. "_priority"). Used per
# request in the cost-calc path, so the f-strings are built once here instead
# of being rebuilt for every model_info key on every call.
_SERVICE_TIER_SUFFIXES: Final[tuple[str, ...]] = tuple(f"_{st.value}" for st in ServiceTier)
# of being rebuilt for every model_info key on every call. Longest-first so a
# substring match resolves "_ultrafast" before "_fast".
_SERVICE_TIER_SUFFIXES: Final[tuple[str, ...]] = tuple(
sorted((f"_{st.value}" for st in ServiceTier), key=len, reverse=True)
)
_SERVICE_TIER_TO_COST_KEY_SUFFIX: Final[Mapping[str, str]] = MappingProxyType(
{
ServiceTier.FLEX.value: ServiceTier.FLEX.value,
ServiceTier.PRIORITY.value: ServiceTier.PRIORITY.value,
ServiceTier.FAST.value: ServiceTier.PRIORITY.value,
ServiceTier.ULTRAFAST.value: ServiceTier.ULTRAFAST.value,
}
)
@ -191,7 +195,7 @@ def _get_service_tier_cost_key(base_key: str, service_tier: str | None) -> str:
Args:
base_key: The base cost key (e.g., "input_cost_per_token")
service_tier: The service tier ("flex", "priority", "fast", or None for standard)
service_tier: The service tier ("flex", "priority", "fast", "ultrafast", or None for standard)
Returns:
str: The cost key to use (e.g., "input_cost_per_token_flex" or "input_cost_per_token")

View file

@ -624,6 +624,12 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
return self.chunk_queue.popleft()
if processed_chunk["type"] == "content_block_delta" and not self._delta_has_content(processed_chunk):
# A tool_use block opens with empty arguments (Bedrock Converse's
# ``contentBlockStart``, OpenAI's ``arguments: ""``), so flush the
# block start queued above instead of waiting for the next upstream
# chunk, which on a trailing-burst provider is the whole generation.
if self.chunk_queue:
return self.chunk_queue.popleft()
continue
if processed_chunk["type"] == "message_delta" and self.sent_content_block_finish is False:
@ -847,6 +853,9 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
if processed_chunk["type"] == "content_block_delta" and not self._delta_has_content(
processed_chunk
):
# See ``__next__``: flush the queued block start (issue #32004).
if self.chunk_queue:
return self.chunk_queue.popleft()
continue
if processed_chunk["type"] == "message_delta" and self.sent_content_block_finish is False:

View file

@ -16,7 +16,7 @@ How it works:
import uuid
from collections.abc import AsyncIterator
from typing import Any, Final
from typing import TYPE_CHECKING, Any, Final
import litellm
import litellm.constants as _c
@ -28,6 +28,9 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
if TYPE_CHECKING:
from litellm.router import Router
ADVISOR_MAX_USES: Final[int] = _c.ADVISOR_MAX_USES
ADVISOR_NATIVE_PROVIDERS: Final[frozenset] = _c.ADVISOR_NATIVE_PROVIDERS
ADVISOR_TOOL_DESCRIPTION: Final[str] = _c.ADVISOR_TOOL_DESCRIPTION
@ -97,6 +100,14 @@ class AdvisorOrchestrationHandler(MessagesInterceptor):
parent_request_id: Final[str] = str(kwargs.pop("litellm_call_id", None) or uuid.uuid4())
metadata_base: Final[dict] = dict(kwargs.pop("metadata", None) or {})
advisor_metadata: Final = {
**metadata_base,
"advisor_sub_call": True,
"parent_request_id": parent_request_id,
}
advisor_router: Final = (
None if (advisor_api_key or advisor_api_base) else _resolve_advisor_router(advisor_model)
)
iteration = 0
while True:
@ -138,20 +149,27 @@ class AdvisorOrchestrationHandler(MessagesInterceptor):
# --- Advisor sub-call (always non-streaming, no tools) ---
try:
advisor_response: AnthropicMessagesResponse = await _call_messages_handler(
model=advisor_model,
messages=advisor_messages,
tools=None,
stream=False,
max_tokens=max_tokens,
custom_llm_provider=None, # let litellm resolve from model name
metadata={
**metadata_base,
"advisor_sub_call": True,
"parent_request_id": parent_request_id,
},
api_key=advisor_api_key,
api_base=advisor_api_base,
advisor_response: AnthropicMessagesResponse = (
await advisor_router.aanthropic_messages(
model=advisor_model,
messages=advisor_messages,
tools=None,
stream=False,
max_tokens=max_tokens,
metadata=advisor_metadata,
)
if advisor_router is not None
else await _call_messages_handler(
model=advisor_model,
messages=advisor_messages,
tools=None,
stream=False,
max_tokens=max_tokens,
custom_llm_provider=None,
metadata=advisor_metadata,
api_key=advisor_api_key,
api_base=advisor_api_base,
)
)
except Exception as advisor_sub_call_exception:
mark_advisor_orchestration_failure(advisor_sub_call_exception)
@ -284,6 +302,11 @@ def _build_advisor_context(
tool_use blocks are excluded because Anthropic requires tool_use to be
immediately followed by tool_result — not the advisor question.
In-sequence system rows (e.g. Claude Code SessionStart hook output) are
excluded: they are executor-directed, and a trailing one becomes invalid
once the question turn is appended after it (a system row must precede an
assistant message or end the array).
"""
question: Final = (advisor_use_block.get("input") or {}).get("question") or (
"Please provide guidance on the current task."
@ -295,7 +318,7 @@ def _build_advisor_context(
for block in raw_content
if isinstance(block, dict) and block.get("type") == "text"
]
result: Final = list(messages)
result: Final = [m for m in messages if m.get("role") != "system"]
if executor_text_blocks:
result.append({"role": "assistant", "content": executor_text_blocks})
result.append({"role": "user", "content": question})
@ -357,6 +380,24 @@ def _inject_max_uses_error(
]
def _resolve_advisor_router(advisor_model: str) -> "Router | None":
"""Return the proxy router when it serves ``advisor_model`` directly or via a wildcard.
Returns ``None`` for SDK callers (no proxy router) and for advisor models the router
doesn't know about, so those keep resolving through ``litellm.anthropic_messages()``
provider inference.
"""
try:
from litellm.proxy.proxy_server import llm_router
except (ImportError, ModuleNotFoundError):
return None
if llm_router is None:
return None
if llm_router.is_recognized_model(advisor_model) or llm_router.pattern_router.route(advisor_model):
return llm_router
return None
async def _call_messages_handler(
model: str,
messages: list[dict],

View file

@ -1,3 +1,4 @@
import copy
import enum
import re
from typing import Any, Final, cast
@ -11,6 +12,7 @@ from litellm._logging import verbose_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
_audio_or_image_in_message_content,
convert_content_list_to_str,
filter_value_from_dict,
)
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
@ -28,6 +30,9 @@ class AzureFoundryErrorStrings(str, enum.Enum):
SET_EXTRA_PARAMETERS_TO_PASS_THROUGH = "Set extra-parameters to 'pass-through'"
NON_OPENAI_SPEC_MESSAGE_FIELDS: Final = ("thinking_blocks", "provider_specific_fields", "cache_control")
class AzureAIStudioConfig(OpenAIConfig):
def get_supported_openai_params(self, model: str) -> list:
model_supports_tool_choice = True # azure ai supports this by default
@ -167,10 +172,23 @@ class AzureAIStudioConfig(OpenAIConfig):
) -> list:
"""
- Azure AI Studio doesn't support content as a list. This handles:
1. Transforms list content to a string.
2. If message contains an image or audio, send as is (user-intended)
1. Strips message fields that are not part of the OpenAI chat-completions
schema (thinking_blocks, provider_specific_fields, cache_control).
Azure AI Foundry backends set additionalProperties=false and reject
these with "Extra inputs are not permitted", which breaks multi-turn
Anthropic-format clients that echo thinking blocks back as history.
2. Transforms list content to a string.
3. If message contains an image or audio, send as is (user-intended)
Operates on a deep copy so the caller's messages keep their thinking blocks
and provider metadata, which a fallback to another provider still needs.
"""
for message in messages:
stripped_messages: Final = copy.deepcopy(messages)
for message in stripped_messages:
message_dict = cast(dict, message) # cast-ok: TypedDict is a runtime dict stripped on our copy
for field in NON_OPENAI_SPEC_MESSAGE_FIELDS:
filter_value_from_dict(message_dict, field)
# Do nothing if the message contains an image or audio
if _audio_or_image_in_message_content(message):
continue
@ -178,7 +196,7 @@ class AzureAIStudioConfig(OpenAIConfig):
texts = convert_content_list_to_str(message=message)
if texts:
message["content"] = texts
return messages
return stripped_messages
def _is_azure_openai_model(self, model: str, api_base: str | None) -> bool:
try:

View file

@ -1,5 +1,6 @@
from abc import abstractmethod
from typing import TYPE_CHECKING, Any
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, NoReturn
import httpx
@ -154,3 +155,75 @@ class BaseVectorStoreConfig:
response: VectorStoreSearchResponse,
) -> tuple[float, float]:
return 0.0, 0.0
class BaseDirectVectorStoreConfig(BaseVectorStoreConfig):
"""
Base config for vector store providers whose datastore has no HTTP API
(e.g. Valkey over RESP). Instead of transforming to an httpx request, the
config executes the search itself via (a)execute_search_vector_store_request.
"""
@abstractmethod
def execute_search_vector_store_request(
self,
vector_store_id: str,
query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object],
timeout: float | httpx.Timeout | None = None,
) -> VectorStoreSearchResponse:
pass
@abstractmethod
async def aexecute_search_vector_store_request(
self,
vector_store_id: str,
query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object],
timeout: float | httpx.Timeout | None = None,
) -> VectorStoreSearchResponse:
pass
def transform_search_vector_store_request(
self,
vector_store_id: str,
query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
) -> NoReturn:
raise NotImplementedError("Direct vector store providers execute the search themselves; no HTTP request shape")
def transform_search_vector_store_response(
self, response: httpx.Response, litellm_logging_obj: LiteLLMLoggingObj
) -> NoReturn:
raise NotImplementedError("Direct vector store providers execute the search themselves; no HTTP response shape")
def transform_create_vector_store_request(
self,
vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams,
api_base: str,
) -> NoReturn:
raise NotImplementedError
def transform_create_vector_store_response(self, response: httpx.Response) -> NoReturn:
raise NotImplementedError
def get_complete_url(
self,
api_base: str | None,
litellm_params: Mapping[str, object],
) -> str:
return api_base or ""
def get_auth_credentials(self, litellm_params: Mapping[str, object]) -> BaseVectorStoreAuthCredentials:
return BaseVectorStoreAuthCredentials()
def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints:
return VectorStoreIndexEndpoints(read=[], write=[]) # mutable-ok: the TypedDict declares list fields

View file

@ -56,7 +56,10 @@ from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfi
from litellm.llms.base_llm.search.transformation import BaseSearchConfig, SearchResponse
from litellm.llms.base_llm.skills.transformation import BaseSkillsAPIConfig
from litellm.llms.base_llm.text_to_speech.transformation import BaseTextToSpeechConfig
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
from litellm.llms.base_llm.vector_store.transformation import (
BaseDirectVectorStoreConfig,
BaseVectorStoreConfig,
)
from litellm.llms.base_llm.vector_store_files.transformation import (
BaseVectorStoreFilesConfig,
)
@ -915,6 +918,7 @@ class BaseLLMHTTPHandler:
)
if provider_config is None:
raise ValueError(f"Provider {custom_llm_provider} does not support embedding")
embedding_extra_body: Final[Mapping[str, object] | None] = optional_params.pop("extra_body", None)
# get config from model, custom llm provider
headers = provider_config.validate_environment(
api_key=api_key,
@ -939,6 +943,8 @@ class BaseLLMHTTPHandler:
optional_params=optional_params,
headers=headers,
)
if embedding_extra_body:
data.update(embedding_extra_body)
# Some providers (e.g. OCI) require request signing after the body is built.
# The default BaseConfig.sign_request returns (headers, None) — a no-op for
@ -9442,6 +9448,24 @@ class BaseLLMHTTPHandler:
client: HTTPHandler | AsyncHTTPHandler | None = None,
_is_async: bool = False,
) -> VectorStoreSearchResponse:
if isinstance(vector_store_provider_config, BaseDirectVectorStoreConfig):
logging_obj.pre_call(
input="",
api_key="",
additional_args={ # mutable-ok: pre_call's additional_args contract is a dict
"query": query,
"vector_store_id": vector_store_id,
},
)
return await vector_store_provider_config.aexecute_search_vector_store_request(
vector_store_id=vector_store_id,
query=query,
vector_store_search_optional_params=vector_store_search_optional_params,
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params), # mutable-ok: snapshot GenericLiteLLMParams into the Mapping shape
timeout=timeout,
)
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders(custom_llm_provider),
@ -9555,6 +9579,24 @@ class BaseLLMHTTPHandler:
client=client,
)
if isinstance(vector_store_provider_config, BaseDirectVectorStoreConfig):
logging_obj.pre_call(
input="",
api_key="",
additional_args={ # mutable-ok: pre_call's additional_args contract is a dict
"query": query,
"vector_store_id": vector_store_id,
},
)
return vector_store_provider_config.execute_search_vector_store_request(
vector_store_id=vector_store_id,
query=query,
vector_store_search_optional_params=vector_store_search_optional_params,
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params), # mutable-ok: snapshot GenericLiteLLMParams into the Mapping shape
timeout=timeout,
)
if client is None or not isinstance(client, HTTPHandler):
sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)})
else:

View file

@ -29,9 +29,12 @@ def get_fireworks_session_id(litellm_params: dict) -> str | None:
return None
AZURE_FOUNDRY_FIREWORKS_MODEL_ID_PREFIX: Final = "FW-"
def resolve_fireworks_resource_name(model: str) -> str:
stripped: Final = model.removeprefix("fireworks_ai/")
if stripped.startswith("accounts/") or "#" in stripped:
if stripped.startswith(("accounts/", AZURE_FOUNDRY_FIREWORKS_MODEL_ID_PREFIX)) or "#" in stripped:
return stripped
if stripped.startswith(("routers/", "models/")):
return f"accounts/fireworks/{stripped}"

View file

View file

@ -0,0 +1,18 @@
"""Shared helpers for Valkey integrations (semantic cache, vector stores)."""
import struct
from collections.abc import Sequence
from typing import Final
from urllib.parse import quote
def build_valkey_url(host: str, port: str, password: str | None = None, ssl: bool = False) -> str:
"""Deliberately reads no environment: callers of the vector store control the
host, so an env-sourced password would be sent to a caller-chosen server."""
credentials: Final = f":{quote(password, safe='')}@" if password else ""
scheme: Final = "rediss" if ssl else "redis"
return f"{scheme}://{credentials}{host}:{port}"
def pack_vector(embedding: Sequence[float]) -> bytes:
return struct.pack(f"<{len(embedding)}f", *embedding)

View file

@ -0,0 +1,3 @@
from litellm.llms.valkey.vector_stores.transformation import ValkeyVectorStoreConfig
__all__ = ("ValkeyVectorStoreConfig",)

View file

@ -0,0 +1,299 @@
"""
Valkey vector store provider.
Valkey's vector search (the valkey-search module) speaks RESP only, no HTTP
API, so this config extends BaseDirectVectorStoreConfig and executes the
FT.SEARCH KNN query itself via redis-py instead of shaping an httpx request.
Documents are HASHes indexed by an FT index named after the vector_store_id.
"""
from collections.abc import Awaitable, Callable, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, NoReturn
import httpx
from pydantic import BaseModel, ConfigDict
import litellm
from litellm.llms.base_llm.vector_store.transformation import BaseDirectVectorStoreConfig
from litellm.llms.valkey.common_utils import build_valkey_url, pack_vector
from litellm.types.utils import EmbeddingResponse
from litellm.types.vector_stores import (
VectorStoreCreateOptionalRequestParams,
VectorStoreResultContent,
VectorStoreSearchOptionalRequestParams,
VectorStoreSearchResponse,
VectorStoreSearchResult,
)
if TYPE_CHECKING:
from redis import Redis
from redis.asyncio import Redis as AsyncRedis
from redis.commands.search.document import Document
from redis.commands.search.query import Query
from redis.commands.search.result import Result
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
DEFAULT_VALKEY_PORT: Final = 6379
DEFAULT_SOCKET_CONNECT_TIMEOUT_SECONDS: Final = 5.0
DEFAULT_SOCKET_TIMEOUT_SECONDS: Final = 30.0
DEFAULT_MAX_NUM_RESULTS: Final = 10
MIN_MAX_NUM_RESULTS: Final = 1
MAX_MAX_NUM_RESULTS: Final = 50
DEFAULT_EMBEDDING_FIELD_NAME: Final = "embedding"
DEFAULT_TEXT_FIELD_NAME: Final = "text"
DISTANCE_FIELD_NAME: Final = "vector_distance"
_EMPTY_EMBEDDING_CONFIG: Final = MappingProxyType({})
_REDIS_INSTALL_HINT: Final = (
"The Valkey vector store requires the 'redis' package. Run 'pip install redis' to install it."
)
_SEARCH_ONLY_MESSAGE: Final = "Valkey vector store is search-only; create indexes with FT.CREATE directly"
def _import_sync_redis() -> "type[Redis]":
try:
from redis import Redis as SyncRedisClient
except ImportError as e:
raise ValueError(_REDIS_INSTALL_HINT) from e
return SyncRedisClient
def _import_async_redis() -> "type[AsyncRedis]":
try:
from redis.asyncio import Redis as AsyncRedisClient
except ImportError as e:
raise ValueError(_REDIS_INSTALL_HINT) from e
return AsyncRedisClient
def _import_query() -> "type[Query]":
try:
from redis.commands.search.query import Query as RedisQuery
except ImportError as e:
raise ValueError(_REDIS_INSTALL_HINT) from e
return RedisQuery
class _ValkeySearchParams(BaseModel):
"""Typed view over the vector store's litellm_params; unrelated keys are ignored."""
model_config = ConfigDict(frozen=True, extra="ignore")
litellm_embedding_model: str | None = None
litellm_embedding_config: Mapping[str, object] | None = None
valkey_host: str | None = None
valkey_port: int | None = None
valkey_password: str | None = None
valkey_ssl: bool | None = None
valkey_text_field: str | None = None
valkey_embedding_field: str | None = None
@property
def text_field(self) -> str:
return self.valkey_text_field or DEFAULT_TEXT_FIELD_NAME
@property
def embedding_field(self) -> str:
return self.valkey_embedding_field or DEFAULT_EMBEDDING_FIELD_NAME
def require_embedding_model(self) -> str:
if not self.litellm_embedding_model:
raise ValueError(
"litellm_embedding_model is required in litellm_params for the Valkey vector store. "
"Example: litellm_params['litellm_embedding_model'] = 'openai/text-embedding-3-small'"
)
return self.litellm_embedding_model
def connection_url(self) -> str:
if not self.valkey_host:
raise ValueError(
"valkey_host is required in litellm_params for the Valkey vector store. "
"Set it on the vector store's litellm_params, e.g. valkey_host: my-valkey.example.com"
)
return build_valkey_url(
host=self.valkey_host,
port=str(self.valkey_port or DEFAULT_VALKEY_PORT),
password=self.valkey_password,
ssl=bool(self.valkey_ssl),
)
class ValkeyVectorStoreConfig(BaseDirectVectorStoreConfig):
def __init__(
self,
sync_client: "Redis | None" = None,
async_client: "AsyncRedis | None" = None,
embedding_fn: Callable[..., EmbeddingResponse] | None = None,
aembedding_fn: Callable[..., Awaitable[EmbeddingResponse]] | None = None,
) -> None:
super().__init__()
self.sync_client = sync_client
self.async_client = async_client
self.embedding_fn = embedding_fn if embedding_fn is not None else litellm.embedding
self.aembedding_fn = aembedding_fn if aembedding_fn is not None else litellm.aembedding
@staticmethod
def _query_text(query: str | Sequence[str]) -> str:
if isinstance(query, str):
return query
if not query:
raise ValueError("query must not be empty")
return " ".join(query)
@staticmethod
def _socket_timeouts(timeout: float | httpx.Timeout | None) -> tuple[float, float]:
if isinstance(timeout, httpx.Timeout):
return (
timeout.connect or DEFAULT_SOCKET_CONNECT_TIMEOUT_SECONDS,
timeout.read or DEFAULT_SOCKET_TIMEOUT_SECONDS,
)
if timeout is not None:
return (min(float(timeout), DEFAULT_SOCKET_CONNECT_TIMEOUT_SECONDS), float(timeout))
return (DEFAULT_SOCKET_CONNECT_TIMEOUT_SECONDS, DEFAULT_SOCKET_TIMEOUT_SECONDS)
@staticmethod
def _knn_limit(vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams) -> int:
requested: Final = vector_store_search_optional_params.get("max_num_results")
if requested is None:
return DEFAULT_MAX_NUM_RESULTS
if not MIN_MAX_NUM_RESULTS <= requested <= MAX_MAX_NUM_RESULTS:
raise ValueError(
f"max_num_results must be between {MIN_MAX_NUM_RESULTS} and {MAX_MAX_NUM_RESULTS}, got {requested}"
)
return requested
@classmethod
def _knn_query(
cls,
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
embedding_field: str,
text_field: str,
) -> "Query":
if vector_store_search_optional_params.get("filters") is not None:
raise ValueError("Valkey vector store does not support the filters parameter yet")
k: Final = cls._knn_limit(vector_store_search_optional_params)
query_cls: Final = _import_query()
knn_expr: Final = f"*=>[KNN {k} @{embedding_field} $vec AS {DISTANCE_FIELD_NAME}]"
# valkey-search rejects SORTBY on the KNN distance alias, so results are
# re-ordered client-side in _to_response instead.
return query_cls(knn_expr).return_fields(text_field, DISTANCE_FIELD_NAME).paging(0, k).dialect(2)
@staticmethod
def _to_result(doc: "Document", text_field: str) -> VectorStoreSearchResult:
content: Final = [ # mutable-ok: VectorStoreSearchResult declares a list of content parts
VectorStoreResultContent(text=str(getattr(doc, text_field, "")), type="text")
]
return VectorStoreSearchResult(
score=1.0 - float(getattr(doc, DISTANCE_FIELD_NAME)),
content=content,
file_id=getattr(doc, "id", None),
filename=getattr(doc, "id", None),
)
@classmethod
def _to_response(cls, search_result: "Result", query_text: str, text_field: str) -> VectorStoreSearchResponse:
docs: Final = getattr(search_result, "docs", None) or ()
data: Final = sorted(
(cls._to_result(doc, text_field) for doc in docs),
key=lambda result: result.get("score") or 0.0,
reverse=True,
)
return VectorStoreSearchResponse(
object="vector_store.search_results.page",
search_query=query_text,
data=data,
)
def execute_search_vector_store_request(
self,
vector_store_id: str,
query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
litellm_logging_obj: "LiteLLMLoggingObj",
litellm_params: Mapping[str, object],
timeout: float | httpx.Timeout | None = None,
) -> VectorStoreSearchResponse:
params: Final = _ValkeySearchParams.model_validate(litellm_params)
query_text: Final = self._query_text(query)
knn: Final = self._knn_query(
vector_store_search_optional_params,
embedding_field=params.embedding_field,
text_field=params.text_field,
)
embedding_response: Final = self.embedding_fn(
model=params.require_embedding_model(),
input=[query_text], # mutable-ok: litellm.embedding's input contract is a list
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
)
vec_params: Final = {"vec": pack_vector(embedding_response.data[0]["embedding"])} # mutable-ok: redis-py API
if self.sync_client is not None:
raw: Final = self.sync_client.ft(vector_store_id).search(knn, query_params=vec_params)
return self._to_response(raw, query_text, params.text_field)
connect_timeout, op_timeout = self._socket_timeouts(timeout)
client: Final = _import_sync_redis().from_url(
params.connection_url(),
socket_connect_timeout=connect_timeout,
socket_timeout=op_timeout,
)
try:
raw_result: Final = client.ft(vector_store_id).search(knn, query_params=vec_params)
return self._to_response(raw_result, query_text, params.text_field)
finally:
client.close()
async def aexecute_search_vector_store_request(
self,
vector_store_id: str,
query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
litellm_logging_obj: "LiteLLMLoggingObj",
litellm_params: Mapping[str, object],
timeout: float | httpx.Timeout | None = None,
) -> VectorStoreSearchResponse:
params: Final = _ValkeySearchParams.model_validate(litellm_params)
query_text: Final = self._query_text(query)
knn: Final = self._knn_query(
vector_store_search_optional_params,
embedding_field=params.embedding_field,
text_field=params.text_field,
)
embedding_response: Final = await self.aembedding_fn(
model=params.require_embedding_model(),
input=[query_text], # mutable-ok: litellm.embedding's input contract is a list
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
)
vec_params: Final = {"vec": pack_vector(embedding_response.data[0]["embedding"])} # mutable-ok: redis-py API
if self.async_client is not None:
raw: Final = await self.async_client.ft(vector_store_id).search( # pyright: ignore[reportGeneralTypeIssues] # types-redis 4.6 stubs shadow redis 5.3.1 and type the async client's ft() as the sync Search, so search() returns a non-awaitable Result; it is a coroutine at runtime
knn, query_params=vec_params
)
return self._to_response(raw, query_text, params.text_field)
connect_timeout, op_timeout = self._socket_timeouts(timeout)
client: Final = _import_async_redis().from_url(
params.connection_url(),
socket_connect_timeout=connect_timeout,
socket_timeout=op_timeout,
)
try:
raw_result: Final = await client.ft(vector_store_id).search( # pyright: ignore[reportGeneralTypeIssues] # types-redis 4.6 stubs shadow redis 5.3.1 and type the async client's ft() as the sync Search, so search() returns a non-awaitable Result; it is a coroutine at runtime
knn, query_params=vec_params
)
return self._to_response(raw_result, query_text, params.text_field)
finally:
await client.aclose()
def transform_create_vector_store_request(
self,
vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams,
api_base: str,
) -> NoReturn:
raise NotImplementedError(_SEARCH_ONLY_MESSAGE)
def transform_create_vector_store_response(self, response: httpx.Response) -> NoReturn:
raise NotImplementedError(_SEARCH_ONLY_MESSAGE)

View file

@ -420,6 +420,8 @@ async def acompletion(
verbosity: Literal["low", "medium", "high"] | None = None,
safety_identifier: str | None = None,
service_tier: str | None = None,
store: bool | None = None,
prompt_cache_key: str | None = None,
# set api_base, api_version, api_key
base_url: str | None = None,
api_version: str | None = None,
@ -585,6 +587,8 @@ async def acompletion(
"verbosity": verbosity,
"safety_identifier": safety_identifier,
"service_tier": service_tier,
"store": store,
"prompt_cache_key": prompt_cache_key,
"extra_headers": extra_headers,
"acompletion": True, # assuming this is a required parameter
"thinking": thinking,
@ -4930,6 +4934,8 @@ def completion(
extra_headers: dict | None = None,
safety_identifier: str | None = None,
service_tier: str | None = None,
store: bool | None = None,
prompt_cache_key: str | None = None,
# soon to be deprecated params by OpenAI
functions: list | None = None,
function_call: str | None = None,
@ -5058,6 +5064,8 @@ def completion(
verbosity=verbosity,
safety_identifier=safety_identifier,
service_tier=service_tier,
store=store,
prompt_cache_key=prompt_cache_key,
base_url=base_url,
api_version=api_version,
api_key=api_key,
@ -5367,6 +5375,8 @@ def completion(
),
"safety_identifier": safety_identifier,
"service_tier": service_tier,
"store": store,
"prompt_cache_key": prompt_cache_key,
"allowed_openai_params": kwargs.get("allowed_openai_params"),
"base_model": base_model,
}

View file

@ -10052,6 +10052,21 @@
"output_cost_per_second": 0.0066027,
"supports_tool_choice": true
},
"bedrock/guardrails": {
"guardrail_cost_per_unit": {
"automatedReasoningPolicyUnits": 0.00017,
"contentPolicyImageUnits": 0.00075,
"contentPolicyUnits": 0.00015,
"contextualGroundingPolicyUnits": 0.0001,
"sensitiveInformationPolicyFreeUnits": 0.0,
"sensitiveInformationPolicyUnits": 0.0001,
"topicPolicyUnits": 0.00015,
"wordPolicyUnits": 0.0
},
"litellm_provider": "bedrock",
"mode": "guardrail",
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"bedrock/ap-northeast-1/1-month-commitment/anthropic.claude-instant-v1": {
"input_cost_per_second": 0.01475,
"litellm_provider": "bedrock",

View file

@ -0,0 +1,6 @@
<?xml version="1.0" encoding="UTF-8"?>
<svg width="64" height="73" viewBox="0 0 64 73" xmlns="http://www.w3.org/2000/svg">
<g id="Group-copy">
<path id="Path" fill="#123678" fill-rule="evenodd" stroke="none" d="M 13.482285 60.694962 L 0.998384 52.884399 L 0.998384 19.502914 L 31.527868 2.001205 L 61.317604 19.532024 L 61.317604 54.64489 L 31.054855 71.68927 L 20.548372 65.115807 L 20.548372 51.041328 L 20.548372 49.119896 L 14.851504 45.555508 L 14.851504 27.453159 L 31.346497 17.99712 L 47.464485 27.482262 L 47.464485 46.451157 L 34.703495 53.638138 L 34.703495 45.998573 C 38.52874 44.52552 41.274452 40.739189 41.274452 36.270489 C 41.274452 30.510658 36.712814 25.88438 31.158138 25.88438 C 25.603172 25.88438 21.041817 30.510658 21.041817 36.270489 C 21.041817 40.739189 23.787249 44.52552 27.612494 45.998573 L 27.612494 60.473576 L 31.261133 62.756348 L 53.635483 50.15464 L 53.635483 23.924595 L 31.477489 10.884869 L 8.680504 23.953705 L 8.680504 48.628967 L 13.482285 51.633297 L 13.482285 60.694962 Z M 31.158138 31.498383 C 33.671822 31.498383 35.660439 33.664162 35.660439 36.270489 C 35.660439 38.876804 33.671822 41.042587 31.158138 41.042587 C 28.644447 41.042587 26.655558 38.876804 26.655558 36.270489 C 26.655558 33.664162 28.644447 31.498383 31.158138 31.498383 Z" />
</g>
</svg>

After

Width:  |  Height:  |  Size: 1.3 KiB

View file

@ -811,6 +811,9 @@ class LiteLLMRoutes(enum.Enum):
"/model/delete",
"/user/daily/activity",
"/user/daily/activity/aggregated",
# Endpoint restricts results to organizations the caller is ORG_ADMIN
# of; a caller who administers none gets an empty result set.
"/organization/daily/activity",
"/user/available_roles", # read-only role metadata; any authenticated user may read
"/user/list", # org admins checked in endpoint; non-admins get 403
"/model/{model_id}/update",
@ -1995,6 +1998,18 @@ class AddTeamCallback(LiteLLMPydanticObjectBase):
return values
class TeamCallbackDeleteResponseData(LiteLLMPydanticObjectBase):
team_id: str
success_callbacks: tuple[str, ...]
failure_callbacks: tuple[str, ...]
class TeamCallbackDeleteResponse(LiteLLMPydanticObjectBase):
status: Literal["success"]
message: str
data: TeamCallbackDeleteResponseData
class TeamCallbackMetadata(LiteLLMPydanticObjectBase):
success_callback: list[str] | None = []
failure_callback: list[str] | None = []

View file

@ -1,14 +1,15 @@
import asyncio
import contextlib
import json
import logging
import math
import time
import traceback
from collections.abc import AsyncGenerator, Callable, Mapping
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping
from datetime import datetime
from functools import lru_cache
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Protocol, overload
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Protocol, TypeAlias, TypeVar, overload
import anyio
import httpx
@ -31,11 +32,13 @@ from litellm.constants import (
UNSAFE_PROXY_RESPONSE_HEADERS,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
from litellm.litellm_core_utils.dd_tracing import NullTracer, tracer
from litellm.litellm_core_utils.get_supported_openai_params import (
get_supported_openai_params,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost
from litellm.litellm_core_utils.llm_response_utils.get_headers import (
get_response_headers,
)
@ -47,7 +50,12 @@ from litellm.proxy.common_utils.callback_utils import (
get_logging_caching_headers,
get_remaining_tokens_and_requests_from_request_data,
)
from litellm.proxy.common_utils.sse_keepalive import wrap_sse_stream_with_keepalive_pings
from litellm.proxy.common_utils.sse_keepalive import (
SSE_COMMENT_PING_BYTES,
coerce_keepalive_interval,
resolve_ttft_keepalive_interval,
wrap_sse_stream_with_keepalive_pings,
)
from litellm.proxy.dd_span_tagger import DDSpanTagger
from litellm.proxy.route_llm_request import route_request
from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails
@ -56,6 +64,100 @@ from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_di
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.router import RouterRateLimitError
_LateResponseT = TypeVar("_LateResponseT", bound=Response)
_LlmCallT = TypeVar("_LlmCallT")
ProxyRouteType: TypeAlias = Literal[
"acompletion",
"aembedding",
"aresponses",
"_arealtime",
"_aresponses_websocket",
"acreate_realtime_client_secret",
"arealtime_calls",
"aget_responses",
"adelete_responses",
"acancel_responses",
"acompact_responses",
"acreate_batch",
"aretrieve_batch",
"alist_batches",
"acancel_batch",
"afile_content",
"afile_retrieve",
"afile_delete",
"atext_completion",
"acreate_fine_tuning_job",
"acancel_fine_tuning_job",
"alist_fine_tuning_jobs",
"aretrieve_fine_tuning_job",
"alist_input_items",
"aimage_edit",
"agenerate_content",
"agenerate_content_stream",
"allm_passthrough_route",
"avector_store_search",
"avector_store_create",
"avector_store_retrieve",
"avector_store_list",
"avector_store_update",
"avector_store_delete",
"avector_store_file_create",
"avector_store_file_list",
"avector_store_file_retrieve",
"avector_store_file_content",
"avector_store_file_update",
"avector_store_file_delete",
"aocr",
"asearch",
"avideo_generation",
"avideo_list",
"avideo_status",
"avideo_content",
"avideo_remix",
"avideo_create_character",
"avideo_get_character",
"avideo_edit",
"avideo_extension",
"acreate_container",
"alist_containers",
"aingest",
"aretrieve_container",
"adelete_container",
"aupload_container_file",
"alist_container_files",
"aretrieve_container_file",
"adelete_container_file",
"aretrieve_container_file_content",
"acreate_skill",
"alist_skills",
"aget_skill",
"adelete_skill",
"anthropic_messages",
"acreate_interaction",
"aget_interaction",
"adelete_interaction",
"acancel_interaction",
"acreate_agent",
"alist_agents",
"aget_agent",
"adelete_agent",
"alist_agent_versions",
"asend_message",
"call_mcp_tool",
"acreate_eval",
"alist_evals",
"aget_eval",
"aupdate_eval",
"adelete_eval",
"acancel_eval",
"acreate_run",
"alist_runs",
"aget_run",
"acancel_run",
"adelete_run",
]
from litellm.types.utils import ServerToolUse
# Type alias for streaming chunk serializer (chunk after hooks + cost injection -> wire format)
@ -559,6 +661,11 @@ class _UpstreamClosingStreamingResponse(StreamingResponse):
super().__init__(content, status_code=status_code, headers=headers, media_type=media_type)
self._upstream_generator = upstream_generator
@property
def upstream_generator(self) -> AsyncGenerator[str, None] | None:
"""The upstream LLM stream, for a caller that has to run this response's cleanup itself."""
return self._upstream_generator
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
try:
await super().__call__(scope, receive, send)
@ -649,6 +756,39 @@ async def _buffer_first_chunk_honoring_disconnect(
raise _ClientDisconnectedBeforeFirstChunk()
def _sse_error_payload(exc: BaseException) -> tuple[int, Mapping[str, object]]:
"""Build the ProxyException-shaped ``{"error": ...}`` body used in SSE error frames.
Matches ``ProxyException.to_dict()`` so streaming and non-streaming error frames
are byte-identical.
"""
# Preserve status code from HTTPException (e.g. guardrail blocks)
error_status: Final = getattr(exc, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR)
raw_detail: Final = _getattr_object(exc, "detail", "Error processing stream start")
message, structured_fields = _serialize_http_exception_detail(raw_detail)
existing_fields: Final = getattr(exc, "provider_specific_fields", None) or {}
merged_fields: Final = {**existing_fields, **structured_fields} if structured_fields else (existing_fields or None)
# Built in one statement then given its one optional key, rather than spread
# conditionally: the spread form costs two extra dict constructions, which
# type-discipline-budget.json's LIT002 ceiling has no room for.
error_obj: Final = {
"message": message,
"type": getattr(exc, "type", "None"),
"param": getattr(exc, "param", "None"),
"code": str(error_status),
}
if merged_fields:
error_obj["provider_specific_fields"] = merged_fields
return error_status, error_obj
def _sse_error_frames(error_obj: Mapping[str, object]) -> tuple[str, str]:
"""The two frames an SSE stream ends with once it can no longer raise."""
return f"data: {json.dumps({'error': error_obj})}\n\n", "data: [DONE]\n\n"
async def create_response(
generator: AsyncGenerator[str, None],
media_type: str,
@ -740,31 +880,11 @@ async def create_response(
# Unexpected error consuming first chunk.
verbose_proxy_logger.exception("Error consuming first chunk from generator: %s", e)
# Preserve status code from HTTPException (e.g., guardrail blocks)
error_status: Final = getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR)
raw_detail: Final = _getattr_object(e, "detail", "Error processing stream start")
message, structured_fields = _serialize_http_exception_detail(raw_detail)
existing_fields: Final = getattr(e, "provider_specific_fields", None) or {}
if structured_fields:
merged_fields: dict | None = {**existing_fields, **structured_fields}
else:
merged_fields = existing_fields or None
# Match ProxyException.to_dict() shape so streaming and non-streaming
# error frames are byte-identical.
error_obj: Final[dict[str, object]] = {
"message": message,
"type": getattr(e, "type", "None"),
"param": getattr(e, "param", "None"),
"code": str(error_status),
}
if merged_fields:
error_obj["provider_specific_fields"] = merged_fields
error_status, error_obj = _sse_error_payload(e)
async def error_gen_message() -> AsyncGenerator[str, None]:
yield f"data: {json.dumps({'error': error_obj})}\n\n"
yield "data: [DONE]\n\n"
for frame in _sse_error_frames(error_obj):
yield frame
return StreamingResponse(
error_gen_message(),
@ -797,6 +917,176 @@ async def create_response(
)
_TTFT_KEEPALIVE_HEADERS: Final[Mapping[str, str]] = MappingProxyType(
{
"Cache-Control": "no-cache",
"X-Accel-Buffering": "no",
}
)
def ttft_keepalive_interval(request_data: Mapping[str, object], llm_router: Router | None = None) -> float | None:
"""The operator's keepalive interval, but only for a request that asked to stream.
Resolved through the deployments the request could land on, so a deployment's
`keepalive_seconds: 0` stays the hard disable it is documented to be rather
than being switched back on by the global default.
"""
if request_data.get("stream") is not True:
return None
requested_model: Final = request_data.get("model")
deployments: Final = (
llm_router.get_model_list(model_name=requested_model) or ()
if llm_router is not None and isinstance(requested_model, str)
else ()
)
return resolve_ttft_keepalive_interval(deployments, litellm.sse_keepalive_ping_interval_seconds)
async def _aclose_late_response(produced: Response) -> None:
"""Run the cleanup Starlette would have run, for a response it never called.
Closing an already-closed async generator is a no-op, so this is safe to call
from both the relay's own teardown and the outer one.
"""
if not isinstance(produced, StreamingResponse):
return
targets: Final = (
(produced.body_iterator, produced.upstream_generator)
if isinstance(produced, _UpstreamClosingStreamingResponse)
else (produced.body_iterator,)
)
for target in targets:
aclose = getattr(target, "aclose", None)
if aclose is None:
continue
try:
await aclose()
except BaseException as exc: # noqa: BLE001 # teardown must not mask why the stream ended
verbose_proxy_logger.debug("error closing relayed streaming generator: %s", exc)
async def _relay_late_response(produced: Response) -> AsyncGenerator[bytes, None]:
"""Replay a Response that was built after a keepalive had already opened the wire."""
if not isinstance(produced, StreamingResponse):
# The status line is already on the wire, so a non-streaming body, an error
# body included, can only reach the client as an SSE frame.
yield b"data: " + (bytes(produced.body) or b"{}") + b"\n\n"
yield b"data: [DONE]\n\n"
return
try:
async for chunk in produced.body_iterator:
yield chunk.encode("utf-8") if isinstance(chunk, str) else bytes(chunk)
finally:
# Starlette never called this response, so the cleanup its __call__ would
# have run has to happen here or the upstream LLM connection leaks.
with anyio.CancelScope(shield=True):
await _aclose_late_response(produced)
async def _sanitized_late_failure(
exc: Exception,
on_late_failure: "Callable[[Exception], Awaitable[HTTPException | None]] | None",
) -> Exception:
"""Report a late failure and return whatever should reach the client.
``post_call_failure_hook`` lets a callback replace the client-facing error, by
returning a replacement or by raising one, and both are used elsewhere in this
module. Serializing the original would leak provider detail a deployment had
configured away, so the hook's answer wins. A callback that fails some other
way is a bug in the callback, not a reason to lose the real error.
"""
if on_late_failure is None:
return exc
try:
replacement: Final = await on_late_failure(exc)
except HTTPException as raised_replacement:
return raised_replacement
except Exception as hook_failure: # noqa: BLE001 # a broken callback must not replace the real error
verbose_proxy_logger.exception("post_call_failure_hook raised while reporting a late failure: %s", hook_failure)
return exc
return replacement if replacement is not None else exc
async def open_sse_before_first_byte(
produce_response: Awaitable[_LateResponseT],
ping_interval_seconds: float | str | None,
media_type: str = "text/event-stream",
on_late_failure: Callable[[Exception], Awaitable[HTTPException | None]] | None = None,
) -> _LateResponseT | StreamingResponse:
"""Write SSE keepalive comments while the upstream LLM call is still in flight.
The whole time-to-first-token is spent inside `produce_response`: the upstream
withholds its response headers until it emits its first token, so nothing has
entered the ASGI response phase yet and the proxy writes zero bytes. An
intermediary with an idle read timeout (AWS ALB and nginx both default to 60s)
then drops a connection that is perfectly healthy.
When `produce_response` does not finish within one interval, the response is
opened immediately and `: ping` comments, which every conformant SSE client
ignores, fill the wire until the real response is ready to be replayed onto it.
Committing the status line that early is the cost: a failure discovered after
the first ping reaches the client as an SSE error frame under a 200 rather than
as an HTTP error status, and LiteLLM's own `x-litellm-*` response headers are
not yet known. Both are why this stays off until an operator sets an interval.
"""
interval: Final = coerce_keepalive_interval(ping_interval_seconds)
if interval is None:
return await produce_response
produce_task: Final = asyncio.ensure_future(produce_response)
await asyncio.wait((produce_task,), timeout=interval)
if produce_task.done():
# Fast path: the upstream answered inside one interval, so nothing was
# written early and this is byte-identical to not being wrapped at all.
return produce_task.result()
async def keepalive_then_relay() -> AsyncGenerator[bytes, None]:
try:
while not produce_task.done():
yield SSE_COMMENT_PING_BYTES
await asyncio.wait((produce_task,), timeout=interval)
try:
produced: Final = produce_task.result()
except Exception as exc: # noqa: BLE001 # the status line is already sent; surface it as a frame
verbose_proxy_logger.exception(
"request failed after its SSE keepalive had opened the response: %s", exc
)
# The caller's own `except` never sees this, so its failure hook
# would never fire and the failure would go unaudited. The hook
# also gets to sanitize what reaches the client, by returning or
# raising a replacement, so its answer decides the frame.
_, error_obj = _sse_error_payload(await _sanitized_late_failure(exc, on_late_failure))
for frame in _sse_error_frames(error_obj):
yield frame.encode()
return
async for chunk in _relay_late_response(produced):
yield chunk
finally:
if not produce_task.done():
produce_task.cancel()
with anyio.CancelScope(shield=True):
with contextlib.suppress(BaseException):
await produce_task
elif not produce_task.cancelled():
# The upstream may have answered while nobody was draining this
# relay, e.g. the client vanished first. Nothing else holds that
# response, so its stream only gets closed here.
with anyio.CancelScope(shield=True):
with contextlib.suppress(BaseException):
await _aclose_late_response(produce_task.result())
verbose_proxy_logger.info(
"no upstream response after %ss, opening the SSE response early and sending keepalives", interval
)
return StreamingResponse(
keepalive_then_relay(),
media_type=media_type,
headers=_TTFT_KEEPALIVE_HEADERS,
)
def _is_azure_model_router_request(model: str) -> bool:
"""
Check if the requested model is an Azure Model Router.
@ -1043,7 +1333,7 @@ def _log_llm_api_exception(e: Exception) -> None:
async def _cancel_llm_call_on_client_disconnect(
request: Request,
llm_api_call: "asyncio.Future[object]",
llm_api_call: "asyncio.Future[_LlmCallT]",
disconnect_event: asyncio.Event,
) -> None:
try:
@ -1062,8 +1352,8 @@ async def _cancel_llm_call_on_client_disconnect(
async def _await_llm_call_cancelling_on_disconnect(
request: Request,
llm_api_call: "asyncio.Future[Any]",
) -> Any:
llm_api_call: "asyncio.Future[_LlmCallT]",
) -> _LlmCallT:
disconnect_event: Final = asyncio.Event()
monitor: Final = asyncio.create_task(_cancel_llm_call_on_client_disconnect(request, llm_api_call, disconnect_event))
try:
@ -1714,100 +2004,11 @@ class ProxyBaseLLMRequestProcessing:
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth,
route_type: Literal[
"acompletion",
"aembedding",
"aresponses",
"_arealtime",
"_aresponses_websocket",
"acreate_realtime_client_secret",
"arealtime_calls",
"aget_responses",
"adelete_responses",
"acancel_responses",
"acompact_responses",
"acreate_batch",
"aretrieve_batch",
"alist_batches",
"acancel_batch",
"afile_content",
"afile_retrieve",
"afile_delete",
"atext_completion",
"acreate_fine_tuning_job",
"acancel_fine_tuning_job",
"alist_fine_tuning_jobs",
"aretrieve_fine_tuning_job",
"alist_input_items",
"aimage_edit",
"agenerate_content",
"agenerate_content_stream",
"allm_passthrough_route",
"avector_store_search",
"avector_store_create",
"avector_store_retrieve",
"avector_store_list",
"avector_store_update",
"avector_store_delete",
"avector_store_file_create",
"avector_store_file_list",
"avector_store_file_retrieve",
"avector_store_file_content",
"avector_store_file_update",
"avector_store_file_delete",
"aocr",
"asearch",
"avideo_generation",
"avideo_list",
"avideo_status",
"avideo_content",
"avideo_remix",
"avideo_create_character",
"avideo_get_character",
"avideo_edit",
"avideo_extension",
"acreate_container",
"alist_containers",
"aingest",
"aretrieve_container",
"adelete_container",
"aupload_container_file",
"alist_container_files",
"aretrieve_container_file",
"adelete_container_file",
"aretrieve_container_file_content",
"acreate_skill",
"alist_skills",
"aget_skill",
"adelete_skill",
"anthropic_messages",
"acreate_interaction",
"aget_interaction",
"adelete_interaction",
"acancel_interaction",
"acreate_agent",
"alist_agents",
"aget_agent",
"adelete_agent",
"alist_agent_versions",
"asend_message",
"call_mcp_tool",
"acreate_eval",
"alist_evals",
"aget_eval",
"aupdate_eval",
"adelete_eval",
"acancel_eval",
"acreate_run",
"alist_runs",
"aget_run",
"acancel_run",
"adelete_run",
],
route_type: ProxyRouteType,
proxy_logging_obj: ProxyLogging,
general_settings: dict,
general_settings: dict[str, object],
proxy_config: ProxyConfig,
select_data_generator: Callable | None = None,
select_data_generator: Callable[..., object] | None = None,
llm_router: Router | None = None,
model: str | None = None,
user_model: str | None = None,
@ -1817,7 +2018,72 @@ class ProxyBaseLLMRequestProcessing:
user_api_base: str | None = None,
version: str | None = None,
is_streaming_request: bool | None = False,
contents: list | None = None, # Add contents parameter
contents: list[object] | None = None,
skip_pre_call_logic: bool = False,
) -> Any:
"""Run the request, sending SSE keepalives while the upstream is still silent.
Everything below this point, the upstream call included, happens before the
proxy can write a byte, so a slow time-to-first-token leaves the response
idle. See ``open_sse_before_first_byte``; unwrapped unless an operator sets
``litellm_settings.sse_keepalive_ping_interval_seconds``.
"""
async def _audit_late_failure(exc: Exception) -> HTTPException | None:
# Once a keepalive is on the wire this can no longer raise, so the
# caller's `except` never runs its own post_call_failure_hook.
return await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=exc,
request_data=self.data,
)
return await open_sse_before_first_byte(
self._process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type=route_type,
proxy_logging_obj=proxy_logging_obj,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
llm_router=llm_router,
model=model,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
is_streaming_request=is_streaming_request,
contents=contents,
skip_pre_call_logic=skip_pre_call_logic,
),
ping_interval_seconds=ttft_keepalive_interval(self.data, llm_router),
on_late_failure=_audit_late_failure,
)
async def _process_llm_request(
self,
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth,
route_type: ProxyRouteType,
proxy_logging_obj: ProxyLogging,
general_settings: dict[str, object],
proxy_config: ProxyConfig,
select_data_generator: Callable[..., object] | None = None,
llm_router: Router | None = None,
model: str | None = None,
user_model: str | None = None,
user_temperature: float | None = None,
user_request_timeout: float | None = None,
user_max_tokens: int | None = None,
user_api_base: str | None = None,
version: str | None = None,
is_streaming_request: bool | None = False,
contents: list[object] | None = None, # Add contents parameter
skip_pre_call_logic: bool = False,
) -> Any:
"""
@ -2039,6 +2305,7 @@ class ProxyBaseLLMRequestProcessing:
return StreamingResponse(
content=generator,
status_code=status.HTTP_200_OK,
media_type=self._passthrough_event_stream_media_type(),
headers=custom_headers,
)
else:
@ -2197,11 +2464,21 @@ class ProxyBaseLLMRequestProcessing:
additional_headers = hidden_params.get("additional_headers", {}) or {}
recover_response_cost: Final = not response_cost and hidden_params.get("response_cost") is None
response_cost_for_headers: Final = (
llm_cost_for_headers: Final = (
self._response_cost_from_logging_obj(response=response, logging_obj=logging_obj) or ""
if recover_response_cost
else response_cost
)
_, request_metadata_bucket = get_or_create_metadata_bucket(self.data)
guardrail_cost_for_headers: Final = guardrail_information_cost(
request_metadata_bucket.get("standard_logging_guardrail_information")
)
response_cost_for_headers: Final = (
(llm_cost_for_headers if isinstance(llm_cost_for_headers, (int, float)) else 0.0)
+ guardrail_cost_for_headers
if guardrail_cost_for_headers > 0
else llm_cost_for_headers
)
fastapi_response.headers.update(
ProxyBaseLLMRequestProcessing.get_custom_headers(
@ -2494,10 +2771,16 @@ class ProxyBaseLLMRequestProcessing:
def _passthrough_event_stream_media_type(self) -> str | None:
"""
Content-type for a buffered passthrough event-stream response, resolved
from the provider handler so the proxy stays provider-agnostic. Mirrors
the upstream content-type the non-streaming path forwards, since the
buffered streaming generator carries no headers of its own.
Content-type for a passthrough event-stream response, resolved from the
provider handler so the proxy stays provider-agnostic. Mirrors the
upstream content-type the non-streaming path forwards, since the
streaming generator carries no headers of its own. Used for both the
buffered (guardrail-rewritten) and the unbuffered relay paths so
clients that enforce the event-stream content-type (e.g. Claude Code on
Bedrock invoke-with-response-stream) see the correct header instead of
no content-type at all, which they fall back to reading as
application/octet-stream. Returns None for providers with no
event-stream media type, leaving the response headers unchanged.
"""
from litellm.llms.pass_through.guardrail_translation.handler import (
LlmPassthroughRouteHandler,

View file

@ -1,15 +1,22 @@
import asyncio
import contextlib
import math
from collections.abc import AsyncGenerator
from collections.abc import AsyncGenerator, Iterable, Mapping
from typing import Final
import anyio
ANTHROPIC_PING_SSE_CHUNK: Final = 'event: ping\ndata: {"type": "ping"}\n\n'
SSE_COMMENT_PING_BYTES: Final = b": ping\n\n"
# The byte form of proxy_server._SSE_FRAME_DELIMITERS, CR-only included: SSE
# terminates a line with CRLF, LF or CR, so a blank line is any of these three.
_SSE_FRAME_DELIMITERS: Final = (b"\r\n\r\n", b"\n\n", b"\r\r")
_SSE_DELIMITER_LOOKBACK: Final = max(len(delimiter) for delimiter in _SSE_FRAME_DELIMITERS)
_STREAM_START_TAIL: Final = b"\n\n"
_SSE_MEDIA_TYPE: Final = "text/event-stream"
def _coerce_interval(ping_interval_seconds: float | str | None) -> float | None:
def coerce_keepalive_interval(ping_interval_seconds: float | str | None) -> float | None:
if ping_interval_seconds is None:
return None
try:
@ -28,7 +35,7 @@ def keepalive_ping_has_fired(elapsed_seconds: float, ping_interval_seconds: floa
the status line is already on the wire. With pings disabled nothing flushes early, so a raise
still carries its real status.
"""
interval: Final = _coerce_interval(ping_interval_seconds)
interval: Final = coerce_keepalive_interval(ping_interval_seconds)
return interval is not None and elapsed_seconds >= interval
@ -36,7 +43,7 @@ def wrap_sse_stream_with_keepalive_pings(
stream: AsyncGenerator[str, None],
ping_interval_seconds: float | str | None,
) -> AsyncGenerator[str, None]:
interval: Final = _coerce_interval(ping_interval_seconds)
interval: Final = coerce_keepalive_interval(ping_interval_seconds)
if interval is None:
return stream
return _keepalive_ping_stream(stream=stream, ping_interval_seconds=interval)
@ -66,3 +73,96 @@ async def _keepalive_ping_stream(
with contextlib.suppress(BaseException):
await pending
await stream.aclose()
def is_sse_content_type(content_type: str | None) -> bool:
return content_type is not None and content_type.split(";", 1)[0].strip().lower() == _SSE_MEDIA_TYPE
def wrap_passthrough_sse_bytes_with_keepalive_pings(
stream: AsyncGenerator[bytes, None],
ping_interval_seconds: float | str | None,
upstream_headers: Mapping[str, str],
) -> AsyncGenerator[bytes, None]:
"""Fill upstream silence on a byte-relaying passthrough stream with SSE comments.
Passthrough routes relay upstream bytes verbatim, so a model that thinks for
longer than an intermediary's idle read timeout has its connection dropped
before the first token. Only streams the upstream itself declares as
``text/event-stream`` are wrapped: a comment spliced into a binary transport
(AWS event streams on ``/bedrock``, protobuf, NDJSON) would corrupt it.
"""
interval: Final = coerce_keepalive_interval(ping_interval_seconds)
if interval is None or not is_sse_content_type(upstream_headers.get("content-type")):
return stream
return _keepalive_ping_byte_stream(stream=stream, ping_interval_seconds=interval)
async def _keepalive_ping_byte_stream(
stream: AsyncGenerator[bytes, None],
ping_interval_seconds: float,
) -> AsyncGenerator[bytes, None]:
pending = asyncio.ensure_future(
stream.__anext__()
) # rebind-ok: re-armed with the next __anext__ after each delivered chunk
# The tail of the bytes relayed so far, long enough to hold any delimiter.
# Seeded as a delimiter because a stream starts at a frame boundary, and kept
# across chunks because a delimiter can be split between two transport reads,
# which testing only the latest chunk would miss for the rest of the stream.
recent_tail = _STREAM_START_TAIL # rebind-ok: rolling window over the relayed bytes
try:
while True:
await asyncio.wait((pending,), timeout=ping_interval_seconds)
if not pending.done():
# The relayed chunks are raw transport reads, not whole SSE
# frames, so an upstream that stalls halfway through a frame
# must not have a comment spliced into it.
if recent_tail.endswith(_SSE_FRAME_DELIMITERS):
yield SSE_COMMENT_PING_BYTES
continue
try:
chunk: bytes = pending.result()
except StopAsyncIteration:
return
if chunk:
recent_tail = (recent_tail + chunk)[-_SSE_DELIMITER_LOOKBACK:]
yield chunk
pending = asyncio.ensure_future(stream.__anext__())
finally:
pending.cancel()
with anyio.CancelScope(shield=True):
with contextlib.suppress(BaseException):
await pending
await stream.aclose()
def resolve_ttft_keepalive_interval(
deployments: Iterable[Mapping[str, object]],
global_interval: float | str | None,
) -> float | None:
"""The keepalive interval to use before the upstream has answered at all.
No deployment has served the request yet, so a per-deployment
``keepalive_seconds`` is only trusted when every candidate under the requested
model carries the same one, which is how the mid-stream engine treats its own
model_name fallback. Otherwise the operator's global default applies.
An explicit ``0`` survives as a disable, since coercion rejects it: that keeps
an operator's documented hard disable working on this path too, rather than
letting the global switch a deployment back on behind their back.
A client-supplied value is deliberately not consulted. Opening the response
early is an operator decision, and a request must not be able to enable it for
a deployment that never did.
"""
configured: Final = frozenset(_keepalive_param(deployment) for deployment in deployments)
agreed: Final = next(iter(configured)) if len(configured) == 1 else None
return coerce_keepalive_interval(global_interval if agreed is None else agreed)
def _keepalive_param(deployment: Mapping[str, object]) -> float | str | None:
params: Final = deployment.get("litellm_params")
if not isinstance(params, Mapping):
return None
value: Final = params.get("keepalive_seconds")
return value if isinstance(value, (int, float, str)) else None

View file

@ -31,6 +31,7 @@ from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS
from litellm.exceptions import ModifyResponseException
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import bedrock_guardrail_cost
from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicMessagesHandler
from litellm.llms.base_llm.guardrail_translation.utils import (
effective_scan_only_tool_results_for_guardrail,
@ -872,6 +873,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
credentials, aws_region_name = self._load_credentials()
allow_chunking: Final = not self._content_uses_contextual_grounding(content)
completed_chunk_usages: Final[list[BedrockGuardrailUsage]] = [] # mutable-ok: billed-chunk usage accumulator
try:
responses: Final = await self._apply_guardrail_content_with_chunking(
content=content,
@ -883,6 +885,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
event_type=event_type,
start_time=start_time,
allow_chunking=allow_chunking,
completed_chunk_usages=completed_chunk_usages,
)
except HTTPException as exc:
if not isinstance(exc.detail, dict):
@ -891,6 +894,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data=request_data,
event_type=event_type,
start_time=start_time,
aws_region_name=aws_region_name,
completed_chunk_usages=completed_chunk_usages,
)
raise
merged_response: Final = self._merge_bedrock_guardrail_responses(responses)
@ -899,6 +904,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data=request_data,
event_type=event_type,
start_time=start_time,
aws_region_name=aws_region_name,
)
return merged_response
@ -913,6 +919,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
event_type: GuardrailEventHooks,
start_time: "datetime",
allow_chunking: bool,
completed_chunk_usages: list[BedrockGuardrailUsage], # mutable-ok: billed-chunk usage accumulator
) -> tuple[BedrockContentChunkResult, ...]:
"""Post `content` to ApplyGuardrail, chunking only if AWS rejects it as too large.
@ -959,6 +966,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data=request_data,
event_type=event_type,
start_time=start_time,
completed_chunk_usages=completed_chunk_usages,
)
return (
BedrockContentChunkResult(
@ -989,6 +997,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
event_type=event_type,
start_time=start_time,
allow_chunking=allow_chunking,
completed_chunk_usages=completed_chunk_usages,
)
for batch in batches
]
@ -1015,6 +1024,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
event_type=event_type,
start_time=start_time,
allow_chunking=allow_chunking,
completed_chunk_usages=completed_chunk_usages,
)
second_results: Final = await self._apply_guardrail_content_with_chunking(
content=second_half,
@ -1026,6 +1036,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
event_type=event_type,
start_time=start_time,
allow_chunking=allow_chunking,
completed_chunk_usages=completed_chunk_usages,
)
combined_results: Final = tuple(first_results) + tuple(second_results)
if is_single_item_text_split:
@ -1045,6 +1056,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper
event_type: GuardrailEventHooks,
start_time: "datetime",
completed_chunk_usages: list[BedrockGuardrailUsage], # mutable-ok: passed through to the single-call layer
) -> BedrockGuardrailResponse:
"""Post one ApplyGuardrail call for `content`, retrying with exponential
backoff on AWS ThrottlingException (HTTP 429).
@ -1072,6 +1084,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data=request_data,
event_type=event_type,
start_time=start_time,
completed_chunk_usages=completed_chunk_usages,
)
except HTTPException as exc:
if (
@ -1093,6 +1106,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper
event_type: GuardrailEventHooks,
start_time: "datetime",
completed_chunk_usages: list[BedrockGuardrailUsage], # mutable-ok: billed-chunk usage accumulator
) -> BedrockGuardrailResponse:
"""Make exactly one signed ApplyGuardrail HTTP call for `content` and
parse the result. Raises HTTPException on a guardrail block or any
@ -1108,7 +1122,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
A block is logged here rather than by the caller: it ends the whole chunking
flow immediately, with no further chunks attempted, so there is no later
merged response for the caller to log instead.
merged response for the caller to log instead. The logged usage still spans
the whole logical request: chunks that passed before the block appended what
AWS billed them to ``completed_chunk_usages``, and the attempt log sums those
with the blocking call's own usage.
"""
bedrock_request_data: Final = { # mutable-ok: outbound JSON request body
**base_request_data,
@ -1151,10 +1168,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data=request_data,
event_type=event_type,
start_time=start_time,
aws_region_name=aws_region_name,
completed_chunk_usages=completed_chunk_usages,
)
raise self._get_http_exception_for_blocked_guardrail(
bedrock_guardrail_response, request_data=request_data
)
response_usage: Final = bedrock_guardrail_response.get("usage")
if isinstance(response_usage, dict):
completed_chunk_usages.append(
response_usage
) # rebind-ok: accumulator threaded from make_bedrock_api_request, recording this billed call
return bedrock_guardrail_response
status_code, detail_message = self._parse_bedrock_guardrail_error_response(httpx_response)
@ -1172,14 +1196,31 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper
event_type: GuardrailEventHooks,
start_time: "datetime",
aws_region_name: str | None,
completed_chunk_usages: Sequence[BedrockGuardrailUsage],
) -> None:
"""Log a single ApplyGuardrail HTTP attempt as-is (its own status,
derived from its own response). Used only for the blocked-content
case, which ends the whole chunking flow immediately."""
tracing_detail: Final = self._build_tracing_detail(BedrockGuardrailResponse(**json_response))
"""Log the blocking ApplyGuardrail attempt, which ends the whole chunking
flow immediately. Its status derives from its own response, but its usage
(and so its cost) spans every billed call of the logical request: the
chunks that passed before the block plus the blocking call itself."""
blocking_usage: Final = json_response.get("usage")
billed_usages: Final[tuple[BedrockGuardrailUsage, ...]] = tuple(completed_chunk_usages) + (
(blocking_usage,) if isinstance(blocking_usage, dict) else ()
)
logged_json_response: Final = (
{ # mutable-ok: raw AWS JSON payload carrying the total billed usage
**json_response,
"usage": self._sum_usage_counters(billed_usages),
}
if completed_chunk_usages
else json_response
)
tracing_detail: Final = self._build_tracing_detail(
BedrockGuardrailResponse(**logged_json_response), aws_region_name=aws_region_name
)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response=json_response,
guardrail_json_response=logged_json_response,
request_data=request_data or {}, # mutable-ok: logging helper requires a dict
guardrail_status=self._get_bedrock_guardrail_response_status(response=httpx_response),
start_time=start_time.timestamp(),
@ -1195,6 +1236,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper
event_type: GuardrailEventHooks,
start_time: "datetime",
aws_region_name: str | None,
) -> None:
"""Log one logical ApplyGuardrail call -- possibly several chunk calls
under the hood -- using its final merged response, so a chunked
@ -1205,7 +1247,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
``Output.__type`` with an exception marker. That marker survives the merge,
so the status is derived from the merged response rather than assumed to be
a success, which is what the pre-chunking code reported for that shape."""
tracing_detail: Final = self._build_tracing_detail(merged_response)
tracing_detail: Final = self._build_tracing_detail(merged_response, aws_region_name=aws_region_name)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response=dict(merged_response), # mutable-ok: logging helper requires a dict
@ -1228,20 +1270,36 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper
event_type: GuardrailEventHooks,
start_time: "datetime",
aws_region_name: str | None,
completed_chunk_usages: Sequence[BedrockGuardrailUsage],
) -> None:
"""Log one logical ApplyGuardrail call that failed end-to-end (an
unrecoverable too-large error, a non-size validation error, or
exhausted throttle retries) as a single failure, rather than logging
every failed attempt chunking made along the way."""
every failed attempt chunking made along the way. Chunk calls AWS
billed before the failure still carry their usage and cost."""
billed_usage: Final = self._sum_usage_counters(completed_chunk_usages) if completed_chunk_usages else None
error_payload: Final = {"error": str(detail)} # mutable-ok: logging helper requires a dict
json_response: Final = (
{**error_payload, "usage": billed_usage} # mutable-ok: logging helper requires a dict
if billed_usage is not None
else error_payload
)
tracing_detail: Final = (
self._build_tracing_detail(BedrockGuardrailResponse(usage=billed_usage), aws_region_name=aws_region_name)
if billed_usage is not None
else None
)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response={"error": str(detail)}, # mutable-ok: logging helper requires a dict
guardrail_json_response=json_response,
request_data=request_data or {}, # mutable-ok: logging helper requires a dict
guardrail_status="guardrail_failed_to_respond",
start_time=start_time.timestamp(),
end_time=datetime.now(timezone.utc).timestamp(),
duration=(datetime.now(timezone.utc) - start_time).total_seconds(),
event_type=event_type,
tracing_detail=tracing_detail or None,
)
@staticmethod
@ -1504,15 +1562,20 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
Keys are taken from the responses rather than from a fixed list, so a counter
this code does not know about (AWS has added several) is still summed and
reported instead of being silently dropped to zero."""
chunk_usages: Final = tuple(
chunk_result.response.get("usage") or {} # mutable-ok: read-only empty fallback
for chunk_result in chunk_results
return BedrockGuardrail._sum_usage_counters(
tuple(
chunk_result.response.get("usage") or {} # mutable-ok: read-only empty fallback
for chunk_result in chunk_results
)
)
@staticmethod
def _sum_usage_counters(usages: Sequence[BedrockGuardrailUsage]) -> BedrockGuardrailUsage:
return cast( # cast-ok: TypedDict assembled from a comprehension
BedrockGuardrailUsage,
{ # mutable-ok: builds the TypedDict payload
key: sum(usage.get(key) or 0 for usage in chunk_usages)
for key in dict.fromkeys(key for usage in chunk_usages for key in usage)
key: sum(usage.get(key) or 0 for usage in usages)
for key in dict.fromkeys(key for usage in usages for key in usage)
},
)
@ -2036,7 +2099,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return (status_code, err)
return (status_code, message)
def _build_tracing_detail(self, response: BedrockGuardrailResponse) -> GuardrailTracingDetail:
def _build_tracing_detail(
self, response: BedrockGuardrailResponse, aws_region_name: str | None
) -> GuardrailTracingDetail:
"""
Build the tracing detail from the raw Bedrock response, before
redaction, so downstream loggers (OTEL, Langfuse, ...) get the
@ -2060,6 +2125,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
}
if usage_units:
tracing_detail["guardrail_usage"] = usage_units
tracing_detail["guardrail_cost"] = bedrock_guardrail_cost(
usage_units=usage_units, aws_region_name=aws_region_name
)
return tracing_detail
def _extract_violation_category_names(self, response: BedrockGuardrailResponse) -> list[str]:

View file

@ -5,7 +5,7 @@ GET /guardrails/usage/overview, /guardrails/usage/detail/:id, /guardrails/usage/
import json
from collections.abc import Callable, Iterable, Mapping, Sequence
from datetime import datetime, timedelta, timezone
from datetime import date, datetime, timedelta, timezone
from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, overload
@ -48,6 +48,40 @@ router: Final = APIRouter()
_EMPTY_UNITS: Final[Mapping[str, int]] = MappingProxyType({})
_USAGE_MAX_RANGE_DAYS: Final = 366
def _resolve_usage_window(start_date: str | None, end_date: str | None) -> tuple[str, str]:
from fastapi import HTTPException, status
now: Final = datetime.now(timezone.utc)
end: Final = end_date or now.strftime("%Y-%m-%d")
start: Final = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d")
try:
parsed: Final = (date.fromisoformat(start), date.fromisoformat(end))
except ValueError:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="start_date and end_date must be in YYYY-MM-DD format",
)
start_obj, end_obj = parsed
if (start_obj.isoformat(), end_obj.isoformat()) != (start, end):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="start_date and end_date must be in YYYY-MM-DD format",
)
if end_obj < start_obj:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="start_date must be on or before end_date",
)
if end_obj - start_obj > timedelta(days=_USAGE_MAX_RANGE_DAYS):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Date range too large; maximum is {_USAGE_MAX_RANGE_DAYS} days",
)
return start, end
def _guardrails_table(
prisma_client: "PrismaClient",
@ -457,9 +491,7 @@ async def guardrails_usage_overview(
rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0, totalUsageUnits=_EMPTY_UNITS
)
now: Final = datetime.now(timezone.utc)
end: Final = end_date or now.strftime("%Y-%m-%d")
start: Final = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d")
start, end = _resolve_usage_window(start_date, end_date)
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
@ -477,7 +509,7 @@ async def guardrails_usage_overview(
)
# Previous period for trend
start_prev: Final = (datetime.strptime(start, "%Y-%m-%d") - timedelta(days=7)).strftime("%Y-%m-%d")
start_prev: Final = (date.fromisoformat(start) - timedelta(days=7)).isoformat()
metrics_prev: Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics] = await _find_daily_guardrail_metrics(
prisma_client, where={"date": {"gte": start_prev, "lt": start}}
)
@ -531,9 +563,7 @@ async def guardrails_usage_detail(
raise HTTPException(status_code=500, detail="Prisma client not initialized")
now: Final = datetime.now(timezone.utc)
end: Final = end_date or now.strftime("%Y-%m-%d")
start: Final = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d")
start, end = _resolve_usage_window(start_date, end_date)
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
@ -556,11 +586,12 @@ async def guardrails_usage_detail(
"date": {"gte": start, "lte": end},
},
)
start_prev: Final = (date.fromisoformat(start) - timedelta(days=7)).isoformat()
metrics_prev: Final[Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics]] = await _find_daily_guardrail_metrics(
prisma_client,
where={
"guardrail_id": {"in": metric_ids},
"date": {"lt": start},
"date": {"gte": start_prev, "lt": start},
},
)
units_where: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsWhereInput] = {
@ -838,9 +869,7 @@ async def policies_usage_overview(
rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0, totalUsageUnits=_EMPTY_UNITS
)
now: Final = datetime.now(timezone.utc)
end: Final = end_date or now.strftime("%Y-%m-%d")
start: Final = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d")
start, end = _resolve_usage_window(start_date, end_date)
try:
policies: Final = await _policies_table(prisma_client).find_many()
@ -851,7 +880,7 @@ async def policies_usage_overview(
prisma_client,
where={
"date": {
"gte": (datetime.strptime(start, "%Y-%m-%d") - timedelta(days=7)).strftime("%Y-%m-%d"),
"gte": (date.fromisoformat(start) - timedelta(days=7)).isoformat(),
"lt": start,
}
},

View file

@ -11,6 +11,7 @@ from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
)
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import (
get_key_object,
@ -184,9 +185,14 @@ class _ProxyDBLogger(CustomLogger):
# recovered cost onto request_data (the usage rides along in
# ``combined_usage_object`` for the token columns), so attribute the
# real partial spend to this failure row instead of zero.
recovered_response_cost = 0.0
if isinstance(request_data.get("combined_usage_object"), litellm.Usage):
recovered_response_cost = max(float(request_data.get("response_cost") or 0.0), 0.0)
recovered_stream_cost: Final = (
max(float(request_data.get("response_cost") or 0.0), 0.0)
if isinstance(request_data.get("combined_usage_object"), litellm.Usage)
else 0.0
)
recovered_response_cost: Final = recovered_stream_cost + guardrail_information_cost(
existing_metadata.get("standard_logging_guardrail_information")
)
await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key_dict.api_key,

View file

@ -296,7 +296,10 @@ _ALLOW_CLIENT_MESSAGE_REDACTION_OPT_OUT_METADATA_KEY: Final = "allow_client_mess
_CLIENT_PRICING_CONTROL_FIELDS: Final = frozenset(CustomPricingLiteLLMParams.model_fields.keys())
# ``model_info`` carries the same pricing fields when read by
# ``use_custom_pricing_for_model``; strip from metadata for the same reason.
_CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info"})
# ``standard_logging_guardrail_information`` is proxy-written telemetry summed
# into response_cost and spend; a client seeding it forges (even negative)
# guardrail cost.
_CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info", "standard_logging_guardrail_information"})
_ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY: Final = "allow_client_pricing_override"
# Request fields whose value, when URL-valued, becomes the outbound destination

View file

@ -559,7 +559,7 @@ async def get_organization_daily_activity(
# Fetch organization aliases for metadata
where_condition: Final = _STR_OBJECT_DICT_ADAPTER.validate_python({})
if org_ids_list:
if org_ids_list is not None:
where_condition["organization_id"] = {"in": list(org_ids_list)}
org_aliases: Final = await _table(OrganizationRepository(prisma_client)).find_many(where=where_condition)

View file

@ -9,7 +9,7 @@ import copy
import json
import traceback
from datetime import datetime, timezone
from typing import Any, Final
from typing import Annotated, Any, Final
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
@ -23,6 +23,8 @@ from litellm.proxy._types import (
LitellmTableNames,
ProxyErrorTypes,
ProxyException,
TeamCallbackDeleteResponse,
TeamCallbackDeleteResponseData,
TeamCallbackMetadata,
UserAPIKeyAuth,
)
@ -209,6 +211,14 @@ async def _emit_team_callback_audit_log(
task.add_done_callback(_log_audit_task_exception)
def _callback_error(status_code: int, message: str) -> HTTPException:
"""Build the ``{"error": ...}`` failure body the team callback endpoints return."""
return HTTPException(
status_code=status_code,
detail={"error": message}, # mutable-ok: the error response body is a JSON object
)
@router.post(
"/team/{team_id:path}/callback",
tags=["team management"],
@ -363,6 +373,151 @@ async def add_team_callbacks(
)
@router.delete(
"/team/{team_id:path}/callback/{callback_name}",
tags=["team management"], # mutable-ok: FastAPI's route decorator takes a list of tags
dependencies=[Depends(user_api_key_auth)], # mutable-ok: FastAPI's route decorator takes a list of dependencies
response_model=TeamCallbackDeleteResponse,
)
@management_endpoint_wrapper
async def delete_team_callback(
http_request: Request,
team_id: str,
callback_name: str,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
litellm_changed_by: Annotated[
str | None,
Header(
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability"
),
] = None,
):
"""
Remove a single callback from a team
The team's other callbacks stay registered and keep firing. Use this instead of
POST /team/{team_id}/disable_logging, which clears every callback on the team at once.
Every entry registered under this callback_name is removed, across callback types, so a
callback registered for both "success" and "failure" is deregistered by one call.
Parameters:
- team_id (str, required): The unique identifier for the team
- callback_name (str, required): The name of the callback to remove, matched exactly as it was
registered with POST /team/{team_id}/callback (e.g. "langfuse", "langsmith", "gcs")
Example curl:
```
curl -X DELETE 'http://localhost:4000/team/dbe2f686-a686-4896-864a-4c3924458709/callback/langsmith' \
-H 'Authorization: Bearer sk-1234'
```
Covers callbacks registered through POST /team/{team_id}/callback and the Admin UI. Teams still
on the deprecated callback_settings metadata shape hold no such entries, so this returns 404 for
them; POST /team/{team_id}/disable_logging remains the way to clear those.
Returns 404 if the team does not exist, or if callback_name is not registered for the team.
"""
try:
from litellm.proxy._types import CommonProxyErrors
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if prisma_client is None:
raise _callback_error(500, CommonProxyErrors.db_not_connected_error.value)
_existing_team: Final = await prisma_client.get_data(
team_id=team_id, table_name="team", query_type="find_unique"
)
if _existing_team is None:
raise _callback_error(404, f"Team id = {team_id} does not exist.")
# IDOR guard: only proxy admins / org admins / team admins of THIS team may
# deregister its callbacks, otherwise any authenticated key holder could
# silence another team's observability integration.
await _verify_team_access(
team_obj=LiteLLM_TeamTable(**_existing_team.model_dump()),
user_api_key_dict=user_api_key_dict,
)
team_metadata: Final = _existing_team.metadata
registered_callbacks: Final = team_metadata.get("logging")
entries: Final = registered_callbacks if isinstance(registered_callbacks, list) else ()
remaining_callbacks: Final = [ # mutable-ok: metadata["logging"] is isinstance-checked for list downstream
entry for entry in entries if not (isinstance(entry, dict) and entry.get("callback_name") == callback_name)
]
if len(remaining_callbacks) == len(entries):
raise _callback_error(404, f"callback_name = {callback_name} is not registered for team_id = {team_id}.")
updated_metadata: Final = {**team_metadata, "logging": remaining_callbacks} # mutable-ok: persisted as JSON
encrypted_metadata: Final = encrypt_callback_vars(updated_metadata)
team_metadata_json: Final = json.dumps(encrypted_metadata)
updated_team: Final = await TeamRepository(prisma_client).table.update(
where={"team_id": team_id}, # mutable-ok: prisma where takes a dict literal
data={"metadata": team_metadata_json}, # mutable-ok: prisma data takes a dict literal
# `object_permission` is included so `_refresh_cached_team` doesn't write a
# cached team with the relation nulled out, see team_model_add for the rationale.
include={"object_permission": True}, # mutable-ok: prisma include takes a dict literal
)
if updated_team is None:
raise _callback_error(404, f"Team id = {team_id} does not exist. Error removing team callback")
# Request-time callback resolution reads the cached team, so without this
# the removed callback keeps firing for live keys until the cache expires.
await _refresh_cached_team(
team_row=updated_team,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await _emit_team_callback_audit_log(
team_id=team_id,
before_metadata=team_metadata,
after_metadata=encrypted_metadata,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
# Report what survives with the same resolution the GET endpoint uses, so a
# caller can confirm in one round trip that its other callbacks are intact.
surviving: Final = _resolve_team_callbacks(encrypted_metadata)
response: Final = TeamCallbackDeleteResponse(
status="success",
message=f"Callback {callback_name} removed for team {team_id}",
data=TeamCallbackDeleteResponseData(
team_id=team_id,
success_callbacks=tuple(surviving.success_callback or ()),
failure_callbacks=tuple(surviving.failure_callback or ()),
),
)
except HTTPException:
# Legitimate 4xx (403 from the access guard, 404 for an unknown team or
# an unregistered callback). Re-raise without the error-level log noise
# the catch-all below would produce.
raise
except ProxyException:
raise
except Exception as e:
verbose_proxy_logger.error("litellm.proxy.proxy_server.delete_team_callback(): Exception occurred - %s", e)
verbose_proxy_logger.debug(traceback.format_exc())
raise ProxyException(
message="Internal Server Error, " + str(e),
type=ProxyErrorTypes.internal_server_error.value,
param=getattr(e, "param", "None"),
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
else:
return response
@router.post(
"/team/{team_id}/disable_logging",
tags=["team management"],

View file

@ -61,11 +61,17 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
open_sse_before_first_byte,
)
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
)
from litellm.proxy.common_utils.sse_keepalive import (
wrap_passthrough_sse_bytes_with_keepalive_pings,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.utils import normalize_route_for_root_path
from litellm.repositories.team_repository import TeamRepository
@ -1173,14 +1179,18 @@ async def pass_through_request(
_response_headers.update(callback_headers)
return StreamingResponse(
PassThroughStreamingHandler.chunk_processor(
response=response,
request_body=_parsed_body,
litellm_logging_obj=logging_obj,
endpoint_type=endpoint_type,
start_time=start_time,
passthrough_success_handler_obj=pass_through_endpoint_logging,
url_route=str(url),
wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=PassThroughStreamingHandler.chunk_processor(
response=response,
request_body=_parsed_body,
litellm_logging_obj=logging_obj,
endpoint_type=endpoint_type,
start_time=start_time,
passthrough_success_handler_obj=pass_through_endpoint_logging,
url_route=str(url),
),
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
upstream_headers=response.headers,
),
headers=_response_headers,
status_code=response.status_code,
@ -1245,14 +1255,18 @@ async def pass_through_request(
_response_headers.update(callback_headers)
return StreamingResponse(
PassThroughStreamingHandler.chunk_processor(
response=response,
request_body=_parsed_body,
litellm_logging_obj=logging_obj,
endpoint_type=endpoint_type,
start_time=start_time,
passthrough_success_handler_obj=pass_through_endpoint_logging,
url_route=str(url),
wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=PassThroughStreamingHandler.chunk_processor(
response=response,
request_body=_parsed_body,
litellm_logging_obj=logging_obj,
endpoint_type=endpoint_type,
start_time=start_time,
passthrough_success_handler_obj=pass_through_endpoint_logging,
url_route=str(url),
),
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
upstream_headers=response.headers,
),
headers=_response_headers,
status_code=response.status_code,
@ -1787,28 +1801,39 @@ def create_pass_through_route(
elif isinstance(custom_body_data, dict):
final_custom_body = custom_body_data
try:
return await pass_through_request(
request=request,
target=full_target,
custom_headers=headers_dict,
user_api_key_dict=user_api_key_dict,
forward_headers=cast(bool | None, param_forward_headers),
merge_query_params=cast(bool | None, param_merge_query_params),
query_params=final_query_params,
default_query_params=cast(dict | None, param_default_query_params),
stream=is_streaming_request or stream,
custom_body=final_custom_body,
cost_per_request=cast(float | None, param_cost_per_request),
custom_llm_provider=custom_llm_provider,
guardrails_config=cast(dict | None, param_guardrails),
timeout=cast(float | None, param_timeout),
)
finally:
if hasattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY):
delattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY)
if hasattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY):
delattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY)
is_stream: Final = bool(is_streaming_request or stream)
async def _relay() -> Response:
try:
return await pass_through_request(
request=request,
target=full_target,
custom_headers=headers_dict,
user_api_key_dict=user_api_key_dict,
forward_headers=cast(bool | None, param_forward_headers),
merge_query_params=cast(bool | None, param_merge_query_params),
query_params=final_query_params,
default_query_params=cast(dict | None, param_default_query_params),
stream=is_stream,
custom_body=final_custom_body,
cost_per_request=cast(float | None, param_cost_per_request),
custom_llm_provider=custom_llm_provider,
guardrails_config=cast(dict | None, param_guardrails),
timeout=cast(float | None, param_timeout),
)
finally:
if hasattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY):
delattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY)
if hasattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY):
delattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY)
# The upstream withholds its response headers until its first token, so
# the whole time-to-first-token is spent inside _relay with nothing on
# the wire. Off unless an operator sets an interval.
return await open_sse_before_first_byte(
_relay(),
ping_interval_seconds=(litellm.sse_keepalive_ping_interval_seconds if is_stream else None),
)
setattr(endpoint_func, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True)
return endpoint_func

View file

@ -657,6 +657,7 @@ from litellm.types.proxy.model_deprecation import (
)
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.router import (
ClassifierPlugin,
DeploymentTypedDict,
RouterGeneralSettings,
RoutingPlugin,
@ -4034,17 +4035,70 @@ def resolve_complexity_router_plugins(
) -> None:
"""
Resolves `complexity_router_config["plugins"]` dotted-path strings to live
instances in place, via `resolve_routing_plugins`.
instances in place, via `resolve_routing_plugins`, and
`complexity_router_config["classifier_plugin"]` via `resolve_classifier_plugin`.
"""
plugin_paths: Final = complexity_router_config.get("plugins")
if not isinstance(plugin_paths, list):
return
if isinstance(plugin_paths, list):
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}",
)
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}",
)
classifier_plugin_path: Final = complexity_router_config.get("classifier_plugin")
if isinstance(classifier_plugin_path, str):
resolved_classifier: Final = resolve_classifier_plugin(
plugin_path=classifier_plugin_path,
config_file_path=config_file_path,
source_label=f"complexity_router_config.classifier_plugin on model {model_name!r}",
)
complexity_router_config["classifier_plugin"] = resolved_classifier # rebind-ok: out-param, resolved in place
def pin_complexity_router_model_id(model: dict) -> None: # mutable-ok: out-param, model_info is stamped in place
"""
Stamps `model_info.id` from the raw litellm_params before plugin resolution swaps
dotted-path strings for live instances. `_delete_deployment` re-reads the raw config
and re-hashes these params to decide which ids the config wants served; an id the
Router derived from the resolved params would never match that hash, so the reconcile
would evict every plugin-bearing deployment one sync after startup.
"""
litellm_params: Final = model.get("litellm_params")
if not isinstance(litellm_params, dict) or not isinstance(litellm_params.get("complexity_router_config"), dict):
return
model_info = model.get("model_info")
if not isinstance(model_info, dict):
model_info = {} # mutable-ok: fresh model_info stamped onto the raw yaml model dict
model["model_info"] = model_info # rebind-ok: out-param, stamped in place
if model_info.get("id") is None:
model_info["id"] = litellm.Router.generate_model_id(
model_group=model.get("model_name", ""),
litellm_params=litellm_params,
)
def resolve_classifier_plugin(
plugin_path: str,
config_file_path: str | None,
source_label: str,
) -> ClassifierPlugin:
"""
Resolves a classifier-plugin dotted path to a live `ClassifierPlugin` instance, with the
same load-time interface check `resolve_routing_plugins` applies to routing plugins: a
sync `def classify` passes the runtime_checkable isinstance and would only fail on the
first classified request, so reject it here where the error names the config key.
"""
resolved: Final = get_instance_fn(value=plugin_path, config_file_path=config_file_path)
if not isinstance(resolved, ClassifierPlugin) or not inspect.iscoroutinefunction(
getattr(resolved, "classify", None)
):
raise ValueError(
f"{source_label} entry {plugin_path!r} resolved to {resolved!r}, which does not "
"implement the ClassifierPlugin interface (an async `classify(context)` method). Fix "
"the referenced module before starting the proxy."
)
return resolved
def _swap_in_model_cost_map(new_model_cost_map: dict) -> int:
@ -5266,6 +5320,7 @@ class ProxyConfig:
for k, v in model["litellm_params"].items():
if isinstance(v, str) and v.startswith("os.environ/"):
model["litellm_params"][k] = get_secret(v)
pin_complexity_router_model_id(model)
complexity_router_config = model["litellm_params"].get("complexity_router_config")
if isinstance(complexity_router_config, dict):
resolve_complexity_router_plugins(
@ -5663,7 +5718,7 @@ class ProxyConfig:
model_id = model.get("model_info", {}).get("id", None)
if model_id is None:
## else - generate stable id's ##
model_id = llm_router._generate_model_id(
model_id = llm_router.generate_model_id(
model_group=model["model_name"],
litellm_params=model["litellm_params"],
)

View file

@ -640,6 +640,15 @@ def _pop_use_chat_completions_api_kw(kwargs: dict[str, object]) -> bool:
return bool(use_cc)
_RESPONSES_ROUTING_PREFIX: Final = "responses/"
def _strip_responses_routing_prefix(model: str) -> str:
if not model.startswith(_RESPONSES_ROUTING_PREFIX):
return model
return model[len(_RESPONSES_ROUTING_PREFIX) :]
def _resolve_model_provider_for_responses(
model: str,
custom_llm_provider: str | None,
@ -649,20 +658,20 @@ def _resolve_model_provider_for_responses(
if custom_llm_provider is not None and not litellm_params.custom_llm_provider:
litellm_params.custom_llm_provider = custom_llm_provider
(
model,
custom_llm_provider,
provider_model,
resolved_provider,
dynamic_api_key,
dynamic_api_base,
) = litellm.get_llm_provider(
model=model,
litellm_params=litellm_params,
)
local_vars["custom_llm_provider"] = custom_llm_provider
local_vars["custom_llm_provider"] = resolved_provider
if dynamic_api_key is not None:
litellm_params.api_key = dynamic_api_key
if dynamic_api_base is not None:
litellm_params.api_base = dynamic_api_base
return model, custom_llm_provider
return _strip_responses_routing_prefix(provider_model), resolved_provider
def _apply_managed_file_id_mapping(
@ -1997,7 +2006,7 @@ async def _aresponses_websocket(
litellm_params_dict: Final = get_litellm_params(**kwargs)
(
model,
provider_model,
_custom_llm_provider,
dynamic_api_key,
dynamic_api_base,
@ -2006,6 +2015,7 @@ async def _aresponses_websocket(
api_base=api_base,
api_key=api_key,
)
resolved_model: Final = _strip_responses_routing_prefix(provider_model)
litellm_params_dict["data_residency"] = infer_openai_data_residency(
_custom_llm_provider,
@ -2014,7 +2024,7 @@ async def _aresponses_websocket(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
model=resolved_model,
user=user,
optional_params={},
litellm_params=litellm_params_dict,
@ -2024,7 +2034,7 @@ async def _aresponses_websocket(
responses_api_provider_config: BaseResponsesAPIConfig | None = None
if _custom_llm_provider is not None:
responses_api_provider_config = ProviderConfigManager.get_provider_responses_api_config(
model=model,
model=resolved_model,
provider=litellm.LlmProviders(_custom_llm_provider),
)
@ -2052,7 +2062,7 @@ async def _aresponses_websocket(
remaining_kwargs: Final = {k: v for k, v in kwargs.items() if k not in _explicit_keys}
await base_llm_http_handler.async_responses_websocket(
model=model,
model=resolved_model,
websocket=websocket,
logging_obj=litellm_logging_obj,
responses_api_provider_config=responses_api_provider_config,

View file

@ -204,6 +204,7 @@ from litellm.types.utils import (
CustomPricingLiteLLMParams,
GenericBudgetConfigType,
LiteLLMBatch,
LlmProviders,
ModelInfo,
ModelResponseStream,
StandardLoggingPayload,
@ -3193,7 +3194,7 @@ class Router:
function_name=function_name,
)
model_group: Final = kwargs.get(metadata_variable_name, {}).get("model_group")
_model_id: Final = self._generate_model_id(model_group=model_group, litellm_params=dynamic_litellm_params)
_model_id: Final = self.generate_model_id(model_group=model_group, litellm_params=dynamic_litellm_params)
original_model_id: Final = model_info.get("id")
model_info["id"] = _model_id
model_info["original_model_id"] = original_model_id
@ -5087,6 +5088,13 @@ class Router:
)
kwargs_copy["file"] = file
if custom_llm_provider == LlmProviders.LITELLM_PROXY.value:
kwargs_copy["extra_body"] = MappingProxyType(
{
**(kwargs_copy.get("extra_body") or MappingProxyType({})),
"target_model_names": stripped_model,
}
)
if (
"gcs_bucket_name" in data
): # TODO: Remove this once we have a better way to handle GCS bucket name: Problem is that we need to pass the gcs_bucket_name to the router for the create_file call but it doesn't show up there
@ -7570,13 +7578,14 @@ class Router:
@staticmethod
def _json_default_stable_id(value: object) -> str:
"""json.dumps default= for _generate_model_id: plain str() on an arbitrary
"""json.dumps default= for generate_model_id: plain str() on an arbitrary
object (e.g. a RoutingPlugin instance) falls back to object.__repr__'s
`<module.Class object at 0x...>`, so the hash -- and deployment id -- would
change every restart. Use the class name instead, stable across restarts."""
return f"{type(value).__module__}.{type(value).__qualname__}"
def _generate_model_id(self, model_group: str, litellm_params: dict):
@staticmethod
def generate_model_id(model_group: str, litellm_params: dict) -> str: # mutable-ok: hashed read-only
"""
Helper function to consistently generate the same id for a deployment
@ -7591,14 +7600,14 @@ class Router:
if isinstance(k, str):
parts.append(k)
elif isinstance(k, dict):
parts.append(json.dumps(k, default=self._json_default_stable_id))
parts.append(json.dumps(k, default=Router._json_default_stable_id))
else:
parts.append(str(k))
if isinstance(v, str):
parts.append(v)
elif isinstance(v, dict):
parts.append(json.dumps(v, default=self._json_default_stable_id))
parts.append(json.dumps(v, default=Router._json_default_stable_id))
else:
parts.append(str(v))
@ -8192,7 +8201,7 @@ class Router:
# check if model info has id
if "id" not in _model_info:
_id = self._generate_model_id(_model_name, _litellm_params)
_id = self.generate_model_id(_model_name, _litellm_params)
_model_info["id"] = _id
if _litellm_params.get("organization", None) is not None and isinstance(
@ -9750,7 +9759,7 @@ class Router:
if model_id is None:
model_name = model.get("model_name", "")
litellm_params = model.get("litellm_params", {})
model_id = self._generate_model_id(model_name, litellm_params)
model_id = self.generate_model_id(model_name, litellm_params)
# Update the model_info in the original list
if "model_info" not in model:
model["model_info"] = {}

View file

@ -128,13 +128,18 @@ class AutoRouter(CustomLogger):
"""
from semantic_router.routers import SemanticRouter
from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages
from litellm.router_strategy.auto_router.litellm_encoder import (
LiteLLMRouterEncoder,
)
from litellm.types.router import PreRoutingHookResponse
if messages is None:
# do nothing, return same inputs
resolved_messages: Final = (
messages
if messages is not None
else resolve_structured_messages(messages=None, request_kwargs=request_kwargs)
)
if resolved_messages is None:
return None
routelayer = self.routelayer
@ -153,7 +158,7 @@ class AutoRouter(CustomLogger):
)
self.routelayer = routelayer
message_content: Final = self._extract_text_from_messages(messages)
message_content: Final = self._extract_text_from_messages(resolved_messages)
route_name: Final = self._matched_route_name(routelayer, message_content)
return PreRoutingHookResponse(

View file

@ -26,8 +26,9 @@ from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast
from pydantic import BaseModel, create_model
from litellm._logging import verbose_router_logger
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
from litellm.constants import EMPTY_MAPPING, RETURN_RAW_MODEL_NAME_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs
from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal_call_metadata
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.types.utils import (
@ -655,6 +656,7 @@ class ClassificationOutcome(NamedTuple):
"heuristic_scorer",
"reasoning_override",
"llm_classifier",
"classifier_plugin",
"classifier_fallback",
"default_model_fallback",
]
@ -1119,15 +1121,18 @@ class ComplexityRouter(CustomLogger):
system_prompt: str | None = None,
request_kwargs: dict[str, Any] | None = None,
messages: Sequence[Mapping[str, object]] | None = None,
raw_messages: list[dict[str, Any]] | None = None, # mutable-ok: same shape _run_routing_plugins receives
) -> ClassificationOutcome:
"""
Classify a prompt by complexity, using the LLM classifier when configured.
Falls back to the local heuristic scorer if classifier_type is "heuristic". If the LLM call
fails, times out, or returns an unparseable response, the configured fallback_tier wins on a
custom tier set, and classifier_fallback otherwise decides between the heuristic scorer and
default_model. The outcome's `cause` reports which path actually ran.
or the classifier plugin fails, times out, or produces no usable tier, the configured
fallback_tier wins on a custom tier set, and classifier_fallback otherwise decides between
the heuristic scorer and default_model. The outcome's `cause` reports which path actually ran.
"""
if self.config.classifier_type == "custom":
return await self._classify_with_plugin(prompt, system_prompt, request_kwargs, raw_messages)
if self.config.classifier_type != "llm" or self.config.classifier_llm_config is None:
tier, score, signals, cause = self._score_and_classify(prompt, system_prompt)
return ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause)
@ -1142,26 +1147,86 @@ class ComplexityRouter(CustomLogger):
classifier_cost=classifier_cost,
)
except Exception as e: # noqa: BLE001 -- external LLM call can fail in many distinct ways (timeout, provider error, validation, parse error); any failure must fall back to the configured fallback path
fallback_tier: Final = self.config.fallback_tier
if fallback_tier is not None:
verbose_router_logger.warning(
"ComplexityRouter: LLM classifier failed (%s), routing to fallback_tier %s", e, fallback_tier
)
return ClassificationOutcome(
tier=fallback_tier,
score=None,
signals=(f"classifier-fallback:{fallback_tier}",),
cause="classifier_fallback",
)
verbose_router_logger.warning(
"ComplexityRouter: LLM classifier failed (%s), falling back to %s",
e,
self.config.classifier_fallback,
return self._classifier_failure_outcome(f"LLM classifier failed ({e})", prompt, system_prompt)
def _classifier_failure_outcome(self, reason: str, prompt: str, system_prompt: str | None) -> ClassificationOutcome:
"""The outcome when the LLM classifier or classifier plugin produced no usable tier:
fallback_tier on a custom tier set, classifier_fallback otherwise."""
fallback_tier: Final = self.config.fallback_tier
if fallback_tier is not None:
verbose_router_logger.warning("ComplexityRouter: %s, routing to fallback_tier %s", reason, fallback_tier)
return ClassificationOutcome(
tier=fallback_tier,
score=None,
signals=(f"classifier-fallback:{fallback_tier}",),
cause="classifier_fallback",
)
if self.config.classifier_fallback == "default_model":
return self._default_model_fallback_outcome()
tier, score, signals, cause = self._score_and_classify(prompt, system_prompt)
return ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause)
verbose_router_logger.warning(
"ComplexityRouter: %s, falling back to %s", reason, self.config.classifier_fallback
)
if self.config.classifier_fallback == "default_model":
return self._default_model_fallback_outcome()
tier, score, signals, cause = self._score_and_classify(prompt, system_prompt)
return ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause)
async def _classify_with_plugin(
self,
prompt: str,
system_prompt: str | None,
request_kwargs: dict[str, Any] | None, # mutable-ok: handed to resolve_structured_messages as-is
raw_messages: list[dict[str, Any]] | None, # mutable-ok: same shape _run_routing_plugins receives
) -> ClassificationOutcome:
from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages
from litellm.types.router import RoutingContext
plugin: Final = self.config.classifier_plugin
if plugin is None:
return self._classifier_failure_outcome("classifier_plugin is not set", prompt, system_prompt)
kwargs: Final = request_kwargs if request_kwargs is not None else EMPTY_MAPPING
pools: Final = self._tier_pools()
try:
context: Final = RoutingContext(
raw_messages=raw_messages or (),
structured_messages=resolve_structured_messages(
messages=raw_messages, request_kwargs=request_kwargs or EMPTY_MAPPING
)
or (),
candidate_models=tuple(model for pool in pools.values() for model in pool),
metadata=kwargs.get(get_metadata_variable_name_from_kwargs(kwargs)) or EMPTY_MAPPING,
)
verdict: Final = await asyncio.wait_for(
plugin.classify(context), timeout=self.config.classifier_plugin_timeout_ms / 1000
)
except asyncio.TimeoutError:
return self._classifier_failure_outcome(
f"classifier plugin timed out after {self.config.classifier_plugin_timeout_ms}ms", prompt, system_prompt
)
except Exception as e: # noqa: BLE001 -- an operator hook can fail in arbitrary ways (network, bug); any failure must fall back rather than fail the request
return self._classifier_failure_outcome(f"classifier plugin failed ({e})", prompt, system_prompt)
if verdict is None:
return self._classifier_failure_outcome("classifier plugin declined to classify", prompt, system_prompt)
if not isinstance(verdict, str):
return self._classifier_failure_outcome(
f"classifier plugin returned a non-string verdict of type {type(verdict).__name__}",
prompt,
system_prompt,
)
tier: Final = self.config.resolve_classified_tier(verdict)
if tier is None:
return self._classifier_failure_outcome(
f"classifier plugin returned unknown tier {verdict!r}", prompt, system_prompt
)
tier_key: Final = _tier_name(tier)
if not pools.get(tier_key):
return self._classifier_failure_outcome(
f"classifier plugin returned tier {tier_key!r}, which has no models configured", prompt, system_prompt
)
return ClassificationOutcome(
tier=tier,
score=None,
signals=(f"classifier-plugin:{tier_key}",),
cause="classifier_plugin",
)
def _default_model_fallback_outcome(self) -> ClassificationOutcome:
"""The classifier-failed outcome for classifier_fallback='default_model'.
@ -1402,7 +1467,7 @@ class ComplexityRouter(CustomLogger):
from litellm.types.router import RoutingContext
tier_key: Final = _tier_name(tier)
metadata_key: Final = "litellm_metadata" if "litellm_metadata" in request_kwargs else "metadata"
metadata_key: Final = get_metadata_variable_name_from_kwargs(request_kwargs)
pool: Final = tuple(self._tier_pools().get(tier_key, ()))
if not pool:
# Nothing for the plugins to filter. Falling through would raise the
@ -2218,7 +2283,9 @@ class ComplexityRouter(CustomLogger):
),
)
outcome: Final = await self.aclassify(user_message, system_prompt, request_kwargs, resolved_messages)
outcome: Final = await self.aclassify(
user_message, system_prompt, request_kwargs, resolved_messages, raw_messages=messages
)
tier, score, signals = outcome.tier, outcome.score, outcome.signals
classified_tier: Final = tier
if escalation_keyword is not None:

View file

@ -10,7 +10,7 @@ from typing import Final, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from litellm.types.router import AdaptiveRouterWeights, RoutingPlugin
from litellm.types.router import AdaptiveRouterWeights, ClassifierPlugin, RoutingPlugin
class ComplexityTier(str, Enum):
@ -434,7 +434,7 @@ class ComplexityRouterConfig(BaseModel):
"becomes that tier's rubric bullet; entries named after a built-in tier may omit the "
"description and inherit the built-in criteria. List order is ascending severity and "
"decides which tier wins when several keyword_tier_rules match. Requires classifier_type "
"'llm', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, "
"'llm' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, "
"adaptive selection, session affinity, plugins, tier_labels, and the calibration-example "
"rubric presets are unavailable with a custom tier set: the first four are built on the "
"built-in tier ladder, and the last two rename or exemplify tiers the set replaces."
@ -535,14 +535,31 @@ class ComplexityRouterConfig(BaseModel):
)
# Classifier strategy
classifier_type: Literal["heuristic", "llm"] = Field(
classifier_type: Literal["heuristic", "llm", "custom"] = Field(
default="heuristic",
description="Classification strategy: local regex/keyword scoring, or an LLM call",
description="Classification strategy: local regex/keyword scoring, an LLM call, or a custom classifier plugin",
)
classifier_llm_config: ClassifierLLMConfig | None = Field(
default=None,
description="Configuration for the LLM classifier; required when classifier_type is 'llm'",
)
classifier_plugin: ClassifierPlugin | None = Field(
default=None,
description=(
"Custom classifier deciding the tier; required when classifier_type is 'custom'. In the proxy "
"config, a dotted path to a ClassifierPlugin instance (resolved at startup, like plugins). Its "
"classify(context) receives the request messages and metadata (caller identity included) and "
"returns the name of the tier to route to, or None to decline and let classifier_fallback decide."
),
)
classifier_plugin_timeout_ms: int = Field(
default=3000,
gt=0,
description=(
"Timeout budget for the classifier plugin call, in milliseconds. On expiry the fallback "
"path decides the tier. Only applies when classifier_type is 'custom'."
),
)
classifier_fallback: Literal["heuristic", "default_model"] = Field(
default="heuristic",
@ -553,7 +570,7 @@ class ComplexityRouterConfig(BaseModel):
"which is what a classifier on some other taxonomy wants: a prompt that grades data "
"sensitivity has no use for a complexity score, and scoring one produces a tier unrelated to "
"what the operator configured. Requires default_model when set to 'default_model'. Only "
"applies when classifier_type is 'llm'."
"applies when classifier_type is 'llm' or 'custom'."
),
)
@ -795,9 +812,16 @@ class ComplexityRouterConfig(BaseModel):
return self
@model_validator(mode="after")
def _validate_llm_classifier_config(self) -> "ComplexityRouterConfig":
def _validate_classifier_config(self) -> "ComplexityRouterConfig":
if self.classifier_type == "llm" and self.classifier_llm_config is None:
raise ValueError("classifier_llm_config is required when classifier_type is 'llm'")
if self.classifier_type == "custom" and self.classifier_plugin is None:
raise ValueError("classifier_plugin is required when classifier_type is 'custom'")
if self.classifier_plugin is not None and self.classifier_type != "custom":
raise ValueError(
f"classifier_plugin is set but classifier_type is {self.classifier_type!r}; "
"the plugin would never run. Set classifier_type 'custom' or remove classifier_plugin"
)
return self
@field_validator("fallback_tier", "classification_prompt")
@ -916,9 +940,10 @@ class ComplexityRouterConfig(BaseModel):
)
if duplicated:
raise ValueError(f"tier_definitions names must be unique (case-insensitive): {', '.join(duplicated)}")
if self.classifier_type != "llm":
if self.classifier_type == "heuristic":
raise ValueError(
"tier_definitions requires classifier_type 'llm': the heuristic scorer only produces the built-in tiers"
"tier_definitions requires classifier_type 'llm' or 'custom': the heuristic scorer only "
"produces the built-in tiers"
)
conflicts: Final = self._tier_definition_conflicts()
if conflicts:

View file

@ -22,6 +22,9 @@ class RequestComplexityRouterConfig(ComplexityRouterConfig):
"""
plugins: None = Field(default=None, description="Not settable over HTTP; routing plugins are runtime objects")
classifier_plugin: None = Field( # pyright: ignore[reportIncompatibleVariableOverride] # narrowing to None is the point: runtime objects are not settable over HTTP
default=None, description="Not settable over HTTP; the classifier plugin is a runtime object"
)
class AutoRouterRoutingTestRequest(BaseModel):

View file

@ -351,6 +351,12 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
milvus_text_field: str | None = None
milvus_db_name: str | None = None
milvus_partition_names: list[str] | None = None
valkey_host: str | None = None
valkey_port: int | None = None
valkey_password: str | None = None
valkey_ssl: bool | None = None
valkey_text_field: str | None = None
valkey_embedding_field: str | None = None
@model_validator(mode="before")
@classmethod
@ -956,6 +962,21 @@ class RoutingPlugin(Protocol):
async def run(self, context: RoutingContext) -> RoutingContext: ...
@runtime_checkable
class ClassifierPlugin(Protocol):
"""Interface a custom classifier must implement to run as the complexity router's classifier_type='custom'.
`classify` returns the name of the tier the request belongs to (a built-in tier value or label,
or a tier_definitions name), or None to decline and let classifier_fallback decide.
The context's `candidate_models` is an informational snapshot of every tier's models, unlike
the narrowing surface RoutingPlugin filters: the returned tier decides the pool, so mutating
the list is a no-op.
"""
async def classify(self, context: RoutingContext) -> str | None: ...
class RequestType(str, enum.Enum):
"""Fixed v0 taxonomy. User-extensible types come in v1."""

View file

@ -196,6 +196,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
input_cost_per_token: Required[float | None]
input_cost_per_token_flex: float | None # OpenAI flex service tier pricing
input_cost_per_token_priority: float | None # OpenAI priority service tier pricing
input_cost_per_token_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing
cache_creation_input_token_cost: float | None
cache_creation_input_token_cost_above_200k_tokens: float | None
cache_creation_input_token_cost_above_272k_tokens: float | None
@ -204,9 +205,11 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
cache_creation_input_token_cost_above_1hr: float | None
cache_creation_input_token_cost_flex: float | None # OpenAI flex service tier pricing
cache_creation_input_token_cost_priority: float | None # OpenAI priority service tier pricing
cache_creation_input_token_cost_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing
cache_read_input_token_cost: float | None
cache_read_input_token_cost_flex: float | None # OpenAI flex service tier pricing
cache_read_input_token_cost_priority: float | None # OpenAI priority service tier pricing
cache_read_input_token_cost_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing
cache_read_input_token_cost_above_200k_tokens: float | None
cache_read_input_token_cost_above_200k_tokens_priority: float | None
cache_read_input_token_cost_above_272k_tokens: float | None
@ -238,6 +241,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
output_cost_per_token: Required[float | None]
output_cost_per_token_flex: float | None # OpenAI flex service tier pricing
output_cost_per_token_priority: float | None # OpenAI priority service tier pricing
output_cost_per_token_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing
regional_processing_uplift_multiplier_eu: (
float | None
) # OpenAI EU data-residency uplift multiplier applied to all token costs (e.g. 1.10 = +10%)
@ -2767,11 +2771,14 @@ RoutingDecisionCause = Literal[
# meant anything that filtered `signals` silently changed what the row claimed.
"reasoning_override",
"llm_classifier",
# The LLM classifier failed on a router with an operator-defined tier set, so the
# request routed to the configured fallback_tier without being classified.
# The operator's classifier plugin (classifier_type 'custom') decided the tier.
"classifier_plugin",
# The LLM classifier or classifier plugin failed on a router with an operator-defined
# tier set, so the request routed to the configured fallback_tier without being classified.
"classifier_fallback",
# The LLM classifier failed and classifier_fallback is 'default_model', so the request
# went to default_model without being classified. Distinct from "default_fallback",
# The LLM classifier or classifier plugin failed and classifier_fallback is
# 'default_model', so the request went to default_model without being classified.
# Distinct from "default_fallback",
# which is a tier having no model configured rather than classification not happening.
"default_model_fallback",
"literal_keyword_match",
@ -3020,6 +3027,11 @@ class StandardLoggingGuardrailInformation(TypedDict, total=False):
provider's counter name (e.g. Bedrock's ``contentPolicyUnits``). Kept as a
sibling of guardrail_response so spend-log prompt redaction never drops it."""
guardrail_cost: ReadOnly[float | None]
"""USD cost of this guardrail invocation, priced from ``guardrail_usage`` by the
provider hook. Summed into the request's ``response_cost`` so it counts against
spend and budgets like token cost."""
class EvalVerdict(TypedDict, total=False):
criterion_name: str
@ -3064,6 +3076,7 @@ class GuardrailTracingDetail(TypedDict, total=False):
violation_categories: list[str] | None
guardrail_action: str | None
guardrail_usage: ReadOnly[Mapping[str, int] | None]
guardrail_cost: ReadOnly[float | None]
StandardLoggingPayloadStatus = Literal["success", "failure"]
@ -3103,8 +3116,9 @@ class CostBreakdown(TypedDict, total=False):
cache_creation_cost: float # Cost of cache-write tokens (premium rate)
output_cost: float # Cost of output/completion tokens (includes reasoning if applicable)
reasoning_cost: float # Cost of reasoning tokens (subset of output_cost)
total_cost: float # Total cost (input + output + tool usage)
total_cost: ReadOnly[float] # Total cost (input + output + tool usage + guardrail)
tool_usage_cost: float # Cost of usage of built-in tools
guardrail_cost: ReadOnly[float] # Cost of guardrail invocations billed by the guardrail provider
additional_costs: dict[str, float] # Free-form additional costs (e.g., {"azure_model_router_flat_cost": 0.00014})
original_cost: float # Cost before discount (optional)
discount_percent: float # Discount percentage applied (e.g., 0.05 = 5%) (optional)
@ -3291,6 +3305,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
# This allows any model_info parameter to be set in litellm_params
input_cost_per_token_flex: float | None = None
input_cost_per_token_priority: float | None = None
input_cost_per_token_ultrafast: float | None = None
cache_creation_input_token_cost_above_1hr: float | None = None
cache_creation_input_token_cost_above_200k_tokens: float | None = None
cache_creation_input_token_cost_above_272k_tokens: float | None = None
@ -3298,9 +3313,11 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
cache_creation_input_token_cost_above_272k_tokens_flex: float | None = None
cache_creation_input_token_cost_flex: float | None = None
cache_creation_input_token_cost_priority: float | None = None
cache_creation_input_token_cost_ultrafast: float | None = None
cache_creation_input_audio_token_cost: float | None = None
cache_read_input_token_cost_flex: float | None = None
cache_read_input_token_cost_priority: float | None = None
cache_read_input_token_cost_ultrafast: float | None = None
cache_read_input_token_cost_above_200k_tokens: float | None = None
cache_read_input_token_cost_above_200k_tokens_priority: float | None = None
cache_read_input_token_cost_above_272k_tokens_priority: float | None = None
@ -3327,6 +3344,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
output_cost_per_token_batches: float | None = None
output_cost_per_token_flex: float | None = None
output_cost_per_token_priority: float | None = None
output_cost_per_token_ultrafast: float | None = None
output_cost_per_audio_token: float | None = None
output_cost_per_token_above_128k_tokens: float | None = None
output_cost_per_token_above_200k_tokens: float | None = None
@ -3476,6 +3494,7 @@ all_litellm_params = (
"bos_token",
"eos_token",
"request_timeout",
"client_side_timeout",
"complete_response",
"self",
"client",
@ -3707,6 +3726,7 @@ class LlmProviders(str, Enum):
NSCALE = "nscale"
PG_VECTOR = "pg_vector"
S3_VECTORS = "s3_vectors"
VALKEY = "valkey"
HELICONE = "helicone"
HYPERBOLIC = "hyperbolic"
RECRAFT = "recraft"
@ -3755,9 +3775,10 @@ LlmProvidersSet: Final = {provider.value for provider in LlmProviders}
OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: set[str] = {
LlmProviders.OPENAI.value,
LlmProviders.HOSTED_VLLM.value,
LlmProviders.LITELLM_PROXY.value,
}
ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "vertex_ai"]
ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai"]
LIST_BATCHES_SUPPORTED_PROVIDERS: Final[frozenset[str]] = frozenset(get_args(ListBatchesSupportedProvider))
@ -3994,6 +4015,7 @@ class ServiceTier(Enum):
FLEX = "flex"
PRIORITY = "priority"
FAST = "fast"
ULTRAFAST = "ultrafast"
class DataResidency(Enum):

View file

@ -2110,7 +2110,7 @@ def encode(model="", text="", custom_tokenizer: dict | None = None):
def decode(
model="",
tokens: list[int] = [],
tokens: Sequence[int] = (),
custom_tokenizer: dict | None = None,
skip_special_tokens: bool = True,
):
@ -2132,7 +2132,7 @@ def decode(
return dec
def _strip_huggingface_special_token_ids(tokenizer: Tokenizer, tokens: list[int]) -> list[int]:
def _strip_huggingface_special_token_ids(tokenizer: Tokenizer, tokens: Sequence[int]) -> Sequence[int]:
try:
added_tokens_decoder: Final = tokenizer.get_added_tokens_decoder()
except Exception:
@ -3972,6 +3972,8 @@ def get_optional_params(
thinking: AnthropicThinkingParam | None = None,
web_search_options: OpenAIWebSearchOptions | None = None,
safety_identifier: str | None = None,
store: bool | None = None,
prompt_cache_key: str | None = None,
base_model: str | None = None,
**kwargs,
):
@ -5578,6 +5580,7 @@ def _get_model_info_helper(
input_cost_per_token=_input_cost_per_token,
input_cost_per_token_flex=_model_info.get("input_cost_per_token_flex", None),
input_cost_per_token_priority=_model_info.get("input_cost_per_token_priority", None),
input_cost_per_token_ultrafast=_model_info.get("input_cost_per_token_ultrafast", None),
cache_creation_input_token_cost=_model_info.get("cache_creation_input_token_cost", None),
cache_creation_input_token_cost_above_200k_tokens=_model_info.get(
"cache_creation_input_token_cost_above_200k_tokens", None
@ -5595,6 +5598,9 @@ def _get_model_info_helper(
cache_creation_input_token_cost_priority=_model_info.get(
"cache_creation_input_token_cost_priority", None
),
cache_creation_input_token_cost_ultrafast=_model_info.get(
"cache_creation_input_token_cost_ultrafast", None
),
cache_read_input_token_cost=_model_info.get("cache_read_input_token_cost", None),
prompt_cache_min_tokens=_model_info.get("prompt_cache_min_tokens", None),
cache_read_input_token_cost_above_200k_tokens=_model_info.get(
@ -5617,6 +5623,7 @@ def _get_model_info_helper(
),
cache_read_input_token_cost_flex=_model_info.get("cache_read_input_token_cost_flex", None),
cache_read_input_token_cost_priority=_model_info.get("cache_read_input_token_cost_priority", None),
cache_read_input_token_cost_ultrafast=_model_info.get("cache_read_input_token_cost_ultrafast", None),
cache_creation_input_token_cost_above_1hr=_model_info.get(
"cache_creation_input_token_cost_above_1hr", None
),
@ -5647,6 +5654,7 @@ def _get_model_info_helper(
output_cost_per_token=_output_cost_per_token,
output_cost_per_token_flex=_model_info.get("output_cost_per_token_flex", None),
output_cost_per_token_priority=_model_info.get("output_cost_per_token_priority", None),
output_cost_per_token_ultrafast=_model_info.get("output_cost_per_token_ultrafast", None),
regional_processing_uplift_multiplier_eu=_model_info.get(
"regional_processing_uplift_multiplier_eu", None
),
@ -8732,6 +8740,12 @@ class ProviderConfigManager:
)
return S3VectorsVectorStoreConfig()
elif litellm.LlmProviders.VALKEY == provider:
from litellm.llms.valkey.vector_stores.transformation import (
ValkeyVectorStoreConfig,
)
return ValkeyVectorStoreConfig()
return None
@staticmethod

View file

@ -10052,6 +10052,21 @@
"output_cost_per_second": 0.0066027,
"supports_tool_choice": true
},
"bedrock/guardrails": {
"guardrail_cost_per_unit": {
"automatedReasoningPolicyUnits": 0.00017,
"contentPolicyImageUnits": 0.00075,
"contentPolicyUnits": 0.00015,
"contextualGroundingPolicyUnits": 0.0001,
"sensitiveInformationPolicyFreeUnits": 0.0,
"sensitiveInformationPolicyUnits": 0.0001,
"topicPolicyUnits": 0.00015,
"wordPolicyUnits": 0.0
},
"litellm_provider": "bedrock",
"mode": "guardrail",
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"bedrock/ap-northeast-1/1-month-commitment/anthropic.claude-instant-v1": {
"input_cost_per_second": 0.01475,
"litellm_provider": "bedrock",

View file

@ -186,6 +186,14 @@
"gemini_native_audio": {
"type": "boolean"
},
"guardrail_cost_per_unit": {
"type": "object",
"description": "USD cost per billable guardrail unit, keyed by the provider's usage counter name (e.g. Bedrock's contentPolicyUnits).",
"additionalProperties": {
"type": "number",
"minimum": 0
}
},
"input_cost_per_audio_per_second": {
"type": "number",
"minimum": 0
@ -361,6 +369,7 @@
"chat",
"completion",
"embedding",
"guardrail",
"image_edit",
"image_generation",
"moderation",

View file

@ -2809,6 +2809,13 @@
"vector_stores_search": true
}
},
"valkey": {
"display_name": "Valkey (`valkey`)",
"url": "https://docs.litellm.ai/docs/providers/valkey_vector_stores",
"endpoints": {
"vector_stores_search": true
}
},
"helicone": {
"display_name": "Helicone (`helicone`)",
"url": "https://docs.litellm.ai/docs/providers/helicone",

View file

@ -33,7 +33,7 @@
"limit": 2
},
"B006": {
"limit": 178
"limit": 177
},
"B008": {
"limit": 503
@ -96,7 +96,7 @@
"limit": 10
},
"DTZ007": {
"limit": 19
"limit": 17
},
"DTZ011": {
"limit": 3

View file

@ -16,6 +16,7 @@ EXCLUDED_PROVIDERS = {
"langfuse", # observability, not LLM provider
"humanloop", # observability, not LLM provider
"pg_vector", # database, not LLM provider
"valkey", # database, not LLM provider
"dotprompt", # prompt management, not provider
"vertex_ai_beta", # beta variant, not needed in main table
}

View file

@ -1303,7 +1303,7 @@ def test_consistent_model_id():
"""
- For a given model group + litellm params, assert the model id is always the same
Test on `_generate_model_id`
Test on `generate_model_id`
Test on `set_model_list`
@ -1317,11 +1317,11 @@ def test_consistent_model_id():
"stream_timeout": 0.001,
}
id1 = Router()._generate_model_id(
id1 = Router().generate_model_id(
model_group=model_group, litellm_params=litellm_params
)
id2 = Router()._generate_model_id(
id2 = Router().generate_model_id(
model_group=model_group, litellm_params=litellm_params
)

View file

@ -1841,8 +1841,8 @@ def test_init_auto_router_deployment_duplicate_model_name(mock_auto_router, mode
router.init_auto_router_deployment(deployment)
def test_generate_model_id_with_deployment_model_name(model_list):
"""Test that _generate_model_id works correctly with deployment model_name and handles None values properly"""
def testgenerate_model_id_with_deployment_model_name(model_list):
"""Test that generate_model_id works correctly with deployment model_name and handles None values properly"""
router = Router(model_list=model_list)
# Test case 1: Normal case with valid model_group and litellm_params
@ -1854,7 +1854,7 @@ def test_generate_model_id_with_deployment_model_name(model_list):
}
try:
result = router._generate_model_id(
result = router.generate_model_id(
model_group=model_group, litellm_params=litellm_params
)
assert isinstance(result, str)
@ -1865,7 +1865,7 @@ def test_generate_model_id_with_deployment_model_name(model_list):
# Test case 2: Edge case with None model_group (this should fail as expected - our fix prevents this from happening)
try:
result = router._generate_model_id(
result = router.generate_model_id(
model_group=None, litellm_params=litellm_params
)
pytest.fail(
@ -1888,7 +1888,7 @@ def test_generate_model_id_with_deployment_model_name(model_list):
}
try:
result = router._generate_model_id(
result = router.generate_model_id(
model_group=model_group, litellm_params=litellm_params_with_none_key
)
assert isinstance(result, str)
@ -1899,7 +1899,7 @@ def test_generate_model_id_with_deployment_model_name(model_list):
# Test case 4: Edge case with empty litellm_params
try:
result = router._generate_model_id(model_group=model_group, litellm_params={})
result = router.generate_model_id(model_group=model_group, litellm_params={})
assert isinstance(result, str)
assert len(result) > 0
print(f"✓ Success with empty litellm_params: {result}")
@ -1907,15 +1907,15 @@ def test_generate_model_id_with_deployment_model_name(model_list):
pytest.fail(f"Failed with empty litellm_params: {e}")
# Test case 5: Verify that the same inputs produce the same result (deterministic)
result1 = router._generate_model_id(
result1 = router.generate_model_id(
model_group=model_group, litellm_params=litellm_params
)
result2 = router._generate_model_id(
result2 = router.generate_model_id(
model_group=model_group, litellm_params=litellm_params
)
assert result1 == result2, "Model ID generation should be deterministic"
print("✓ All _generate_model_id tests passed!")
print("✓ All generate_model_id tests passed!")
def test_handle_clientside_credential_with_deployment_model_name(model_list):
@ -1945,13 +1945,13 @@ def test_handle_clientside_credential_with_deployment_model_name(model_list):
# Test that the method doesn't fail when metadata is empty
try:
# This would normally call _generate_model_id internally
# This would normally call generate_model_id internally
# We're testing that the fix prevents the TypeError
model_group = deployment["model_name"] # This is what our fix does
assert model_group == "gpt-4.1"
# Verify that _generate_model_id works with this model_group
result = router._generate_model_id(
# Verify that generate_model_id works with this model_group
result = router.generate_model_id(
model_group=model_group, litellm_params=dynamic_litellm_params
)
assert isinstance(result, str)

View file

@ -89,6 +89,20 @@ def _semantic_cache():
)
@pytest.mark.parametrize(
"cache_type",
[LiteLLMCacheType.REDIS_SEMANTIC, LiteLLMCacheType.VALKEY_SEMANTIC],
)
def test_semantic_cache_embedding_max_input_tokens_reaches_backend(cache_type):
cache = Cache(
type=cache_type,
redis_url="redis://localhost:6379",
similarity_threshold=0.8,
semantic_cache_embedding_max_input_tokens=2048,
)
assert cache.cache.embedding_max_input_tokens == 2048
def test_semantic_cache_key_excludes_prompt_so_paraphrases_share_a_bucket():
cache = _semantic_cache()
tenant = {"user_api_key": "hash-abc"}

View file

@ -4,9 +4,12 @@ from unittest.mock import MagicMock
sys.path.insert(0, os.path.abspath("../../.."))
import litellm
from litellm.caching._embedding_router import (
build_router_embedding_metadata,
resolve_embedding_max_input_tokens,
resolve_embedding_router,
truncate_embedding_input,
)
@ -65,3 +68,40 @@ def test_build_metadata_handles_none_and_does_not_mutate_input():
assert md == {"user_api_key": "sk-x", "semantic-cache-embedding": True}
assert original == {"user_api_key": "sk-x"}
assert build_router_embedding_metadata(None) == {"semantic-cache-embedding": True}
def test_resolve_max_input_tokens_prefers_configured_over_deployment():
router = MagicMock()
router.get_configured_token_limits.return_value = (8191, None)
assert resolve_embedding_max_input_tokens(512, "sem-embed", router) == 512
router.get_configured_token_limits.assert_not_called()
def test_resolve_max_input_tokens_falls_back_to_deployment_limit():
router = MagicMock()
router.get_configured_token_limits.return_value = (8191, 4096)
assert resolve_embedding_max_input_tokens(None, "sem-embed", router) == 8191
router.get_configured_token_limits.assert_called_once_with("sem-embed")
def test_resolve_max_input_tokens_is_none_without_router_or_deployment_limit():
router = MagicMock()
router.get_configured_token_limits.return_value = (None, None)
assert resolve_embedding_max_input_tokens(None, "sem-embed", router) is None
assert resolve_embedding_max_input_tokens(None, "sem-embed", None) is None
def test_truncate_embedding_input_keeps_prompt_within_limit():
prompt = "The quick brown fox jumps over the lazy dog"
assert truncate_embedding_input(prompt, "sem-embed", None) == prompt
assert truncate_embedding_input(prompt, "sem-embed", 100) == prompt
token_count = len(litellm.encode(model="sem-embed", text=prompt))
assert truncate_embedding_input(prompt, "sem-embed", token_count) == prompt
def test_truncate_embedding_input_cuts_prompt_to_token_limit():
prompt = " ".join(f"word{i}" for i in range(400))
truncated = truncate_embedding_input(prompt, "sem-embed", 50)
assert prompt.startswith(truncated)
assert len(truncated) < len(prompt)
assert len(litellm.encode(model="sem-embed", text=truncated)) == 50

View file

@ -43,6 +43,7 @@ def test_qdrant_semantic_cache_initialization(monkeypatch):
qdrant_api_base="http://test.qdrant.local",
qdrant_api_key="test_key",
similarity_threshold=0.8,
embedding_max_input_tokens=512,
)
# Verify the cache was initialized with correct parameters
@ -50,6 +51,7 @@ def test_qdrant_semantic_cache_initialization(monkeypatch):
assert qdrant_cache.qdrant_api_base == "http://test.qdrant.local"
assert qdrant_cache.qdrant_api_key == "test_key"
assert qdrant_cache.similarity_threshold == 0.8
assert qdrant_cache.embedding_max_input_tokens == 512
mock_sync_client_instance.put.assert_called_once_with(
url="http://test.qdrant.local/collections/test_collection/index",
headers={
@ -832,6 +834,7 @@ def test_qdrant_sync_get_cache_routes_through_router(monkeypatch):
cache.sync_client.post.return_value = search_response
router = MagicMock()
router.get_configured_token_limits.return_value = (None, None)
router.embedding = MagicMock(
return_value={"data": [{"embedding": [0.3, 0.3, 0.3]}]}
)
@ -892,6 +895,7 @@ async def test_qdrant_async_embedding_forwards_full_metadata(monkeypatch):
cache.embedding_model = "sem-embed"
router = MagicMock()
router.get_configured_token_limits.return_value = (None, None)
router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]})
monkeypatch.setitem(
sys.modules,
@ -908,3 +912,57 @@ async def test_qdrant_async_embedding_forwards_full_metadata(monkeypatch):
assert md["user_api_key"] == "sk-x"
assert md["user_api_key_team_id"] == "team-1"
assert md["semantic-cache-embedding"] is True
LONG_PROMPT = " ".join(f"token{i}" for i in range(300))
def _token_count(model, text):
import litellm
return len(litellm.encode(model=model, text=text))
def test_qdrant_get_embedding_truncates_to_deployment_max_input_tokens(monkeypatch):
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
cache = QdrantSemanticCache.__new__(QdrantSemanticCache)
cache.embedding_model = "sem-embed"
router = MagicMock()
router.get_configured_token_limits.return_value = (5, None)
router.embedding = MagicMock(return_value={"data": [{"embedding": [0.5, 0.6]}]})
monkeypatch.setitem(
sys.modules,
"litellm.proxy.proxy_server",
_router_proxy_module(router, "sem-embed"),
)
cache._get_embedding(LONG_PROMPT)
sent_input = router.embedding.call_args.kwargs["input"]
assert LONG_PROMPT.startswith(sent_input)
assert _token_count("sem-embed", sent_input) == 5
@pytest.mark.asyncio
async def test_qdrant_async_embedding_explicit_limit_beats_deployment_limit(monkeypatch):
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
cache = QdrantSemanticCache.__new__(QdrantSemanticCache)
cache.embedding_model = "sem-embed"
cache.embedding_max_input_tokens = 3
router = MagicMock()
router.get_configured_token_limits.return_value = (8191, None)
router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]})
monkeypatch.setitem(
sys.modules,
"litellm.proxy.proxy_server",
_router_proxy_module(router, "sem-embed"),
)
await cache._get_async_embedding(LONG_PROMPT)
sent_input = router.aembedding.call_args.kwargs["input"]
assert _token_count("sem-embed", sent_input) == 3

View file

@ -901,6 +901,7 @@ def test_redis_get_embedding_routes_through_router(monkeypatch):
cache.embedding_model = "sem-embed"
router = MagicMock()
router.get_configured_token_limits.return_value = (None, None)
router.embedding = MagicMock(return_value={"data": [{"embedding": [0.5, 0.6]}]})
fake_proxy = types.ModuleType("litellm.proxy.proxy_server")
fake_proxy.llm_router = router
@ -1145,6 +1146,7 @@ async def test_redis_async_embedding_forwards_full_metadata(monkeypatch):
cache.embedding_model = "sem-embed"
router = MagicMock()
router.get_configured_token_limits.return_value = (None, None)
router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]})
fake_proxy = types.ModuleType("litellm.proxy.proxy_server")
fake_proxy.llm_router = router
@ -1162,6 +1164,100 @@ async def test_redis_async_embedding_forwards_full_metadata(monkeypatch):
assert md["semantic-cache-embedding"] is True
LONG_PROMPT = " ".join(f"token{i}" for i in range(300))
def _proxy_with_router(monkeypatch: pytest.MonkeyPatch, router: MagicMock, model_name: str) -> None:
import sys
import types
fake_proxy = types.ModuleType("litellm.proxy.proxy_server")
fake_proxy.llm_router = router
fake_proxy.llm_model_list = [{"model_name": model_name}]
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy)
def _token_count(model: str, text: str) -> int:
import litellm
return len(litellm.encode(model=model, text=text))
def test_redis_get_embedding_truncates_to_deployment_max_input_tokens(monkeypatch):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
cache = RedisSemanticCache.__new__(RedisSemanticCache)
cache.embedding_model = "sem-embed"
router = MagicMock()
router.get_configured_token_limits.return_value = (5, None)
router.embedding = MagicMock(return_value={"data": [{"embedding": [0.5, 0.6]}]})
_proxy_with_router(monkeypatch, router, "sem-embed")
assert cache._get_embedding(LONG_PROMPT) == [0.5, 0.6]
sent_input = router.embedding.call_args.kwargs["input"]
assert LONG_PROMPT.startswith(sent_input)
assert _token_count("sem-embed", sent_input) == 5
assert _token_count("sem-embed", LONG_PROMPT) > 5
@pytest.mark.asyncio
async def test_redis_async_embedding_explicit_limit_beats_deployment_limit(monkeypatch):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
cache = RedisSemanticCache.__new__(RedisSemanticCache)
cache.embedding_model = "sem-embed"
cache.embedding_max_input_tokens = 3
router = MagicMock()
router.get_configured_token_limits.return_value = (8191, None)
router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]})
_proxy_with_router(monkeypatch, router, "sem-embed")
assert await cache._get_async_embedding(LONG_PROMPT) == [0.1, 0.2]
sent_input = router.aembedding.call_args.kwargs["input"]
assert _token_count("sem-embed", sent_input) == 3
def test_redis_get_embedding_truncates_direct_path_with_explicit_limit(monkeypatch):
import sys
import types
from litellm.caching.redis_semantic_cache import RedisSemanticCache
cache = RedisSemanticCache.__new__(RedisSemanticCache)
cache.embedding_model = "text-embedding-3-small"
cache.embedding_max_input_tokens = 4
fake_proxy = types.ModuleType("litellm.proxy.proxy_server")
fake_proxy.llm_router = None
fake_proxy.llm_model_list = None
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy)
with patch(
"litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2]}]}
) as direct_embed:
cache._get_embedding(LONG_PROMPT)
sent_input = direct_embed.call_args.kwargs["input"]
assert _token_count("text-embedding-3-small", sent_input) == 4
def test_redis_semantic_cache_init_stores_embedding_max_input_tokens(monkeypatch):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
cache = RedisSemanticCache(
redis_url="redis://localhost:6379",
similarity_threshold=0.8,
embedding_max_input_tokens=512,
)
assert cache.embedding_max_input_tokens == 512
default_cache = RedisSemanticCache(redis_url="redis://localhost:6379", similarity_threshold=0.8)
assert default_cache.embedding_max_input_tokens is None
def test_redis_init_defers_redisvl_construction(monkeypatch):
semantic_cache_mock = MagicMock()
custom_vectorizer_mock = MagicMock()

View file

@ -105,6 +105,17 @@ def test_init_requires_similarity_threshold():
ValkeySemanticCache(sync_client=MagicMock(), async_client=AsyncMock())
def test_init_stores_embedding_max_input_tokens():
cache = ValkeySemanticCache(
similarity_threshold=0.8,
sync_client=MagicMock(),
async_client=AsyncMock(),
embedding_max_input_tokens=512,
)
assert cache.embedding_max_input_tokens == 512
assert _make_cache().embedding_max_input_tokens is None
def test_init_rejects_cluster_startup_nodes():
with pytest.raises(ValueError, match="cluster-mode-enabled"):
ValkeySemanticCache(

View file

@ -0,0 +1,113 @@
import os
import pytest
import litellm
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import (
bedrock_guardrail_cost,
cost_breakdown_with_guardrail,
guardrail_information_cost,
)
@pytest.fixture
def synthetic_cost_map(monkeypatch):
monkeypatch.setattr(
litellm,
"model_cost",
{
"bedrock/guardrails": {
"guardrail_cost_per_unit": {
"contentPolicyUnits": 0.00015,
"topicPolicyUnits": 0.00015,
"wordPolicyUnits": 0.0,
}
},
"bedrock/eu-west-1/guardrails": {"guardrail_cost_per_unit": {"contentPolicyUnits": 0.0002}},
"bedrock/us-west-2/guardrails": {"guardrail_cost_per_unit": "malformed"},
},
)
def test_bedrock_guardrail_cost_prices_each_counter(synthetic_cost_map):
cost = bedrock_guardrail_cost(
usage_units={"contentPolicyUnits": 2, "topicPolicyUnits": 1, "wordPolicyUnits": 5},
aws_region_name="us-east-1",
)
assert cost == pytest.approx(0.00045)
def test_bedrock_guardrail_cost_prefers_regional_entry(synthetic_cost_map):
cost = bedrock_guardrail_cost(usage_units={"contentPolicyUnits": 1}, aws_region_name="eu-west-1")
assert cost == pytest.approx(0.0002)
def test_bedrock_guardrail_cost_unknown_counter_is_free(synthetic_cost_map):
assert bedrock_guardrail_cost(usage_units={"someFutureCounter": 3}, aws_region_name="us-east-1") == 0.0
def test_bedrock_guardrail_cost_malformed_regional_entry_falls_back(synthetic_cost_map):
cost = bedrock_guardrail_cost(usage_units={"contentPolicyUnits": 1}, aws_region_name="us-west-2")
assert cost == pytest.approx(0.00015)
def test_bedrock_guardrail_cost_no_pricing_entry(monkeypatch):
monkeypatch.setattr(litellm, "model_cost", {})
assert bedrock_guardrail_cost(usage_units={"contentPolicyUnits": 1}, aws_region_name="us-east-1") == 0.0
def test_shipped_bedrock_guardrail_prices_match_aws_pricing_page():
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
assert litellm.model_cost["bedrock/guardrails"]["guardrail_cost_per_unit"] == {
"automatedReasoningPolicyUnits": 0.00017,
"contentPolicyImageUnits": 0.00075,
"contentPolicyUnits": 0.00015,
"contextualGroundingPolicyUnits": 0.0001,
"sensitiveInformationPolicyFreeUnits": 0.0,
"sensitiveInformationPolicyUnits": 0.0001,
"topicPolicyUnits": 0.00015,
"wordPolicyUnits": 0.0,
}
assert "bedrock/guardrails" not in litellm.bedrock_models
def test_guardrail_information_cost_sums_entries():
entries = [
{"guardrail_name": "a", "guardrail_cost": 0.0003},
{"guardrail_name": "b", "guardrail_cost": None},
{"guardrail_name": "c"},
{"guardrail_name": "d", "guardrail_cost": 0.0001},
]
assert guardrail_information_cost(entries) == pytest.approx(0.0004)
def test_guardrail_information_cost_single_entry_and_garbage():
assert guardrail_information_cost({"guardrail_cost": 0.0001}) == pytest.approx(0.0001)
assert guardrail_information_cost(None) == 0.0
assert guardrail_information_cost("not-guardrail-info") == 0.0
assert guardrail_information_cost([{"guardrail_cost": "bad"}]) == 0.0
def test_guardrail_information_cost_ignores_negative_and_non_finite():
entries = [
{"guardrail_name": "forged-negative", "guardrail_cost": -0.005},
{"guardrail_name": "forged-nan", "guardrail_cost": float("nan")},
{"guardrail_name": "forged-inf", "guardrail_cost": float("inf")},
{"guardrail_name": "real", "guardrail_cost": 0.0003},
]
assert guardrail_information_cost(entries) == pytest.approx(0.0003)
assert guardrail_information_cost({"guardrail_cost": -1.0}) == 0.0
def test_cost_breakdown_with_guardrail_merges_and_creates():
assert cost_breakdown_with_guardrail(None, 0.0) is None
untouched = {"input_cost": 0.1, "total_cost": 0.4}
assert cost_breakdown_with_guardrail(untouched, 0.0) is untouched
merged = cost_breakdown_with_guardrail({"input_cost": 0.1, "total_cost": 0.4}, 0.0003)
assert merged is not None
assert merged["guardrail_cost"] == pytest.approx(0.0003)
assert merged["total_cost"] == pytest.approx(0.4003)
assert merged["input_cost"] == pytest.approx(0.1)
created = cost_breakdown_with_guardrail(None, 0.0003)
assert created == {"guardrail_cost": 0.0003, "total_cost": 0.0003}

View file

@ -1781,6 +1781,81 @@ def test_service_tier_fallback_pricing():
), f"Standard completion cost mismatch: {std_cost[1]} vs {expected_standard_completion}"
def test_service_tier_ultrafast_pricing():
"""An ultrafast request bills the *_ultrafast rates for all token types.
Regression for the ultrafast service tier being absent from ServiceTier:
the cost-key lookup silently returned the standard keys, undercounting
every ultrafast request.
"""
cached_tokens = 200
cache_write_tokens = 300
text_tokens = 500
usage = Usage(
prompt_tokens=text_tokens + cached_tokens + cache_write_tokens,
completion_tokens=400,
total_tokens=text_tokens + cached_tokens + cache_write_tokens + 400,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=cached_tokens, cache_write_tokens=cache_write_tokens
),
)
model_info: ModelInfo = {
"key": "gpt-5.6-sol",
"input_cost_per_token": 5e-06,
"output_cost_per_token": 3e-05,
"cache_creation_input_token_cost": 6.25e-06,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token_ultrafast": 5e-05,
"output_cost_per_token_ultrafast": 3e-04,
"cache_creation_input_token_cost_ultrafast": 6.25e-05,
"cache_read_input_token_cost_ultrafast": 5e-06,
}
prompt_cost, completion_cost = generic_cost_per_token(
model="gpt-5.6-sol",
usage=usage,
custom_llm_provider="openai",
service_tier="ultrafast",
model_info=model_info,
)
expected_prompt_cost = (
text_tokens * 5e-05 + cached_tokens * 5e-06 + cache_write_tokens * 6.25e-05
)
assert prompt_cost == pytest.approx(expected_prompt_cost)
assert completion_cost == pytest.approx(400 * 3e-04)
def test_service_tier_ultrafast_fallback_pricing():
"""Without *_ultrafast keys an ultrafast request bills the standard rate, not zero.
Guards the suffix fallback in _get_cost_per_unit: "_fast" is a substring of
"_ultrafast", so a shortest-first suffix match would strip the wrong suffix
and price the request at 0.
"""
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
std_prompt_cost, std_completion_cost = generic_cost_per_token(
model="gpt-5.6-sol",
usage=usage,
custom_llm_provider="openai",
service_tier=None,
)
ultrafast_prompt_cost, ultrafast_completion_cost = generic_cost_per_token(
model="gpt-5.6-sol",
usage=usage,
custom_llm_provider="openai",
service_tier="ultrafast",
)
assert std_prompt_cost + std_completion_cost > 0
assert ultrafast_prompt_cost == pytest.approx(std_prompt_cost)
assert ultrafast_completion_cost == pytest.approx(std_completion_cost)
@pytest.mark.parametrize(
"model",
[
@ -2322,7 +2397,11 @@ def test_service_tier_suffixes_constant_in_sync_with_enum():
from litellm.litellm_core_utils.llm_cost_calc.utils import _SERVICE_TIER_SUFFIXES
from litellm.types.utils import ServiceTier
assert _SERVICE_TIER_SUFFIXES == tuple(f"_{st.value}" for st in ServiceTier)
assert set(_SERVICE_TIER_SUFFIXES) == {f"_{st.value}" for st in ServiceTier}
# longest-first so a substring match resolves "_ultrafast" before "_fast"
assert list(_SERVICE_TIER_SUFFIXES) == sorted(
_SERVICE_TIER_SUFFIXES, key=len, reverse=True
)
def test_get_cost_per_unit_falls_back_from_service_tier_key_to_base():

View file

@ -4826,3 +4826,87 @@ async def test_restore_correlation_context_works_across_asyncio_task_boundary():
finally:
trace_id_var.set("")
session_id_var.set("")
def _build_success_payload(logging_obj, kwargs):
import datetime
from litellm.litellm_core_utils.litellm_logging import (
get_standard_logging_object_payload,
)
now = datetime.datetime.now()
return get_standard_logging_object_payload(
kwargs=kwargs,
init_response_obj={},
start_time=now,
end_time=now,
logging_obj=logging_obj,
status="success",
)
def _guardrail_kwargs(response_cost):
return {
"litellm_call_id": "guardrail-cost-call",
"model": "gpt-4o",
"messages": [],
"response_cost": response_cost,
"litellm_params": {
"metadata": {
"standard_logging_guardrail_information": [
{
"guardrail_name": "bedrock-pre",
"guardrail_status": "success",
"guardrail_usage": {"topicPolicyUnits": 1, "contentPolicyUnits": 1},
"guardrail_cost": 0.0003,
},
{"guardrail_name": "no-usage-guardrail", "guardrail_status": "success"},
]
}
},
}
def test_payload_response_cost_includes_guardrail_cost(logging_obj):
"""LIT-5651: provider-billed guardrail cost must count in response_cost."""
payload = _build_success_payload(logging_obj, _guardrail_kwargs(response_cost=0.0000429))
assert payload is not None
assert payload["response_cost"] == pytest.approx(0.0003429)
assert payload["cost_breakdown"] is not None
assert payload["cost_breakdown"]["guardrail_cost"] == pytest.approx(0.0003)
assert payload["cost_breakdown"]["total_cost"] == pytest.approx(0.0003)
assert payload["hidden_params"]["response_cost"] == pytest.approx(0.0000429)
def test_payload_guardrail_cost_merges_into_existing_cost_breakdown(logging_obj):
logging_obj.set_cost_breakdown(
input_cost=0.00003,
output_cost=0.0000129,
total_cost=0.0000429,
cost_for_built_in_tools_cost_usd_dollar=0.0,
)
payload = _build_success_payload(logging_obj, _guardrail_kwargs(response_cost=0.0000429))
assert payload is not None
assert payload["response_cost"] == pytest.approx(0.0003429)
assert payload["cost_breakdown"]["guardrail_cost"] == pytest.approx(0.0003)
assert payload["cost_breakdown"]["total_cost"] == pytest.approx(0.0003429)
assert payload["cost_breakdown"]["input_cost"] == pytest.approx(0.00003)
assert logging_obj.cost_breakdown["total_cost"] == pytest.approx(0.0000429)
def test_payload_without_guardrail_cost_is_unchanged(logging_obj):
kwargs = {
"litellm_call_id": "no-guardrail-call",
"model": "gpt-4o",
"messages": [],
"response_cost": 0.0000429,
"litellm_params": {"metadata": {}},
}
payload = _build_success_payload(logging_obj, kwargs)
assert payload is not None
assert payload["response_cost"] == pytest.approx(0.0000429)
assert payload["cost_breakdown"] is None

View file

@ -916,3 +916,117 @@ def test_mixed_finish_chunk_emits_usage_once_sync():
assert message_deltas[0]["usage"]["output_tokens"] == 7
assert _text_deltas(events) == ["Hi"]
_assert_deltas_match_their_block_type(events)
class _CountingSyncStream:
"""Sync stream recording how many upstream chunks have been pulled."""
def __init__(self, items: List[MagicMock]):
self._items = list(items)
self.pulled = 0
def __iter__(self):
return self
def __next__(self):
if self.pulled >= len(self._items):
raise StopIteration
item = self._items[self.pulled]
self.pulled += 1
return item
class _CountingAsyncStream(_CountingSyncStream):
def __aiter__(self):
return self
async def __anext__(self):
try:
return next(self)
except StopIteration:
raise StopAsyncIteration
def _bedrock_tool_open_then_args() -> List[MagicMock]:
"""The Bedrock Converse shape: ``contentBlockStart`` names the tool and
carries empty arguments, the arguments arrive in later events.
"""
return [
_tool_chunk("call_1", "Write", ""),
_tool_chunk("call_1", None, '{"file_text":'),
_tool_chunk("call_1", None, ' "hello"}'),
_make_chunk(Delta(content=None), finish_reason="tool_calls"),
]
def test_tool_block_start_emitted_without_awaiting_the_next_chunk_sync():
"""Regression test for issue #32004.
A tool_use block opened by a chunk whose delta is empty (Bedrock Converse
sends the tool id/name and its arguments in separate events) must emit
``content_block_start`` off that chunk alone. Holding it until the next
upstream chunk arrives means a provider that delivers tool arguments as a
trailing burst leaves the client with nothing after ``message_start`` for
the whole generation, tripping client and load-balancer idle timeouts.
"""
stream = _CountingSyncStream(_bedrock_tool_open_then_args())
wrapper = AnthropicStreamWrapper(completion_stream=stream, model="claude-x")
assert next(wrapper)["type"] == "message_start"
assert stream.pulled == 0
start = next(wrapper)
assert start["type"] == "content_block_start"
assert start["content_block"] == {
"type": "tool_use",
"id": "call_1",
"name": "Write",
"input": {},
}
assert stream.pulled == 1, (
f"content_block_start was withheld until {stream.pulled} upstream chunks had arrived"
)
@pytest.mark.asyncio
async def test_tool_block_start_emitted_without_awaiting_the_next_chunk_async():
"""Async twin of the sync regression test above (issue #32004)."""
stream = _CountingAsyncStream(_bedrock_tool_open_then_args())
wrapper = AnthropicStreamWrapper(completion_stream=stream, model="claude-x")
assert (await wrapper.__anext__())["type"] == "message_start"
assert stream.pulled == 0
start = await wrapper.__anext__()
assert start["type"] == "content_block_start"
assert start["content_block"]["name"] == "Write"
assert stream.pulled == 1, (
f"content_block_start was withheld until {stream.pulled} upstream chunks had arrived"
)
@pytest.mark.parametrize("is_async", [False, True])
@pytest.mark.asyncio
async def test_tool_block_start_flush_does_not_duplicate_or_drop_events(is_async: bool):
"""Flushing the queued ``content_block_start`` early must not duplicate it,
lose the empty opening delta's successors, or break event ordering.
"""
chunks = _bedrock_tool_open_then_args()
if is_async:
wrapper = AnthropicStreamWrapper(completion_stream=_AsyncStream(chunks), model="claude-x")
events = await _drain_async(wrapper)
else:
wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x")
events = _drain_sync(wrapper)
assert [e["type"] for e in events] == [
"message_start",
"content_block_start",
"content_block_delta",
"content_block_delta",
"content_block_stop",
"message_delta",
"message_stop",
]
assert _input_json_deltas(events) == ['{"file_text":', ' "hello"}']
_assert_deltas_match_their_block_type(events)

View file

@ -1041,3 +1041,309 @@ async def test_executor_failure_is_not_tagged():
)
assert is_advisor_orchestration_failure(exc_info.value) is False
# ---------------------------------------------------------------------------
# 15. The advisor sub-call resolves through the proxy router when the advisor
# model is configured in model_list, instead of dialing the public
# Anthropic API (regression for LIT-5307).
# ---------------------------------------------------------------------------
def _router_with_advisor_deployment(
recorder, advisor_model="claude-opus-4-8", deployment_model=None, model_group_alias=None
):
"""Build a Router whose only deployment is the advisor model on Foundry.
The recorder replaces ``litellm.anthropic_messages`` before construction
because Router binds it at init time, so the returned Router exercises the
real deployment-resolution path and records what it dispatched.
"""
import litellm
from litellm.router import Router
with patch("litellm.anthropic_messages", new=recorder):
return Router(
model_list=[
{
"model_name": advisor_model,
"litellm_params": {
"model": deployment_model or f"azure_ai/{advisor_model}",
"api_base": "http://127.0.0.1:1/foundry",
"api_key": "fake-foundry-key",
},
}
],
model_group_alias=model_group_alias,
num_retries=0,
)
@pytest.mark.asyncio
async def test_advisor_sub_call_routes_through_proxy_router():
import litellm.proxy.proxy_server as proxy_server
from litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor import (
AdvisorOrchestrationHandler,
)
router_calls = []
async def recorder(**kwargs):
router_calls.append(kwargs)
return _make_text_response("Use trial division.", model="claude-opus-4-8")
router = _router_with_advisor_deployment(recorder)
call_count = 0
async def mock_call(model, messages, tools, stream, max_tokens, **kwargs):
nonlocal call_count
call_count += 1
if call_count == 1:
return _make_advisor_tool_use_response()
return _make_text_response("Final answer.")
with (
patch(
"litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor._call_messages_handler",
side_effect=mock_call,
),
patch.object(proxy_server, "llm_router", router),
):
h = AdvisorOrchestrationHandler()
result = await h.handle(
model="executor-model",
messages=MESSAGES,
tools=[{**ADVISOR_TOOL, "model": "claude-opus-4-8"}],
stream=False,
max_tokens=512,
custom_llm_provider="azure_ai",
)
assert call_count == 2
assert len(router_calls) == 1
assert router_calls[0]["model"] == "azure_ai/claude-opus-4-8"
assert router_calls[0]["api_base"] == "http://127.0.0.1:1/foundry"
assert router_calls[0]["api_key"] == "fake-foundry-key"
assert "Final answer." in result["content"][0]["text"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("router_kwargs", "advisor_model"),
[
pytest.param({"model_group_alias": {"advisor": "claude-opus-4-8"}}, "advisor", id="model_group_alias"),
pytest.param(
{"advisor_model": "azure_ai/*", "deployment_model": "azure_ai/*"},
"azure_ai/claude-opus-4-8",
id="wildcard",
),
],
)
async def test_advisor_sub_call_routes_through_router_for_alias_and_wildcard(router_kwargs, advisor_model):
"""Alias and wildcard advisor models resolve through the router like exact model_list matches."""
import litellm.proxy.proxy_server as proxy_server
from litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor import (
AdvisorOrchestrationHandler,
)
router_calls = []
async def recorder(**kwargs):
router_calls.append(kwargs)
return _make_text_response("Use trial division.", model="claude-opus-4-8")
router = _router_with_advisor_deployment(recorder, **router_kwargs)
call_count = 0
async def mock_call(model, messages, tools, stream, max_tokens, **kwargs):
nonlocal call_count
call_count += 1
if call_count == 1:
return _make_advisor_tool_use_response()
return _make_text_response("Final answer.")
with (
patch(
"litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor._call_messages_handler",
side_effect=mock_call,
),
patch.object(proxy_server, "llm_router", router),
):
h = AdvisorOrchestrationHandler()
await h.handle(
model="executor-model",
messages=MESSAGES,
tools=[{**ADVISOR_TOOL, "model": advisor_model}],
stream=False,
max_tokens=512,
custom_llm_provider="azure_ai",
)
assert call_count == 2
assert len(router_calls) == 1
assert router_calls[0]["model"] == "azure_ai/claude-opus-4-8"
assert router_calls[0]["api_base"] == "http://127.0.0.1:1/foundry"
assert router_calls[0]["api_key"] == "fake-foundry-key"
@pytest.mark.asyncio
async def test_advisor_sub_call_bypasses_router_for_unconfigured_model():
"""An advisor model the router doesn't know about keeps the SDK-level path."""
import litellm.proxy.proxy_server as proxy_server
from litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor import (
AdvisorOrchestrationHandler,
)
router_calls = []
async def recorder(**kwargs):
router_calls.append(kwargs)
return _make_text_response("should not be used")
router = _router_with_advisor_deployment(recorder, advisor_model="some-other-model")
call_count = 0
async def mock_call(model, messages, tools, stream, max_tokens, **kwargs):
nonlocal call_count
call_count += 1
if call_count == 1:
return _make_advisor_tool_use_response()
if tools is None:
return _make_text_response("Advice.", model="claude-opus-4-8")
return _make_text_response("Final answer.")
with (
patch(
"litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor._call_messages_handler",
side_effect=mock_call,
),
patch.object(proxy_server, "llm_router", router),
):
h = AdvisorOrchestrationHandler()
await h.handle(
model="executor-model",
messages=MESSAGES,
tools=[{**ADVISOR_TOOL, "model": "claude-opus-4-8"}],
stream=False,
max_tokens=512,
custom_llm_provider="azure_ai",
)
assert router_calls == []
assert call_count == 3
@pytest.mark.asyncio
async def test_advisor_sub_call_client_override_bypasses_router():
"""A caller-supplied api_key/api_base override must not be re-routed."""
import litellm
import litellm.proxy.proxy_server as proxy_server
from litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor import (
AdvisorOrchestrationHandler,
)
router_calls = []
async def recorder(**kwargs):
router_calls.append(kwargs)
return _make_text_response("should not be used")
router = _router_with_advisor_deployment(recorder)
sub_calls = []
async def mock_call(model, messages, tools, stream, max_tokens, **kwargs):
sub_calls.append({"model": model, "tools": tools, **kwargs})
if len(sub_calls) == 1:
return _make_advisor_tool_use_response()
if tools is None:
return _make_text_response("Advice.", model="claude-opus-4-8")
return _make_text_response("Final answer.")
advisor_tool = {
**ADVISOR_TOOL,
"model": "claude-opus-4-8",
"api_key": "client-key",
"api_base": "https://client.example.com",
}
with (
patch(
"litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor._call_messages_handler",
side_effect=mock_call,
),
patch.object(proxy_server, "llm_router", router),
patch.dict(proxy_server.general_settings, {"allow_client_side_credentials": True}),
patch.object(litellm, "user_url_validation", False),
):
h = AdvisorOrchestrationHandler()
await h.handle(
model="executor-model",
messages=MESSAGES,
tools=[advisor_tool],
stream=False,
max_tokens=512,
custom_llm_provider="azure_ai",
)
assert router_calls == []
advisor_sub_calls = [c for c in sub_calls if c["tools"] is None]
assert len(advisor_sub_calls) == 1
assert advisor_sub_calls[0]["api_key"] == "client-key"
assert advisor_sub_calls[0]["api_base"] == "https://client.example.com"
# ---------------------------------------------------------------------------
# 16. In-sequence system rows (e.g. Claude Code SessionStart hook output) are
# excluded from the advisor sub-call context but kept for the executor: a
# trailing system row followed by the appended question turn is rejected
# upstream ("role 'system' must precede an 'assistant' message or end the
# array").
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_advisor_context_excludes_in_sequence_system_rows():
from litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor import (
AdvisorOrchestrationHandler,
)
messages_with_system_row = [
*MESSAGES,
{"role": "system", "content": "SessionStart hook output: prefer functional style."},
]
sub_calls = []
async def mock_call(model, messages, tools, stream, max_tokens, **kwargs):
sub_calls.append({"messages": messages, "tools": tools})
if len(sub_calls) == 1:
return _make_advisor_tool_use_response()
if tools is None:
return _make_text_response("Advice.", model="claude-opus-4-6")
return _make_text_response("Final answer.")
with patch(
"litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor._call_messages_handler",
side_effect=mock_call,
):
h = AdvisorOrchestrationHandler()
await h.handle(
model="openai/gpt-4o-mini",
messages=messages_with_system_row,
tools=[ADVISOR_TOOL],
stream=False,
max_tokens=512,
custom_llm_provider="openai",
)
assert len(sub_calls) == 3
advisor_messages = sub_calls[1]["messages"]
assert sub_calls[1]["tools"] is None
assert [m["role"] for m in advisor_messages if m["role"] == "system"] == []
assert advisor_messages[-1]["role"] == "user"
executor_roles = [m["role"] for m in sub_calls[0]["messages"]]
assert "system" in executor_roles

View file

@ -266,3 +266,125 @@ def test_drop_tool_level_extra_fields_strips_copilot_mcp_server_name():
assert "copilot_mcp_server_name" not in tool
assert result["tools"][0]["type"] == "function"
assert result["tools"][1]["function"]["name"] == "read_file"
def _find_key_anywhere(obj, key: str) -> bool:
if isinstance(obj, dict):
if key in obj:
return True
return any(_find_key_anywhere(v, key) for v in obj.values())
if isinstance(obj, list):
return any(_find_key_anywhere(item, key) for item in obj)
return False
def test_azure_ai_strips_non_openai_spec_message_fields():
"""
Regression for https://github.com/BerriAI/litellm/issues/33961.
Azure AI Foundry backends set additionalProperties=false, so any message
field outside the OpenAI chat-completions schema causes a 400 "Extra inputs
are not permitted". Anthropic-format clients (e.g. Claude Code) echo prior
assistant turns back as history carrying thinking_blocks, a nested thought
signature at tool_calls[].function.provider_specific_fields, and Anthropic
cache_control annotations. transform_request must strip all of these before
the request reaches the upstream.
"""
config = AzureAIStudioConfig()
messages = [
{"role": "user", "content": "Read a file."},
{
"role": "assistant",
"content": "I can help.",
"thinking_blocks": [
{
"type": "thinking",
"thinking": "The user wants me to read a file.",
"signature": "",
"cache_control": {"type": "ephemeral"},
}
],
"provider_specific_fields": {"thought_signature": "sig-top"},
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "read_file",
"arguments": "{}",
"provider_specific_fields": {"thought_signature": "sig-nested"},
},
}
],
},
{"role": "user", "content": "go ahead"},
]
request = config.transform_request(
model="fw-glm-5.2",
messages=messages,
optional_params={},
litellm_params={},
headers={},
)
transformed_messages = request["messages"]
assert not _find_key_anywhere(transformed_messages, "thinking_blocks")
assert not _find_key_anywhere(transformed_messages, "provider_specific_fields")
assert not _find_key_anywhere(transformed_messages, "cache_control")
assistant_message = transformed_messages[1]
assert assistant_message["content"] == "I can help."
assert assistant_message["tool_calls"][0]["function"]["name"] == "read_file"
def test_azure_ai_stripping_does_not_mutate_caller_messages():
"""
The stripping must not touch the caller's messages. LiteLLM reuses the same
message objects when falling back to another provider, so stripping in place
would hand the fallback a conversation history with its thinking blocks and
provider metadata already destroyed.
"""
config = AzureAIStudioConfig()
messages = [
{"role": "user", "content": "Read a file."},
{
"role": "assistant",
"content": "I can help.",
"thinking_blocks": [
{"type": "thinking", "thinking": "Reading the file.", "signature": "sig"}
],
"provider_specific_fields": {"thought_signature": "sig-top"},
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "read_file",
"arguments": "{}",
"provider_specific_fields": {"thought_signature": "sig-nested"},
},
}
],
},
]
request = config.transform_request(
model="fw-glm-5.2",
messages=messages,
optional_params={},
litellm_params={},
headers={},
)
assert not _find_key_anywhere(request["messages"], "thinking_blocks")
original_assistant = messages[1]
assert original_assistant["thinking_blocks"][0]["thinking"] == "Reading the file."
assert original_assistant["provider_specific_fields"] == {"thought_signature": "sig-top"}
assert original_assistant["tool_calls"][0]["function"]["provider_specific_fields"] == {
"thought_signature": "sig-nested"
}

View file

@ -2071,3 +2071,90 @@ async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_reques
retry_authorization = posts[1]["headers"]["Authorization"]
assert retry_authorization.startswith("AWS4-HMAC-SHA256")
assert retry_authorization != first_attempt_headers["Authorization"]
def _make_stub_direct_vector_store_config(response):
from litellm.llms.base_llm.vector_store.transformation import (
BaseDirectVectorStoreConfig,
)
class StubDirectVectorStoreConfig(BaseDirectVectorStoreConfig):
def __init__(self):
super().__init__()
self.sync_calls = []
self.async_calls = []
def execute_search_vector_store_request(self, **kwargs):
self.sync_calls.append(kwargs)
return response
async def aexecute_search_vector_store_request(self, **kwargs):
self.async_calls.append(kwargs)
return response
return StubDirectVectorStoreConfig()
def test_vector_store_search_handler_direct_config_sync_skips_http():
handler = BaseLLMHTTPHandler()
stub_response = {"object": "vector_store.search_results.page", "search_query": "q", "data": []}
config = _make_stub_direct_vector_store_config(stub_response)
logging_obj = Mock()
with patch("litellm.llms.custom_httpx.llm_http_handler._get_httpx_client") as mock_get_client:
result = handler.vector_store_search_handler(
vector_store_id="vs_direct",
query="q",
vector_store_search_optional_params={"max_num_results": 4},
vector_store_provider_config=config,
custom_llm_provider="valkey",
litellm_params=GenericLiteLLMParams(valkey_host="localhost"),
logging_obj=logging_obj,
timeout=12.5,
_is_async=False,
)
assert result is stub_response
mock_get_client.assert_not_called()
assert len(config.sync_calls) == 1
call = config.sync_calls[0]
assert call["vector_store_id"] == "vs_direct"
assert call["query"] == "q"
assert call["timeout"] == 12.5
assert call["vector_store_search_optional_params"] == {"max_num_results": 4}
assert isinstance(call["litellm_params"], dict)
assert call["litellm_params"]["valkey_host"] == "localhost"
pre_call_args = logging_obj.pre_call.call_args.kwargs["additional_args"]
assert pre_call_args["query"] == "q"
assert pre_call_args["vector_store_id"] == "vs_direct"
@pytest.mark.asyncio
async def test_vector_store_search_handler_direct_config_async_skips_http():
handler = BaseLLMHTTPHandler()
stub_response = {"object": "vector_store.search_results.page", "search_query": "q", "data": []}
config = _make_stub_direct_vector_store_config(stub_response)
logging_obj = Mock()
with patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client") as mock_get_client:
result = await handler.vector_store_search_handler(
vector_store_id="vs_direct",
query=["q1", "q2"],
vector_store_search_optional_params={},
vector_store_provider_config=config,
custom_llm_provider="valkey",
litellm_params=GenericLiteLLMParams(valkey_host="localhost"),
logging_obj=logging_obj,
timeout=7.0,
_is_async=True,
)
assert result is stub_response
mock_get_client.assert_not_called()
assert len(config.async_calls) == 1
assert config.async_calls[0]["query"] == ["q1", "q2"]
assert config.async_calls[0]["litellm_params"]["valkey_host"] == "localhost"
assert config.async_calls[0]["timeout"] == 7.0
pre_call_args = logging_obj.pre_call.call_args.kwargs["additional_args"]
assert pre_call_args["query"] == ["q1", "q2"]
assert pre_call_args["vector_store_id"] == "vs_direct"

View file

@ -39,6 +39,9 @@ from litellm.llms.fireworks_ai.common_utils import resolve_fireworks_resource_na
"glm-4p6#accounts/gitlab/deployments/2fb7764c",
"glm-4p6#accounts/gitlab/deployments/2fb7764c",
),
("FW-Kimi-K3", "FW-Kimi-K3"),
("fireworks_ai/FW-Kimi-K3", "FW-Kimi-K3"),
("FW-GLM-5.2-Fast", "FW-GLM-5.2-Fast"),
],
)
def test_resolve_fireworks_resource_name(model, expected):

View file

@ -8,7 +8,7 @@ especially ensuring that encoding_format is not included when not provided.
import json
import os
import sys
from unittest.mock import MagicMock, Mock, patch
from unittest.mock import Mock, patch
import pytest
@ -289,6 +289,49 @@ class TestHostedVLLMEmbeddingTransformation:
assert sent_data["model"] == "BAAI/bge-small-en-v1.5"
assert sent_data["input"] == ["Hello world"]
@pytest.mark.parametrize(
"provider_params",
[
{"extra_body": {"truncate": "END", "input_type": "query"}},
{"truncate": "END", "input_type": "query"},
],
)
def test_provider_params_are_sent_at_the_top_level_of_the_request(self, provider_params: dict[str, object]) -> None:
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
with patch.object(HTTPHandler, "post") as mock_post:
mock_response = Mock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.json.return_value = {
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}],
"model": "nvidia/nv-embedqa-e5-v5",
"usage": {"prompt_tokens": 2, "total_tokens": 2},
}
mock_response.text = json.dumps(mock_response.json.return_value)
mock_post.return_value = mock_response
litellm.embedding(
model="hosted_vllm/nvidia/nv-embedqa-e5-v5",
input=["Hello world"],
api_base="https://integrate.api.nvidia.com/v1",
api_key="fake-key",
client=client,
caching=False,
**provider_params,
)
sent_data = json.loads(mock_post.call_args.kwargs["data"])
assert sent_data["truncate"] == "END"
assert sent_data["input_type"] == "query"
assert "extra_body" not in sent_data
assert sent_data["model"] == "nvidia/nv-embedqa-e5-v5"
assert sent_data["input"] == ["Hello world"]
if __name__ == "__main__":
pytest.main([__file__, "-v", "-s"])

View file

@ -0,0 +1,376 @@
import struct
import sys
from types import SimpleNamespace
from typing import Final
from unittest.mock import MagicMock, patch
from urllib.parse import unquote, urlsplit
import httpx
import pytest
from litellm.llms.valkey.vector_stores.transformation import (
ValkeyVectorStoreConfig,
_ValkeySearchParams,
)
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
class FakeSearchIndex:
def __init__(self, result):
self.result = result
self.searched_query = None
self.searched_query_params = None
def search(self, query, query_params=None):
self.searched_query = query
self.searched_query_params = query_params
return self.result
class FakeRedis:
def __init__(self, result=None):
self.index = FakeSearchIndex(result if result is not None else SimpleNamespace(docs=[]))
self.ft_index_name = None
def ft(self, index_name):
self.ft_index_name = index_name
return self.index
class FakeAsyncSearchIndex(FakeSearchIndex):
async def search(self, query, query_params=None):
self.searched_query = query
self.searched_query_params = query_params
return self.result
class FakeAsyncRedis(FakeRedis):
def __init__(self, result=None):
super().__init__(result)
self.index = FakeAsyncSearchIndex(self.index.result)
class FakeEmbeddingFn:
def __init__(self, embedding):
self.embedding = embedding
self.captured_kwargs = None
def __call__(self, **kwargs):
self.captured_kwargs = kwargs
return SimpleNamespace(data=[{"embedding": self.embedding}])
class FakeAsyncEmbeddingFn(FakeEmbeddingFn):
async def __call__(self, **kwargs):
self.captured_kwargs = kwargs
return SimpleNamespace(data=[{"embedding": self.embedding}])
def _doc(doc_id, distance, **fields):
return SimpleNamespace(id=doc_id, vector_distance=str(distance), **fields)
def _search(config, client=None, query="what is litellm", optional_params=None, litellm_params=None):
return config.execute_search_vector_store_request(
vector_store_id="my_index",
query=query,
vector_store_search_optional_params=optional_params or {},
litellm_logging_obj=MagicMock(),
litellm_params={"litellm_embedding_model": "openai/text-embedding-3-small", **(litellm_params or {})},
)
def test_sync_search_builds_knn_query_with_packed_vector():
embedding_fn = FakeEmbeddingFn([0.1, 0.2, 0.3])
client = FakeRedis()
config = ValkeyVectorStoreConfig(sync_client=client, embedding_fn=embedding_fn)
_search(config, optional_params={"max_num_results": 5})
assert client.ft_index_name == "my_index"
assert client.index.searched_query.query_string() == "*=>[KNN 5 @embedding $vec AS vector_distance]"
args = client.index.searched_query.get_args()
assert args[args.index("DIALECT") + 1] == 2
assert args[args.index("LIMIT") : args.index("LIMIT") + 3] == ["LIMIT", 0, 5]
return_args = args[args.index("RETURN") : args.index("RETURN") + 4]
assert return_args == ["RETURN", 2, "text", "vector_distance"]
assert client.index.searched_query_params == {"vec": struct.pack("<3f", 0.1, 0.2, 0.3)}
def test_sync_search_defaults_to_10_results():
client = FakeRedis()
config = ValkeyVectorStoreConfig(sync_client=client, embedding_fn=FakeEmbeddingFn([1.0]))
_search(config)
assert client.index.searched_query.query_string() == "*=>[KNN 10 @embedding $vec AS vector_distance]"
def test_sync_search_honors_custom_field_names():
client = FakeRedis(result=SimpleNamespace(docs=[_doc("doc:1", 0.5, chunk="custom text")]))
config = ValkeyVectorStoreConfig(sync_client=client, embedding_fn=FakeEmbeddingFn([1.0]))
response = _search(
config,
litellm_params={"valkey_embedding_field": "emb", "valkey_text_field": "chunk"},
)
assert client.index.searched_query.query_string() == "*=>[KNN 10 @emb $vec AS vector_distance]"
assert "chunk" in client.index.searched_query.get_args()
assert response["data"][0]["content"][0]["text"] == "custom text"
def test_sync_search_maps_response_with_inverted_score_sorted_best_first():
client = FakeRedis(
result=SimpleNamespace(docs=[_doc("doc:2", 0.75, text="bye"), _doc("doc:1", 0.25, text="hello world")])
)
config = ValkeyVectorStoreConfig(sync_client=client, embedding_fn=FakeEmbeddingFn([1.0]))
response = _search(config)
assert response["object"] == "vector_store.search_results.page"
assert response["search_query"] == "what is litellm"
assert response["data"][0]["score"] == pytest.approx(0.75)
assert response["data"][0]["content"] == [{"text": "hello world", "type": "text"}]
assert response["data"][0]["file_id"] == "doc:1"
assert response["data"][0]["filename"] == "doc:1"
assert response["data"][1]["score"] == pytest.approx(0.25)
assert response["data"][1]["file_id"] == "doc:2"
def test_sync_search_list_query_joins_all_elements():
embedding_fn = FakeEmbeddingFn([1.0])
config = ValkeyVectorStoreConfig(sync_client=FakeRedis(), embedding_fn=embedding_fn)
response = _search(config, query=["first query", "second query"])
assert embedding_fn.captured_kwargs["input"] == ["first query second query"]
assert response["search_query"] == "first query second query"
def test_socket_timeouts_default_to_bounded_values():
assert ValkeyVectorStoreConfig._socket_timeouts(None) == (5.0, 30.0)
def test_socket_timeouts_derive_from_numeric_request_timeout():
assert ValkeyVectorStoreConfig._socket_timeouts(2.0) == (2.0, 2.0)
assert ValkeyVectorStoreConfig._socket_timeouts(120.0) == (5.0, 120.0)
def test_socket_timeouts_derive_from_httpx_timeout():
timeout = httpx.Timeout(connect=3.0, read=7.0, write=1.0, pool=1.0)
assert ValkeyVectorStoreConfig._socket_timeouts(timeout) == (3.0, 7.0)
def test_sync_search_expands_embedding_config_into_kwargs():
embedding_fn = FakeEmbeddingFn([1.0])
config = ValkeyVectorStoreConfig(sync_client=FakeRedis(), embedding_fn=embedding_fn)
_search(
config,
litellm_params={"litellm_embedding_config": {"api_key": "sk-test", "api_base": "https://embed.example.com"}},
)
assert embedding_fn.captured_kwargs == {
"model": "openai/text-embedding-3-small",
"input": ["what is litellm"],
"api_key": "sk-test",
"api_base": "https://embed.example.com",
}
def test_sync_search_requires_embedding_model():
config = ValkeyVectorStoreConfig(sync_client=FakeRedis(), embedding_fn=FakeEmbeddingFn([1.0]))
with pytest.raises(ValueError, match="litellm_embedding_model is required"):
config.execute_search_vector_store_request(
vector_store_id="my_index",
query="q",
vector_store_search_optional_params={},
litellm_logging_obj=MagicMock(),
litellm_params={},
)
def test_sync_search_requires_valkey_host_without_injected_client(monkeypatch):
monkeypatch.delenv("VALKEY_HOST", raising=False)
monkeypatch.delenv("REDIS_HOST", raising=False)
config = ValkeyVectorStoreConfig(embedding_fn=FakeEmbeddingFn([1.0]))
with pytest.raises(ValueError, match="valkey_host is required"):
_search(config)
_VALKEY_ENV_VARS: Final = (
"VALKEY_HOST",
"VALKEY_PORT",
"VALKEY_PASSWORD",
"REDIS_HOST",
"REDIS_PORT",
"REDIS_PASSWORD",
)
def test_connection_url_building(monkeypatch):
for var in _VALKEY_ENV_VARS:
monkeypatch.delenv(var, raising=False)
full: Final = _ValkeySearchParams.model_validate(
{"valkey_host": "h", "valkey_port": 6380, "valkey_password": "p", "valkey_ssl": True}
)
assert full.connection_url() == "rediss://:p@h:6380"
minimal: Final = _ValkeySearchParams.model_validate({"valkey_host": "h", "valkey_password": ""})
assert minimal.connection_url() == "redis://h:6379"
def test_connection_url_never_borrows_gateway_credentials_from_the_environment(monkeypatch):
monkeypatch.setenv("VALKEY_HOST", "gateway-valkey.internal")
monkeypatch.setenv("VALKEY_PORT", "6380")
monkeypatch.setenv("VALKEY_PASSWORD", "gateway-secret")
monkeypatch.setenv("REDIS_HOST", "gateway-redis.internal")
monkeypatch.setenv("REDIS_PORT", "6381")
monkeypatch.setenv("REDIS_PASSWORD", "gateway-redis-secret")
caller_controlled: Final = _ValkeySearchParams.model_validate({"valkey_host": "attacker.example.com"})
assert caller_controlled.connection_url() == "redis://attacker.example.com:6379"
def test_connection_url_percent_encodes_the_password():
password: Final = "p@ss/w#rd%1:x"
params: Final = _ValkeySearchParams.model_validate({"valkey_host": "h", "valkey_password": password})
parsed: Final = urlsplit(params.connection_url())
assert parsed.hostname == "h"
assert parsed.port == 6379
assert unquote(parsed.password or "") == password
def test_connection_url_accepts_string_booleans_from_the_ui_select():
params: Final = _ValkeySearchParams.model_validate(
{"valkey_host": "h", "valkey_port": "6380", "valkey_ssl": "true"}
)
assert params.connection_url() == "rediss://h:6380"
assert _ValkeySearchParams.model_validate({"valkey_host": "h", "valkey_ssl": "false"}).connection_url() == (
"redis://h:6379"
)
def test_search_rejects_filters():
embedding_fn = FakeEmbeddingFn([1.0])
config = ValkeyVectorStoreConfig(sync_client=FakeRedis(), embedding_fn=embedding_fn)
with pytest.raises(ValueError, match="does not support the filters parameter"):
_search(config, optional_params={"filters": {"category": "docs"}})
assert embedding_fn.captured_kwargs is None
@pytest.mark.asyncio
async def test_async_search_rejects_filters():
aembedding_fn = FakeAsyncEmbeddingFn([1.0])
config = ValkeyVectorStoreConfig(async_client=FakeAsyncRedis(), aembedding_fn=aembedding_fn)
with pytest.raises(ValueError, match="does not support the filters parameter"):
await config.aexecute_search_vector_store_request(
vector_store_id="my_index",
query="q",
vector_store_search_optional_params={"filters": {"category": "docs"}},
litellm_logging_obj=MagicMock(),
litellm_params={"litellm_embedding_model": "openai/text-embedding-3-small"},
)
assert aembedding_fn.captured_kwargs is None
def test_search_rejects_empty_query():
config = ValkeyVectorStoreConfig(sync_client=FakeRedis(), embedding_fn=FakeEmbeddingFn([1.0]))
with pytest.raises(ValueError, match="query must not be empty"):
_search(config, query=[])
@pytest.mark.parametrize("max_num_results", [0, -1, 51])
def test_search_rejects_out_of_range_max_num_results(max_num_results):
embedding_fn = FakeEmbeddingFn([1.0])
config = ValkeyVectorStoreConfig(sync_client=FakeRedis(), embedding_fn=embedding_fn)
with pytest.raises(ValueError, match="max_num_results must be between 1 and 50"):
_search(config, optional_params={"max_num_results": max_num_results})
assert embedding_fn.captured_kwargs is None
def test_search_allows_max_num_results_at_the_upper_bound():
client = FakeRedis()
config = ValkeyVectorStoreConfig(sync_client=client, embedding_fn=FakeEmbeddingFn([1.0]))
_search(config, optional_params={"max_num_results": 50})
assert client.index.searched_query.query_string() == "*=>[KNN 50 @embedding $vec AS vector_distance]"
def test_search_treats_an_explicit_null_max_num_results_as_the_default():
client = FakeRedis()
config = ValkeyVectorStoreConfig(sync_client=client, embedding_fn=FakeEmbeddingFn([1.0]))
_search(config, optional_params={"max_num_results": None})
assert client.index.searched_query.query_string() == "*=>[KNN 10 @embedding $vec AS vector_distance]"
def test_missing_redis_dependency_raises_actionable_error():
config = ValkeyVectorStoreConfig(sync_client=FakeRedis(), embedding_fn=FakeEmbeddingFn([1.0]))
blocked = {name: None for name in list(sys.modules) if name == "redis" or name.startswith("redis.")}
with patch.dict(sys.modules, blocked):
with pytest.raises(ValueError, match="pip install redis"):
_search(config)
@pytest.mark.asyncio
async def test_async_search_builds_knn_query_and_maps_response():
aembedding_fn = FakeAsyncEmbeddingFn([0.5, 0.5])
client = FakeAsyncRedis(result=SimpleNamespace(docs=[_doc("doc:9", 0.1, text="async hit")]))
config = ValkeyVectorStoreConfig(async_client=client, aembedding_fn=aembedding_fn)
response = await config.aexecute_search_vector_store_request(
vector_store_id="my_index",
query=["async query", "part two"],
vector_store_search_optional_params={"max_num_results": 3},
litellm_logging_obj=MagicMock(),
litellm_params={
"litellm_embedding_model": "openai/text-embedding-3-small",
"litellm_embedding_config": {"api_key": "sk-async"},
},
)
assert client.ft_index_name == "my_index"
assert client.index.searched_query.query_string() == "*=>[KNN 3 @embedding $vec AS vector_distance]"
assert client.index.searched_query_params == {"vec": struct.pack("<2f", 0.5, 0.5)}
assert aembedding_fn.captured_kwargs == {
"model": "openai/text-embedding-3-small",
"input": ["async query part two"],
"api_key": "sk-async",
}
assert response["search_query"] == "async query part two"
assert response["data"][0]["score"] == pytest.approx(0.9)
assert response["data"][0]["content"] == [{"text": "async hit", "type": "text"}]
assert response["data"][0]["file_id"] == "doc:9"
def test_create_vector_store_is_not_supported():
config = ValkeyVectorStoreConfig()
with pytest.raises(NotImplementedError, match="search-only"):
config.transform_create_vector_store_request(vector_store_create_optional_params={}, api_base="")
def test_provider_config_manager_returns_valkey_config():
config = ProviderConfigManager.get_provider_vector_stores_config(provider=LlmProviders.VALKEY, api_type=None)
assert isinstance(config, ValkeyVectorStoreConfig)

View file

@ -1,5 +1,6 @@
import os
import sys
from datetime import datetime
from unittest.mock import MagicMock, patch
sys.path.insert(
@ -9,7 +10,14 @@ sys.path.insert(
import pytest
from fastapi import HTTPException, Request
from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import (
LiteLLM_OrganizationMembershipTable,
LiteLLM_UserTable,
LiteLLMRoutes,
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin
from litellm.proxy.auth.route_checks import RouteChecks
@ -3298,3 +3306,76 @@ def test_user_daily_activity_aggregated_not_covered_by_prefix_match():
route="/user/daily/activity/aggregated",
allowed_routes=["/user/daily/activity"],
)
@pytest.mark.parametrize(
"user_role",
[
LitellmUserRoles.INTERNAL_USER.value,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
],
)
def test_organization_daily_activity_reachable_by_non_admin_roles(user_role):
"""The Organization Usage dashboard calls /organization/daily/activity, whose
handler restricts results to organizations the caller is ORG_ADMIN of (and
403s on any other org). That scoping is unreachable unless the route layer
lets a non-proxy-admin through first: the route belongs to no info /
management / org_admin_only list, so self_managed_routes is the only entry
granting it, and dropping it 401s every org admin's Organization Usage view
before the handler ever runs.
"""
user_obj = LiteLLM_UserTable(
user_id="test_user",
user_email="test@example.com",
user_role=user_role,
)
valid_token = UserAPIKeyAuth(user_id="test_user", user_role=user_role)
request = MagicMock(spec=Request)
request.query_params = {}
RouteChecks.non_proxy_admin_allowed_routes_check(
user_obj=user_obj,
_user_role=user_role,
route="/organization/daily/activity",
request=request,
valid_token=valid_token,
request_data={},
)
def test_organization_daily_activity_not_granted_by_org_admin_request_data_branch():
"""The org-admin branch of the route gate cannot grant this route, so the
self_managed_routes entry is load-bearing rather than redundant.
Query params do reach request_data, so the reason is not body-vs-query: it
is the key name. _user_is_org_admin reads ``organization_id`` (singular) and
``organizations``, while this endpoint's filter is ``organization_ids``
(plural), and the dashboard's first page load sends no organization filter
at all. Both shapes are pinned below because renaming the query param would
otherwise silently change which gate is doing the work.
"""
user_obj = LiteLLM_UserTable(
user_id="test_user",
user_email="test@example.com",
user_role=LitellmUserRoles.INTERNAL_USER.value,
organization_memberships=[
LiteLLM_OrganizationMembershipTable(
user_id="test_user",
organization_id="org-a",
user_role=LitellmUserRoles.ORG_ADMIN.value,
created_at=datetime.now(),
updated_at=datetime.now(),
)
],
)
# The dashboard's default page load: no organization filter at all.
assert not _user_is_org_admin(request_data={}, user_object=user_obj)
# The filtered load, naming an org this user really does administer.
assert not _user_is_org_admin(request_data={"organization_ids": "org-a"}, user_object=user_obj)
# The key name the helper would have had to see to grant it.
assert _user_is_org_admin(request_data={"organization_id": "org-a"}, user_object=user_obj)
assert not RouteChecks.check_route_access(
route="/organization/daily/activity",
allowed_routes=LiteLLMRoutes.org_admin_only_routes.value,
)

View file

@ -8,6 +8,9 @@ from fastapi.responses import StreamingResponse
from litellm.proxy.common_request_processing import create_response
from litellm.proxy.common_utils.sse_keepalive import (
ANTHROPIC_PING_SSE_CHUNK,
SSE_COMMENT_PING_BYTES,
resolve_ttft_keepalive_interval,
wrap_passthrough_sse_bytes_with_keepalive_pings,
wrap_sse_stream_with_keepalive_pings,
)
@ -156,3 +159,229 @@ async def test_create_response_streams_ping_first_for_slow_upstream():
collected: Final = [chunk async for chunk in response.body_iterator]
assert collected[0] == ANTHROPIC_PING_SSE_CHUNK
assert collected[-1] == MESSAGE_START_CHUNK
SSE_FRAME_BYTES: Final = b'event: content_block_delta\ndata: {"type": "content_block_delta"}\n\n'
BEDROCK_EVENT_STREAM_CONTENT_TYPE: Final = "application/vnd.amazon.eventstream"
@pytest.mark.asyncio
async def test_passthrough_ping_emitted_while_waiting_for_the_first_upstream_byte():
async def slow_start_stream() -> AsyncGenerator[bytes, None]:
await asyncio.sleep(0.2)
yield SSE_FRAME_BYTES
wrapped: Final = wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=slow_start_stream(),
ping_interval_seconds=0.05,
upstream_headers={"content-type": "text/event-stream"},
)
collected: Final = [chunk async for chunk in wrapped]
assert collected[0] == SSE_COMMENT_PING_BYTES
assert collected[-1] == SSE_FRAME_BYTES
assert b"".join(c for c in collected if c != SSE_COMMENT_PING_BYTES) == SSE_FRAME_BYTES
@pytest.mark.asyncio
@pytest.mark.parametrize("content_type", ["text/event-stream", "text/event-stream; charset=utf-8", "TEXT/Event-Stream"])
async def test_passthrough_wraps_every_spelling_of_the_sse_content_type(content_type: str):
async def slow_start_stream() -> AsyncGenerator[bytes, None]:
await asyncio.sleep(0.2)
yield SSE_FRAME_BYTES
wrapped: Final = wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=slow_start_stream(),
ping_interval_seconds=0.05,
upstream_headers={"content-type": content_type},
)
collected: Final = [chunk async for chunk in wrapped]
assert SSE_COMMENT_PING_BYTES in collected
@pytest.mark.asyncio
@pytest.mark.parametrize(
"content_type",
[BEDROCK_EVENT_STREAM_CONTENT_TYPE, "application/json", "application/x-ndjson", None, "text/event-streamish"],
)
async def test_passthrough_leaves_a_non_sse_transport_untouched(content_type: str | None):
"""A comment spliced into a binary transport (e.g. an AWS event stream) corrupts it."""
async def any_stream() -> AsyncGenerator[bytes, None]:
yield SSE_FRAME_BYTES
stream: Final = any_stream()
assert (
wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=stream,
ping_interval_seconds=0.05,
upstream_headers={} if content_type is None else {"content-type": content_type},
)
is stream
)
await stream.aclose()
@pytest.mark.asyncio
async def test_passthrough_ping_is_never_spliced_into_a_half_delivered_frame():
"""Relayed chunks are raw transport reads, so an upstream can stall mid-frame."""
async def stalls_mid_frame() -> AsyncGenerator[bytes, None]:
yield b'event: content_block_delta\ndata: {"partial":'
await asyncio.sleep(0.3)
yield b"1}\n\n"
wrapped: Final = wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=stalls_mid_frame(),
ping_interval_seconds=0.05,
upstream_headers={"content-type": "text/event-stream"},
)
collected: Final = [chunk async for chunk in wrapped]
assert SSE_COMMENT_PING_BYTES not in collected
assert b"".join(collected) == b'event: content_block_delta\ndata: {"partial":1}\n\n'
@pytest.mark.asyncio
async def test_passthrough_ping_resumes_once_the_stalled_frame_completes():
async def stalls_mid_frame_then_at_boundary() -> AsyncGenerator[bytes, None]:
yield b'event: content_block_delta\ndata: {"partial":'
await asyncio.sleep(0.2)
yield b"1}\n\n"
await asyncio.sleep(0.2)
yield SSE_FRAME_BYTES
wrapped: Final = wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=stalls_mid_frame_then_at_boundary(),
ping_interval_seconds=0.05,
upstream_headers={"content-type": "text/event-stream"},
)
collected: Final = [chunk async for chunk in wrapped]
ping_index: Final = collected.index(SSE_COMMENT_PING_BYTES)
assert collected[:ping_index] == [b'event: content_block_delta\ndata: {"partial":', b"1}\n\n"]
assert collected[-1] == SSE_FRAME_BYTES
@pytest.mark.asyncio
@pytest.mark.parametrize("bad_interval", [None, 0, "abc", float("inf"), float("nan"), "-3"])
async def test_passthrough_invalid_or_disabled_interval_returns_stream_unwrapped(bad_interval: float | str | None):
async def any_stream() -> AsyncGenerator[bytes, None]:
yield SSE_FRAME_BYTES
stream: Final = any_stream()
assert (
wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=stream,
ping_interval_seconds=bad_interval,
upstream_headers={"content-type": "text/event-stream"},
)
is stream
)
await stream.aclose()
@pytest.mark.asyncio
async def test_passthrough_aclose_mid_silence_cancels_upstream_and_runs_its_cleanup():
upstream_cleaned_up: Final = asyncio.Event()
async def hung_stream() -> AsyncGenerator[bytes, None]:
try:
yield SSE_FRAME_BYTES
await asyncio.Event().wait()
finally:
upstream_cleaned_up.set()
wrapped: Final = wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=hung_stream(),
ping_interval_seconds=0.05,
upstream_headers={"content-type": "text/event-stream"},
)
assert await wrapped.__anext__() == SSE_FRAME_BYTES
assert await wrapped.__anext__() == SSE_COMMENT_PING_BYTES
await wrapped.aclose()
assert upstream_cleaned_up.is_set()
@pytest.mark.asyncio
async def test_passthrough_upstream_exception_propagates():
async def failing_stream() -> AsyncGenerator[bytes, None]:
yield SSE_FRAME_BYTES
raise ValueError("upstream broke")
wrapped: Final = wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=failing_stream(),
ping_interval_seconds=5.0,
upstream_headers={"content-type": "text/event-stream"},
)
assert await wrapped.__anext__() == SSE_FRAME_BYTES
with pytest.raises(ValueError, match="upstream broke"):
await wrapped.__anext__()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"split_frame",
[
(b'data: {"a": 1}\n', b"\n"),
(b'data: {"a": 1}\r\n', b"\r\n"),
(b'data: {"a": 1}\r', b"\n\r\n"),
(b'data: {"a": 1}\r', b"\r"),
(b'data: {"a": 1}\r\r', b""),
(b'data: {"a": 1}\n\n', b""),
],
ids=["lf-split", "crlf-split", "crlf-mixed-split", "cr-only-split", "cr-only-whole", "not-split"],
)
async def test_passthrough_sees_a_frame_delimiter_split_across_transport_chunks(split_frame):
"""A raw transport read can end mid-delimiter. Testing only the latest chunk
would leave the stream looking permanently mid-frame, silently disabling the
keepalive the operator configured."""
async def split_delimiter_stream() -> AsyncGenerator[bytes, None]:
for part in split_frame:
if part:
yield part
await asyncio.sleep(0.3)
yield SSE_FRAME_BYTES
wrapped: Final = wrap_passthrough_sse_bytes_with_keepalive_pings(
stream=split_delimiter_stream(),
ping_interval_seconds=0.05,
upstream_headers={"content-type": "text/event-stream"},
)
collected: Final = [chunk async for chunk in wrapped]
assert SSE_COMMENT_PING_BYTES in collected
assert b"".join(c for c in collected if c != SSE_COMMENT_PING_BYTES) == b"".join(split_frame) + SSE_FRAME_BYTES
def _deployment(keepalive_seconds=..., model="openai/gpt-4o"):
params = {"model": model}
if keepalive_seconds is not ...:
params["keepalive_seconds"] = keepalive_seconds
return {"model_name": "m", "litellm_params": params}
@pytest.mark.parametrize(
"deployments, global_interval, expected, why",
[
([], 30.0, 30.0, "no deployments known, the global applies"),
([_deployment()], 30.0, 30.0, "nothing configured, the global applies"),
([_deployment(0)], 30.0, None, "an operator's explicit 0 is a hard disable the global cannot lift"),
([_deployment("0")], 30.0, None, "the same, written as a yaml string"),
([_deployment(15)], 30.0, 15.0, "a deployment value wins over the global"),
([_deployment(15), _deployment(15)], 30.0, 15.0, "agreeing deployments are trusted"),
([_deployment(15), _deployment(60)], 30.0, 30.0, "disagreeing deployments fall back to the global"),
([_deployment(0), _deployment(30)], 30.0, 30.0, "a partial disable is not trusted before one is chosen"),
([_deployment(15)], None, 15.0, "a deployment value applies with no global set"),
([_deployment()], None, None, "nothing anywhere leaves it off"),
],
)
def test_ttft_interval_resolves_through_the_deployments_it_could_land_on(
deployments, global_interval, expected, why
):
assert resolve_ttft_keepalive_interval(deployments, global_interval) == expected, why

View file

@ -5079,24 +5079,200 @@ async def test_apply_guardrail_failure_logs_a_dict_not_a_bare_string():
assert "error" in logged
def test_build_tracing_detail_surfaces_usage_counters():
"""LIT-5650: the billable usage block Bedrock returns per ApplyGuardrail call must
land on the tracing detail as guardrail_usage so it reaches spend logs as a
sibling of guardrail_response (which default redaction replaces wholesale)."""
def test_build_tracing_detail_surfaces_usage_counters_and_cost(monkeypatch):
"""LIT-5650/LIT-5651: AWS-billed usage must land as guardrail_usage priced into guardrail_cost."""
monkeypatch.setattr(
litellm,
"model_cost",
{
"bedrock/guardrails": {
"guardrail_cost_per_unit": {
"topicPolicyUnits": 0.00015,
"contentPolicyUnits": 0.00015,
"wordPolicyUnits": 0.0,
}
}
},
)
guardrail = BedrockGuardrail(guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT")
detail = guardrail._build_tracing_detail(
{
"action": "GUARDRAIL_INTERVENED",
"usage": {"topicPolicyUnits": 1, "contentPolicyUnits": 2, "wordPolicyUnits": 0, "oddball": "not-an-int"},
}
},
aws_region_name="us-east-1",
)
assert detail["guardrail_usage"] == {"topicPolicyUnits": 1, "contentPolicyUnits": 2, "wordPolicyUnits": 0}
assert detail["guardrail_cost"] == pytest.approx(0.00045)
def test_build_tracing_detail_omits_guardrail_usage_when_bedrock_reports_none():
guardrail = BedrockGuardrail(guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT")
assert "guardrail_usage" not in guardrail._build_tracing_detail({"action": "NONE"})
assert "guardrail_usage" not in guardrail._build_tracing_detail({"action": "NONE", "usage": {}})
for detail in (
guardrail._build_tracing_detail({"action": "NONE"}, aws_region_name="us-east-1"),
guardrail._build_tracing_detail({"action": "NONE", "usage": {}}, aws_region_name="us-east-1"),
):
assert "guardrail_usage" not in detail
assert "guardrail_cost" not in detail
@pytest.mark.asyncio
async def test_blocked_chunk_logs_usage_and_cost_of_prior_passed_chunks(monkeypatch):
"""LIT-5651 regression: a block on a later chunk must still bill the chunks AWS already processed."""
monkeypatch.setattr(
litellm,
"model_cost",
{
"bedrock/guardrails": {
"guardrail_cost_per_unit": {
"contentPolicyUnits": 0.00015,
"wordPolicyUnits": 0.0,
}
}
},
)
guardrail = BedrockGuardrail(
guardrailIdentifier="test-guardrail",
guardrailVersion="DRAFT",
chunk_budget_chars=40,
)
too_large_response = MagicMock()
too_large_response.status_code = 429
too_large_response.json.return_value = {
"message": "Input text size (60 text units) exceeds the maximum allowed (1 text units) for the content filter policy"
}
passed_chunk_response = MagicMock()
passed_chunk_response.status_code = 200
passed_chunk_response.json.return_value = {
"action": "NONE",
"outputs": [],
"assessments": [],
"usage": {"contentPolicyUnits": 2, "wordPolicyUnits": 1},
}
blocked_chunk_response = MagicMock()
blocked_chunk_response.status_code = 200
blocked_chunk_response.json.return_value = {
"action": "GUARDRAIL_INTERVENED",
"assessments": [{"contentPolicy": {"filters": [{"type": "HATE", "confidence": "HIGH", "action": "BLOCKED"}]}}],
"outputs": [{"text": "Content blocked"}],
"usage": {"contentPolicyUnits": 3},
}
mock_credentials = MagicMock()
mock_credentials.access_key = "test-access-key"
mock_credentials.secret_key = "test-secret-key"
mock_credentials.token = None
request_data = {
"model": "gpt-4o",
"messages": [
{"role": "user", "content": "a" * 30},
{"role": "user", "content": "b" * 30},
],
}
with (
patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post,
patch.object(guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")),
patch.object(guardrail, "_prepare_request", return_value=MagicMock()),
):
mock_post.side_effect = [too_large_response, passed_chunk_response, blocked_chunk_response]
with pytest.raises(HTTPException):
await guardrail.make_bedrock_api_request(
source="INPUT",
messages=request_data["messages"],
request_data=request_data,
)
assert mock_post.call_count == 3
logged_entries = request_data["metadata"]["standard_logging_guardrail_information"]
assert len(logged_entries) == 1
logged = logged_entries[0]
assert logged["guardrail_usage"] == {"contentPolicyUnits": 5, "wordPolicyUnits": 1}
assert logged["guardrail_cost"] == pytest.approx(0.00075)
assert logged["guardrail_response"]["usage"] == {"contentPolicyUnits": 5, "wordPolicyUnits": 1}
@pytest.mark.asyncio
async def test_terminal_failure_logs_usage_and_cost_of_prior_passed_chunks(monkeypatch):
"""LIT-5651 regression: a terminal failure on a later chunk must still bill the chunks AWS already processed."""
monkeypatch.setattr(
litellm,
"model_cost",
{
"bedrock/guardrails": {
"guardrail_cost_per_unit": {
"contentPolicyUnits": 0.00015,
"wordPolicyUnits": 0.0,
}
}
},
)
guardrail = BedrockGuardrail(
guardrailIdentifier="test-guardrail",
guardrailVersion="DRAFT",
chunk_budget_chars=40,
)
too_large_response = MagicMock()
too_large_response.status_code = 429
too_large_response.json.return_value = {
"message": "Input text size (60 text units) exceeds the maximum allowed (1 text units) for the content filter policy"
}
passed_chunk_response = MagicMock()
passed_chunk_response.status_code = 200
passed_chunk_response.json.return_value = {
"action": "NONE",
"outputs": [],
"assessments": [],
"usage": {"contentPolicyUnits": 2, "wordPolicyUnits": 1},
}
failed_chunk_response = MagicMock()
failed_chunk_response.status_code = 400
failed_chunk_response.json.return_value = {"message": "ValidationException: guardrail is in a failed state"}
mock_credentials = MagicMock()
mock_credentials.access_key = "test-access-key"
mock_credentials.secret_key = "test-secret-key"
mock_credentials.token = None
request_data = {
"model": "gpt-4o",
"messages": [
{"role": "user", "content": "a" * 30},
{"role": "user", "content": "b" * 30},
],
}
with (
patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post,
patch.object(guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")),
patch.object(guardrail, "_prepare_request", return_value=MagicMock()),
):
mock_post.side_effect = [too_large_response, passed_chunk_response, failed_chunk_response]
with pytest.raises(HTTPException):
await guardrail.make_bedrock_api_request(
source="INPUT",
messages=request_data["messages"],
request_data=request_data,
)
assert mock_post.call_count == 3
logged_entries = request_data["metadata"]["standard_logging_guardrail_information"]
assert len(logged_entries) == 1
logged = logged_entries[0]
assert logged["guardrail_status"] == "guardrail_failed_to_respond"
assert logged["guardrail_usage"] == {"contentPolicyUnits": 2, "wordPolicyUnits": 1}
assert logged["guardrail_cost"] == pytest.approx(0.0003)
assert logged["guardrail_response"]["usage"] == {"contentPolicyUnits": 2, "wordPolicyUnits": 1}
assert "error" in logged["guardrail_response"]

View file

@ -386,25 +386,30 @@ def test_flush_deferred_async_logging_noop_when_no_closure_stored():
def test_proxy_finally_block_routes_through_flush_helper():
"""
Source-level contract: the proxy's `base_process_llm_request` finally
block must delegate to `_flush_deferred_async_logging` rather than
inlining the gating logic. Inlining is what allowed the duplicate
Success+Failure spend log to slip in originally — this guards the
refactor.
Source-level contract: the proxy's request-processing finally block must
delegate to `_flush_deferred_async_logging` rather than inlining the gating
logic. Inlining is what allowed the duplicate Success+Failure spend log to
slip in originally — this guards the refactor.
Both halves of the request path are inspected: `base_process_llm_request` is
the public entry point and `_process_llm_request` holds the body, so neither
may inline the reset regardless of which one carries the finally block.
"""
import inspect
src = inspect.getsource(ProxyBaseLLMRequestProcessing.base_process_llm_request)
src = inspect.getsource(ProxyBaseLLMRequestProcessing._process_llm_request) + inspect.getsource(
ProxyBaseLLMRequestProcessing.base_process_llm_request
)
assert "_flush_deferred_async_logging" in src, (
"base_process_llm_request must call _flush_deferred_async_logging "
"from its finally block — do not inline the gating logic."
"the request path must call _flush_deferred_async_logging from its "
"finally block — do not inline the gating logic."
)
# Belt-and-braces: the inlined `_enqueue_deferred_logging = None` reset
# was the symptom of the duplicate-log bug; assert it stays inside the
# helper, not in the request-processing function.
assert "_enqueue_deferred_logging = None" not in src, (
"Reset of _enqueue_deferred_logging must live inside "
"_flush_deferred_async_logging, not in base_process_llm_request."
"_flush_deferred_async_logging, not in the request path."
)

View file

@ -27,6 +27,7 @@ from litellm.proxy.guardrails.usage_endpoints import (
guardrails_usage_detail,
guardrails_usage_logs,
guardrails_usage_overview,
policies_usage_overview,
)
from litellm.types.guardrails import Guardrail, LitellmParams
@ -356,3 +357,82 @@ async def test_logs_resolves_config_guardrail_logical_name():
)
where = prisma.db.litellm_spendlogguardrailindex.find_many.call_args.kwargs["where"]
assert where["guardrail_id"] == {"in": ["yaml-uuid", "yaml-pii"]}
# ---- date window cap (LIT-5762) ---------------------------------------------
@pytest.mark.asyncio
async def test_overview_rejects_range_over_max_days():
prisma = _prisma()
handler = _config_handler()
p1, p2 = _patches(prisma, handler)
with p1, p2, pytest.raises(HTTPException) as exc:
await guardrails_usage_overview(start_date="2020-01-01", end_date=END, user_api_key_dict=ADMIN)
assert exc.value.status_code == 400
assert "366" in str(exc.value.detail)
@pytest.mark.asyncio
async def test_overview_accepts_range_at_exactly_max_days():
prisma = _prisma()
handler = _config_handler()
p1, p2 = _patches(prisma, handler)
with p1, p2:
resp = await guardrails_usage_overview(start_date="2025-04-26", end_date="2026-04-27", user_api_key_dict=ADMIN)
assert resp.totalRequests == 0
@pytest.mark.asyncio
async def test_overview_rejects_malformed_dates():
prisma = _prisma()
handler = _config_handler()
p1, p2 = _patches(prisma, handler)
with p1, p2, pytest.raises(HTTPException) as exc:
await guardrails_usage_overview(start_date="not-a-date", end_date=END, user_api_key_dict=ADMIN)
assert exc.value.status_code == 400
@pytest.mark.asyncio
async def test_overview_rejects_non_canonical_date_format():
prisma = _prisma()
handler = _config_handler()
p1, p2 = _patches(prisma, handler)
with p1, p2, pytest.raises(HTTPException) as exc:
await guardrails_usage_overview(start_date="20260420", end_date=END, user_api_key_dict=ADMIN)
assert exc.value.status_code == 400
assert "YYYY-MM-DD" in str(exc.value.detail)
@pytest.mark.asyncio
async def test_detail_rejects_reversed_dates():
prisma = _prisma(find_unique=_db_row())
handler = _config_handler()
p1, p2 = _patches(prisma, handler)
with p1, p2, pytest.raises(HTTPException) as exc:
await guardrails_usage_detail(guardrail_id="db-1", start_date=END, end_date=START, user_api_key_dict=ADMIN)
assert exc.value.status_code == 400
@pytest.mark.asyncio
async def test_policies_overview_rejects_range_over_max_days():
prisma = _prisma()
handler = _config_handler()
p1, p2 = _patches(prisma, handler)
with p1, p2, pytest.raises(HTTPException) as exc:
await policies_usage_overview(start_date="2020-01-01", end_date=END, user_api_key_dict=ADMIN)
assert exc.value.status_code == 400
@pytest.mark.asyncio
async def test_detail_prev_trend_query_is_bounded():
"""Regression: the trend query scanned every metrics row before start_date."""
prisma = _prisma(find_unique=_db_row())
handler = _config_handler()
p1, p2 = _patches(prisma, handler)
with p1, p2:
await guardrails_usage_detail(guardrail_id="db-1", start_date=START, end_date=END, user_api_key_dict=ADMIN)
wheres = [c.kwargs["where"] for c in prisma.db.litellm_dailyguardrailmetrics.find_many.await_args_list]
prev_wheres = [w for w in wheres if "lt" in w.get("date", {})]
assert prev_wheres
assert all("gte" in w["date"] for w in prev_wheres)

View file

@ -17,7 +17,7 @@ from litellm.proxy.hooks.proxy_track_cost_callback import (
_should_track_cost_callback,
_update_database_and_spend_counters,
)
from litellm.types.utils import CallTypes
from litellm.types.utils import CallTypes, Usage
@pytest.mark.asyncio
@ -152,6 +152,70 @@ async def test_async_post_call_failure_hook_does_not_clobber_guardrail_info_in_m
assert metadata["standard_logging_guardrail_information"] == metadata_bucket_info
@pytest.mark.asyncio
async def test_async_post_call_failure_hook_bills_guardrail_cost_on_blocked_request():
"""LIT-5651: a request blocked by a guardrail never reaches the LLM, but the
guardrail invocation itself is billed by the provider. The failure row must
charge that cost against the key instead of recording zero spend."""
logger = _ProxyDBLogger()
request_data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {
"standard_logging_guardrail_information": [
{
"guardrail_name": "bedrock-guard",
"guardrail_status": "guardrail_intervened",
"guardrail_usage": {"topicPolicyUnits": 1, "contentPolicyUnits": 1},
"guardrail_cost": 0.0003,
}
]
},
"proxy_server_request": {"request_id": "test_request_id"},
}
with patch(
"litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database",
new_callable=AsyncMock,
) as mock_update_database:
await logger.async_post_call_failure_hook(
request_data=request_data,
original_exception=Exception("Violated guardrail policy"),
user_api_key_dict=UserAPIKeyAuth(api_key="test_api_key"),
)
assert mock_update_database.call_args[1]["response_cost"] == pytest.approx(0.0003)
@pytest.mark.asyncio
async def test_async_post_call_failure_hook_adds_guardrail_cost_to_recovered_stream_cost():
logger = _ProxyDBLogger()
request_data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {
"standard_logging_guardrail_information": [
{"guardrail_name": "bedrock-guard", "guardrail_status": "success", "guardrail_cost": 0.0003}
]
},
"proxy_server_request": {"request_id": "test_request_id"},
"combined_usage_object": Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
"response_cost": 0.001,
}
with patch(
"litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database",
new_callable=AsyncMock,
) as mock_update_database:
await logger.async_post_call_failure_hook(
request_data=request_data,
original_exception=Exception("stream broke mid-flight"),
user_api_key_dict=UserAPIKeyAuth(api_key="test_api_key"),
)
assert mock_update_database.call_args[1]["response_cost"] == pytest.approx(0.0013)
@pytest.mark.asyncio
async def test_async_post_call_failure_hook_non_llm_route():
# Setup

View file

@ -286,6 +286,13 @@ def test_semantic_matching_without_an_embedding_model_is_rejected():
_request("what is 2+2", semantic_keyword_matching=True)
def test_classifier_plugin_is_not_settable_over_http():
"""classifier_plugin holds a live runtime object, closed off like `plugins`; a plugin-mode
config is therefore unrepresentable in a request body."""
with pytest.raises(ValidationError):
_request("what is 2+2", classifier_type="custom", classifier_plugin="my_module.instance")
class TestAutoRouterBenchmarks:
from litellm.proxy.management_endpoints.auto_router_endpoints import _SessionAggRow

View file

@ -992,3 +992,51 @@ def test_build_budget_write_data_clears_reset_at_with_null_duration():
data = build_budget_write_data({"budget_duration": None}, "admin-1")
assert data["budget_duration"] is None
assert data["budget_reset_at"] is None
@pytest.mark.asyncio
async def test_get_organization_daily_activity_non_admin_without_org_admin_role_sees_nothing(
monkeypatch,
):
"""A caller who is ORG_ADMIN of no organization must resolve to an EMPTY id
list, never to None. None means "no entity filter" downstream, i.e. every
organization's spend, so the natural simplification of falling back to None
on an empty membership set turns a scoping rule into a proxy-wide leak. The
organization-alias lookup must be scoped by that same empty list rather than
reading the whole table.
"""
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints import organization_endpoints
from litellm.proxy.management_endpoints.organization_endpoints import (
get_organization_daily_activity,
)
mock_prisma_client = AsyncMock()
org_table_find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_organizationtable.find_many = org_table_find_many
mock_prisma_client.db.litellm_organizationmembership.find_many = AsyncMock(return_value=[])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.organization_endpoints._user_has_admin_view",
lambda _: False,
)
get_daily_activity_mock = AsyncMock(return_value=MagicMock(name="SpendAnalyticsPaginatedResponse"))
monkeypatch.setattr(organization_endpoints, "get_daily_activity", get_daily_activity_mock)
auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="no-orgs-user")
await get_organization_daily_activity(
organization_ids=None,
start_date="2024-04-01",
end_date="2024-04-30",
model=None,
api_key=None,
page=1,
page_size=10,
exclude_organization_ids=None,
user_api_key_dict=auth,
)
assert get_daily_activity_mock.call_args.kwargs["entity_id"] == []
assert org_table_find_many.call_args.kwargs["where"] == {"organization_id": {"in": []}}

View file

@ -21,6 +21,7 @@ from litellm.proxy._types import (
)
from litellm.proxy.management_endpoints.team_callback_endpoints import (
add_team_callbacks,
delete_team_callback,
disable_team_logging,
get_team_callbacks,
)
@ -942,3 +943,465 @@ async def test_disable_team_logging_leaves_team_re_enablable():
written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"])
assert [entry["callback_name"] for entry in written["logging"]] == ["langfuse"]
def _two_callback_metadata() -> dict:
"""A team with two tenants' integrations registered, the LIT-5161 shape."""
return {
"logging": [
{
"callback_name": "langsmith",
"callback_type": "success",
"callback_vars": {
"langsmith_api_key": "ls-demo",
"langsmith_project": "demo",
},
},
{
"callback_name": "langfuse",
"callback_type": "success",
"callback_vars": {
"langfuse_public_key": "pk-demo",
"langfuse_secret_key": "sk-demo",
},
},
]
}
@pytest.mark.asyncio
async def test_delete_team_callback_rejects_unauthorized_caller(patched_prisma, unauthorized_caller):
with pytest.raises(HTTPException) as exc:
await delete_team_callback(
http_request=Mock(spec=Request),
team_id="team-victim",
callback_name="langsmith",
user_api_key_dict=unauthorized_caller,
)
assert exc.value.status_code == 403
patched_prisma.db.litellm_teamtable.update.assert_not_called()
@pytest.mark.asyncio
async def test_delete_team_callback_removes_only_the_named_callback():
"""The ticket's scenario: one tenant deregisters without touching the others.
disable_logging is the only other removal route and it drops every callback
on the team, so the surviving entry has to come through this write intact,
credentials included.
"""
mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=_two_callback_metadata()))
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.master_key", None),
):
response = await delete_team_callback(
http_request=MagicMock(spec=Request),
team_id="team-1",
callback_name="langsmith",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"])
assert [entry["callback_name"] for entry in written["logging"]] == ["langfuse"]
assert written["logging"][0]["callback_vars"].keys() == {
"langfuse_public_key",
"langfuse_secret_key",
}
assert response.status == "success"
assert response.data.team_id == "team-1"
assert response.data.success_callbacks == ("langfuse",)
assert response.data.failure_callbacks == ()
@pytest.mark.asyncio
async def test_delete_team_callback_leaves_the_other_callback_firing():
"""The survivor has to still be live, not merely still stored.
Asks the real request-time resolver what the written row would do, the same
way the disable_logging regression test does.
"""
from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata
mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=_two_callback_metadata()))
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.master_key", None),
):
await delete_team_callback(
http_request=MagicMock(spec=Request),
team_id="team-1",
callback_name="langsmith",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"])
resolved = _get_dynamic_logging_metadata(
UserAPIKeyAuth(api_key="hashed", team_id="team-1", team_metadata=written),
proxy_config=MagicMock(**{"load_team_config.return_value": {}}),
)
assert resolved is not None
assert resolved.success_callback == ["langfuse"]
assert "langsmith" not in resolved.success_callback
assert resolved.callback_vars.get("langfuse_public_key") == "pk-demo"
@pytest.mark.asyncio
async def test_delete_team_callback_removes_every_type_under_that_name():
"""A callback registered for both events is deregistered by one call.
add_team_callbacks keys its duplicate check on (callback_name, callback_type),
so the same destination can hold a success entry and a failure entry. Removing
only one of them would leave the team still sending to it.
"""
metadata = {
"logging": [
{
"callback_name": "langfuse",
"callback_type": "success",
"callback_vars": {"langfuse_public_key": "pk-demo"},
},
{
"callback_name": "langsmith",
"callback_type": "success",
"callback_vars": {"langsmith_project": "demo"},
},
{
"callback_name": "langfuse",
"callback_type": "failure",
"callback_vars": {"langfuse_public_key": "pk-demo"},
},
]
}
mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=metadata))
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.master_key", None),
):
response = await delete_team_callback(
http_request=MagicMock(spec=Request),
team_id="team-1",
callback_name="langfuse",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"])
assert [entry["callback_name"] for entry in written["logging"]] == ["langsmith"]
assert response.data.success_callbacks == ("langsmith",)
assert response.data.failure_callbacks == ()
@pytest.mark.asyncio
async def test_delete_team_callback_404s_for_unregistered_callback():
"""An unregistered name must not rewrite the team's metadata."""
mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=_two_callback_metadata()))
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.master_key", None),
):
with pytest.raises(HTTPException) as exc:
await delete_team_callback(
http_request=MagicMock(spec=Request),
team_id="team-1",
callback_name="gcs",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
assert exc.value.status_code == 404
assert exc.value.detail == {"error": "callback_name = gcs is not registered for team_id = team-1."}
mock_prisma.db.litellm_teamtable.update.assert_not_called()
@pytest.mark.asyncio
async def test_delete_team_callback_404s_when_team_has_no_logging_slot():
"""A team on the deprecated callback_settings shape holds no logging entries."""
metadata = {
"callback_settings": {
"success_callback": ["langfuse"],
"failure_callback": [],
"callback_vars": {"langfuse_public_key": "pk-demo"},
}
}
mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=metadata))
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.master_key", None),
):
with pytest.raises(HTTPException) as exc:
await delete_team_callback(
http_request=MagicMock(spec=Request),
team_id="team-1",
callback_name="langfuse",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
assert exc.value.status_code == 404
mock_prisma.db.litellm_teamtable.update.assert_not_called()
@pytest.mark.asyncio
async def test_delete_team_callback_404s_for_unknown_team():
mock_prisma = MagicMock()
mock_prisma.get_data = AsyncMock(return_value=None)
mock_prisma.db.litellm_teamtable.update = AsyncMock()
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
with pytest.raises(HTTPException) as exc:
await delete_team_callback(
http_request=MagicMock(spec=Request),
team_id="team-missing",
callback_name="langfuse",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
assert exc.value.status_code == 404
mock_prisma.db.litellm_teamtable.update.assert_not_called()
@pytest.mark.asyncio
async def test_delete_team_callback_keeps_last_removal_from_reviving_legacy_shape():
"""Removing the last entry must leave metadata["logging"] present and empty.
Request-time resolution selects the logging branch on key presence, so
dropping the key would fall through to a legacy callback_settings block and
silently re-enable a destination the caller just removed.
"""
from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata
metadata = {
"logging": [
{
"callback_name": "langsmith",
"callback_type": "success",
"callback_vars": {"langsmith_project": "demo"},
}
],
"callback_settings": {
"success_callback": ["langfuse"],
"failure_callback": [],
"callback_vars": {"langfuse_public_key": "pk-legacy"},
},
}
mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=metadata))
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.master_key", None),
):
response = await delete_team_callback(
http_request=MagicMock(spec=Request),
team_id="team-1",
callback_name="langsmith",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"])
assert written["logging"] == []
assert response.data.success_callbacks == ()
resolved = _get_dynamic_logging_metadata(
UserAPIKeyAuth(api_key="hashed", team_id="team-1", team_metadata=written),
proxy_config=MagicMock(**{"load_team_config.return_value": {}}),
)
assert not (resolved.success_callback if resolved else None)
@pytest.mark.asyncio
async def test_delete_team_callback_refreshes_cached_team(stub_team_cache_refresh):
"""The DB write alone leaves the removed callback firing.
Auth serves a cached team object and request-time callback resolution reads
the metadata off it, so a key already in flight keeps sending to the removed
destination until the cache entry expires.
"""
mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=_two_callback_metadata()))
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.master_key", None),
):
await delete_team_callback(
http_request=MagicMock(spec=Request),
team_id="team-1",
callback_name="langsmith",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
stub_team_cache_refresh.assert_awaited_once()
refreshed = stub_team_cache_refresh.await_args.kwargs["team_row"]
assert refreshed is mock_prisma.db.litellm_teamtable.update.return_value
# The row fed to the cache has to carry object_permission, or the refresh
# publishes a team whose tool allowlists look empty, which reads as
# unrestricted on the search-tool and MCP-tool checks.
update_kwargs = mock_prisma.db.litellm_teamtable.update.await_args.kwargs
assert update_kwargs["include"]["object_permission"] is True
@pytest.mark.asyncio
async def test_delete_team_callback_emits_redacted_audit_log(monkeypatch):
"""The audit row records the removal without becoming a credential sink."""
monkeypatch.setattr(litellm, "store_audit_logs", True)
mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=_two_callback_metadata()))
audit_calls = []
async def capture(request_data):
audit_calls.append(request_data)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
patch("litellm.proxy.proxy_server.master_key", None),
patch(
"litellm.proxy.management_helpers.audit_logs.create_audit_log_for_update",
new=capture,
),
):
await delete_team_callback(
http_request=MagicMock(spec=Request),
team_id="team-1",
callback_name="langsmith",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
import asyncio
for _ in range(3):
await asyncio.sleep(0)
assert len(audit_calls) == 1
log = audit_calls[0]
assert log.table_name == LitellmTableNames.TEAM_TABLE_NAME
assert log.object_id == "team-1"
assert log.action == "updated"
before = json.loads(log.before_value)
after = json.loads(log.updated_values)
assert [entry["callback_name"] for entry in before["metadata"]["logging"]] == [
"langsmith",
"langfuse",
]
assert [entry["callback_name"] for entry in after["metadata"]["logging"]] == ["langfuse"]
assert "ls-demo" not in log.before_value
assert "sk-demo" not in log.updated_values
@pytest.mark.asyncio
async def test_delete_team_callback_encrypts_surviving_callback_vars(monkeypatch):
"""The write must not downgrade the survivors' stored credentials to plaintext."""
from litellm.proxy.common_utils.callback_utils import decrypt_callback_vars
monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt-32-bytes-aaaaaaaaaaaaaa")
mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=_two_callback_metadata()))
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.master_key", None),
):
await delete_team_callback(
http_request=MagicMock(spec=Request),
team_id="team-1",
callback_name="langsmith",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"])
stored = written["logging"][0]["callback_vars"]
assert stored["langfuse_secret_key"] != "sk-demo"
assert decrypt_callback_vars(written)["logging"][0]["callback_vars"]["langfuse_secret_key"] == "sk-demo"
@pytest.mark.asyncio
async def test_delete_team_callback_keeps_entries_it_cannot_parse():
"""A malformed entry is left alone rather than crashing the removal.
metadata["logging"] is free-form JSON that /team/update will persist as given,
so the filter has to tolerate an entry that is not a callback dict.
"""
metadata = {
"logging": [
"not-a-callback-entry",
{
"callback_name": "langsmith",
"callback_type": "success",
"callback_vars": {"langsmith_project": "demo"},
},
]
}
mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=metadata))
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.master_key", None),
):
await delete_team_callback(
http_request=MagicMock(spec=Request),
team_id="team-1",
callback_name="langsmith",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"])
assert written["logging"] == ["not-a-callback-entry"]
@pytest.mark.asyncio
async def test_delete_team_callback_route_accepts_team_ids_containing_slashes():
"""The route has to reach the same team ids POST and GET /team/{team_id}/callback do.
Those siblings declare team_id with the path converter, so a team registered under an
id with a slash can add and list callbacks. Without the same converter here the delete
404s at the routing layer for exactly those teams.
"""
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.team_callback_endpoints import router
team_id = "tenant/eu-west"
metadata = {
"logging": [
{
"callback_name": "langsmith",
"callback_type": "success",
"callback_vars": {"langsmith_project": "demo"},
},
{
"callback_name": "langfuse",
"callback_type": "success",
"callback_vars": {"langfuse_public_key": "pk-demo"},
},
]
}
mock_prisma = _patch_prisma(_team_row(team_id=team_id, metadata=metadata))
app = FastAPI()
app.include_router(router)
app.dependency_overrides[user_api_key_auth] = _admin_auth
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.master_key", None),
):
response = TestClient(app).delete(f"/team/{team_id}/callback/langfuse")
assert response.status_code == 200
assert response.json()["data"]["success_callbacks"] == ["langsmith"]
written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"])
assert [entry["callback_name"] for entry in written["logging"]] == ["langsmith"]

View file

@ -41,6 +41,10 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
import litellm
MESSAGE_START_SSE_FRAME = b'event: message_start\ndata: {"type": "message_start"}\n\n'
# Test is_multipart
def test_is_multipart():
@ -5104,3 +5108,185 @@ async def test_passthrough_body_cannot_forge_budget_reservation():
increment_spend_counters.assert_awaited_once()
assert increment_spend_counters.await_args.kwargs["budget_reservation"] is None
async def _drive_streaming_pass_through(
upstream_content_type, chunk_delay_seconds, client_asked_for_stream=True
):
"""Drive pass_through_request against an upstream that stalls before its first byte.
``client_asked_for_stream`` picks which of pass_through_request's two streaming
dispatch branches runs: the up-front one, and the one that only discovers the
response is a stream from its content-type.
"""
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
PassThroughStreamingHandler,
)
with ExitStack() as stack:
mock_proxy_logging = stack.enter_context(
patch("litellm.proxy.proxy_server.proxy_logging_obj")
)
mock_get_client = stack.enter_context(
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
)
)
mock_chunk_processor = stack.enter_context(
patch.object(PassThroughStreamingHandler, "chunk_processor")
)
mock_proxy_logging.pre_call_hook = AsyncMock(
return_value={"model": "claude-3", "stream": True}
if client_asked_for_stream
else {"model": "claude-3"}
)
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
upstream_response = MagicMock()
upstream_response.status_code = 200
upstream_response.headers = {"content-type": upstream_content_type}
upstream_response.raise_for_status = MagicMock()
async_client = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
async def _slow_first_chunk(*args, **kwargs):
await asyncio.sleep(chunk_delay_seconds)
yield MESSAGE_START_SSE_FRAME
mock_chunk_processor.return_value = _slow_first_chunk()
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/v1/messages"
mock_request.body = AsyncMock(
return_value=b'{"model": "claude-3", "stream": true}'
if client_asked_for_stream
else b'{"model": "claude-3"}'
)
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
response = await pass_through_request(
request=mock_request,
target="http://target-api.com/v1/messages",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=client_asked_for_stream,
)
return [chunk async for chunk in response.body_iterator]
@pytest.mark.asyncio
@pytest.mark.parametrize("client_asked_for_stream", [True, False])
async def test_pass_through_sse_stream_emits_keepalive_before_the_first_upstream_byte(
client_asked_for_stream,
):
"""
Regression for #34819: a passthrough SSE stream wrote zero bytes during the
model's time-to-first-token, so an intermediary with an idle read timeout
(ALB, nginx) dropped a healthy connection before any token arrived.
Both dispatch branches are covered: a request that declared stream=true, and
one whose response is only recognised as a stream from its content-type.
"""
with patch.object(litellm, "sse_keepalive_ping_interval_seconds", 0.05):
collected = await _drive_streaming_pass_through(
upstream_content_type="text/event-stream",
chunk_delay_seconds=0.2,
client_asked_for_stream=client_asked_for_stream,
)
assert collected[0] == b": ping\n\n"
assert collected[-1] == MESSAGE_START_SSE_FRAME
@pytest.mark.asyncio
async def test_pass_through_sse_stream_stays_silent_when_keepalive_is_unconfigured():
with patch.object(litellm, "sse_keepalive_ping_interval_seconds", None):
collected = await _drive_streaming_pass_through(
upstream_content_type="text/event-stream", chunk_delay_seconds=0.2
)
assert collected == [MESSAGE_START_SSE_FRAME]
@pytest.mark.asyncio
async def test_pass_through_binary_event_stream_is_never_given_an_sse_comment():
"""An AWS event stream is a binary transport: a ": ping" frame would corrupt it."""
with patch.object(litellm, "sse_keepalive_ping_interval_seconds", 0.05):
collected = await _drive_streaming_pass_through(
upstream_content_type="application/vnd.amazon.eventstream",
chunk_delay_seconds=0.2,
)
assert collected == [MESSAGE_START_SSE_FRAME]
@pytest.mark.asyncio
@pytest.mark.parametrize("configured_interval, expect_ping", [(0.05, True), (None, False)])
async def test_pass_through_route_pings_while_the_upstream_call_is_still_running(
configured_interval, expect_ping
):
"""The upstream withholds its response headers until its first token, so the
whole time-to-first-token is spent inside pass_through_request with nothing on
the wire (issue #34819)."""
from fastapi import Response
from fastapi.responses import StreamingResponse
module = "litellm.proxy.pass_through_endpoints.pass_through_endpoints"
async def _relayed():
yield MESSAGE_START_SSE_FRAME
async def slow_pass_through(**kwargs):
await asyncio.sleep(0.25)
return StreamingResponse(_relayed(), media_type="text/event-stream")
with ExitStack() as stack:
stack.enter_context(
patch(
f"{module}.InitPassThroughEndpointHelpers.is_registered_pass_through_route",
return_value=True,
)
)
stack.enter_context(
patch(
f"{module}.InitPassThroughEndpointHelpers.get_registered_pass_through_route",
return_value=None,
)
)
stack.enter_context(patch(f"{module}.pass_through_request", slow_pass_through))
stack.enter_context(
patch.object(litellm, "sse_keepalive_ping_interval_seconds", configured_interval)
)
endpoint_func = create_pass_through_route(
endpoint="/v1/messages",
target="https://api.anthropic.com/v1/messages",
custom_headers={},
is_streaming_request=True,
)
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = httpx.URL("http://test-proxy.com/v1/messages")
mock_request.scope = {}
mock_request.body = AsyncMock(return_value=b'{"model": "claude-3"}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
mock_request.state = SimpleNamespace()
response = await endpoint_func(
request=mock_request,
fastapi_response=Response(),
user_api_key_dict=MagicMock(),
)
collected = [chunk async for chunk in response.body_iterator]
assert (collected[0] == b": ping\n\n") is expect_ping
assert collected[-1] in (MESSAGE_START_SSE_FRAME, MESSAGE_START_SSE_FRAME.decode())

View file

@ -188,6 +188,77 @@ def test_resolve_complexity_router_plugins_rejects_synchronous_run_method(tmp_pa
)
def test_resolve_complexity_router_plugins_resolves_classifier_plugin_dotted_path(tmp_path):
plugin_file = tmp_path / "my_classifier.py"
plugin_file.write_text(
"class _Classifier:\n"
" async def classify(self, context):\n"
" return 'SIMPLE'\n"
"\n"
"my_classifier_instance = _Classifier()\n"
)
config: dict[str, Any] = {
"classifier_type": "custom",
"classifier_plugin": "my_classifier.my_classifier_instance",
}
resolve_complexity_router_plugins(
model_name="smart-router",
complexity_router_config=config,
config_file_path=str(tmp_path / "config.yaml"),
)
assert hasattr(config["classifier_plugin"], "classify")
assert type(config["classifier_plugin"]).__name__ == "_Classifier"
def test_resolve_complexity_router_plugins_rejects_non_classifier_object(tmp_path):
plugin_file = tmp_path / "bad_classifier.py"
plugin_file.write_text("not_a_classifier = object()\n")
config: dict[str, Any] = {"classifier_plugin": "bad_classifier.not_a_classifier"}
with pytest.raises(ValueError, match="does not implement the ClassifierPlugin interface"):
resolve_complexity_router_plugins(
model_name="smart-router",
complexity_router_config=config,
config_file_path=str(tmp_path / "config.yaml"),
)
def test_resolve_complexity_router_plugins_rejects_synchronous_classify_method(tmp_path):
"""A synchronous `classify` passes the runtime_checkable isinstance and would only fail on
the first classified request, so reject it at config load like the sync-run case above."""
plugin_file = tmp_path / "sync_classifier.py"
plugin_file.write_text(
"class _SyncClassifier:\n"
" def classify(self, context):\n"
" return 'SIMPLE'\n"
"\n"
"sync_classifier_instance = _SyncClassifier()\n"
)
config: dict[str, Any] = {"classifier_plugin": "sync_classifier.sync_classifier_instance"}
with pytest.raises(ValueError, match="does not implement the ClassifierPlugin interface"):
resolve_complexity_router_plugins(
model_name="smart-router",
complexity_router_config=config,
config_file_path=str(tmp_path / "config.yaml"),
)
def test_resolve_complexity_router_plugins_leaves_live_classifier_instance_alone():
class _Classifier:
async def classify(self, context):
return "SIMPLE"
instance = _Classifier()
config: dict[str, Any] = {"classifier_plugin": instance}
resolve_complexity_router_plugins(
model_name="smart-router", complexity_router_config=config, config_file_path=None
)
assert config["classifier_plugin"] is instance
# ---------------------------------------------------------------------------
# resolve_routing_plugins
# ---------------------------------------------------------------------------

View file

@ -28,6 +28,9 @@ from litellm.proxy.common_request_processing import (
_get_cost_breakdown_from_logging_obj,
_has_attribute_error_in_chain,
_is_azure_model_router_request,
_UpstreamClosingStreamingResponse,
open_sse_before_first_byte,
ttft_keepalive_interval,
_override_openai_response_model,
_parse_event_data_for_error,
_resolve_per_request_model_group_alias,
@ -4511,6 +4514,61 @@ class TestAllmPassthroughStreamingProviderGate:
assert streamed == chunks
mock_handler.assert_not_awaited()
@pytest.mark.asyncio
async def test_bedrock_invoke_stream_sets_event_stream_content_type(self, monkeypatch):
"""
Regression for LIT-4561. The unbuffered Bedrock event-stream relay
(invoke-with-response-stream, no post-call guardrail rewriting) must set
content-type: application/vnd.amazon.eventstream instead of emitting no
content-type header at all, which trips Claude Code's content-type guard
added in 2.1.208
"""
processing_obj = self._build_processing_obj(
"bedrock", "model/us.anthropic.claude-sonnet-4-20250514-v1:0/invoke-with-response-stream"
)
chunks = [b"raw-1", b"raw-2"]
with patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails",
return_value=False,
), patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails_for_passthrough",
return_value=False,
):
result = await self._run(processing_obj, monkeypatch, chunks)
assert isinstance(result, StreamingResponse)
assert result.media_type == "application/vnd.amazon.eventstream"
assert result.headers["content-type"] == "application/vnd.amazon.eventstream"
streamed = [chunk async for chunk in result.body_iterator]
assert streamed == chunks
@pytest.mark.asyncio
async def test_non_bedrock_stream_keeps_default_content_type(self, monkeypatch):
"""
A provider with no registered event-stream media type must not have one
forced onto its unbuffered stream, so the response default is unchanged
"""
processing_obj = self._build_processing_obj("anthropic")
chunks = [b"chunk-1", b"chunk-2"]
with patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails",
return_value=False,
), patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails_for_passthrough",
return_value=False,
):
result = await self._run(processing_obj, monkeypatch, chunks)
assert isinstance(result, StreamingResponse)
assert result.media_type is None
assert "content-type" not in result.headers
class TestResponseCostHeaderForTypedDictResponses:
"""
@ -6057,3 +6115,525 @@ class TestProcessChunkWithCostInjection:
)
assert ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(chunk, "gpt-4o-mini") == chunk
# ---------------------------------------------------------------------------
# SSE keepalive during the time-to-first-token (issue #34819)
# ---------------------------------------------------------------------------
TTFT_PING = b": ping\n\n"
async def _drain(response):
return [chunk async for chunk in response.body_iterator]
def _sse_response(chunks, upstream_generator=None):
async def gen():
for chunk in chunks:
yield chunk
if upstream_generator is None:
return StreamingResponse(gen(), media_type="text/event-stream")
return _UpstreamClosingStreamingResponse(
gen(),
media_type="text/event-stream",
upstream_generator=upstream_generator,
)
@pytest.mark.asyncio
async def test_ttft_keepalive_fills_the_wire_while_the_upstream_is_still_silent():
"""Regression for #34819. The upstream withholds its headers until the first
token, so the whole wait happens before a byte can be written and an
idle-timeout hop drops a healthy connection."""
async def slow_upstream():
await asyncio.sleep(0.35)
return _sse_response(['data: {"first": true}\n\n'])
response = await open_sse_before_first_byte(slow_upstream(), ping_interval_seconds=0.05)
assert isinstance(response, StreamingResponse)
assert response.headers["x-accel-buffering"] == "no"
collected = await _drain(response)
assert collected[0] == TTFT_PING
assert collected.count(TTFT_PING) >= 3
assert collected[-1] == b'data: {"first": true}\n\n'
@pytest.mark.asyncio
async def test_ttft_keepalive_is_a_no_op_when_the_upstream_answers_in_time():
produced = _sse_response(['data: {"fast": true}\n\n'])
async def fast_upstream():
return produced
response = await open_sse_before_first_byte(fast_upstream(), ping_interval_seconds=5.0)
assert response is produced
assert await _drain(response) == ['data: {"fast": true}\n\n']
@pytest.mark.asyncio
@pytest.mark.parametrize("interval", [None, 0, "", "abc", float("inf"), float("nan"), -1])
async def test_ttft_keepalive_unconfigured_leaves_the_call_completely_untouched(interval):
produced = _sse_response(['data: {"x": 1}\n\n'])
started_at = asyncio.get_running_loop().time()
async def slow_upstream():
await asyncio.sleep(0.15)
return produced
response = await open_sse_before_first_byte(slow_upstream(), ping_interval_seconds=interval)
assert response is produced
assert asyncio.get_running_loop().time() - started_at >= 0.15
@pytest.mark.asyncio
async def test_ttft_keepalive_reraises_a_fast_failure_so_it_keeps_its_http_status():
async def fast_failure():
raise HTTPException(status_code=429, detail="rate limited")
with pytest.raises(HTTPException) as excinfo:
await open_sse_before_first_byte(fast_failure(), ping_interval_seconds=5.0)
assert excinfo.value.status_code == 429
@pytest.mark.asyncio
async def test_ttft_keepalive_delivers_a_late_failure_as_an_sse_frame():
"""Once a ping is on the wire the status line is committed, so a failure
discovered afterwards can only reach the client as a frame."""
async def slow_failure():
await asyncio.sleep(0.2)
raise HTTPException(status_code=429, detail="rate limited")
response = await open_sse_before_first_byte(slow_failure(), ping_interval_seconds=0.05)
collected = await _drain(response)
assert collected[0] == TTFT_PING
assert collected[-1] == b"data: [DONE]\n\n"
error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip())
assert error_frame["error"]["code"] == "429"
assert error_frame["error"]["message"] == "rate limited"
@pytest.mark.asyncio
async def test_ttft_keepalive_relays_a_late_non_streaming_body_as_an_sse_frame():
async def slow_json():
await asyncio.sleep(0.2)
return JSONResponse(status_code=400, content={"error": {"message": "bad request"}})
response = await open_sse_before_first_byte(slow_json(), ping_interval_seconds=0.05)
collected = await _drain(response)
assert collected[0] == TTFT_PING
assert json.loads(collected[-2].decode().removeprefix("data: ").strip()) == {"error": {"message": "bad request"}}
assert collected[-1] == b"data: [DONE]\n\n"
@pytest.mark.asyncio
async def test_ttft_keepalive_closes_the_upstream_stream_it_relayed():
"""Starlette never calls the produced response, so its own cleanup never runs
and the upstream LLM connection would leak."""
upstream_closed = asyncio.Event()
async def upstream():
try:
yield 'data: {"a": 1}\n\n'
finally:
upstream_closed.set()
upstream_gen = upstream()
# Started, as create_response leaves it: aclose() on a never-started generator
# skips its body, so an unstarted fixture cannot tell cleanup from no cleanup.
await upstream_gen.__anext__()
async def slow_upstream():
await asyncio.sleep(0.2)
return _sse_response(['data: {"a": 1}\n\n'], upstream_generator=upstream_gen)
response = await open_sse_before_first_byte(slow_upstream(), ping_interval_seconds=0.05)
await _drain(response)
assert upstream_closed.is_set()
@pytest.mark.asyncio
async def test_ttft_keepalive_cancels_the_in_flight_call_when_the_client_gives_up():
upstream_cancelled = asyncio.Event()
async def never_answers():
try:
await asyncio.Event().wait()
except asyncio.CancelledError:
upstream_cancelled.set()
raise
response = await open_sse_before_first_byte(never_answers(), ping_interval_seconds=0.05)
assert await response.body_iterator.__anext__() == TTFT_PING
await response.body_iterator.aclose()
await asyncio.sleep(0)
assert upstream_cancelled.is_set()
@pytest.mark.parametrize(
"request_data, global_interval, expected",
[
({"stream": True}, 30.0, 30.0),
({"stream": True}, None, None),
({"stream": False}, 30.0, None),
({}, 30.0, None),
({"stream": "true"}, 30.0, None),
],
)
def test_ttft_keepalive_interval_only_arms_for_a_streaming_request(request_data, global_interval, expected):
with patch.object(litellm, "sse_keepalive_ping_interval_seconds", global_interval):
assert ttft_keepalive_interval(request_data) == expected
@pytest.mark.asyncio
@pytest.mark.parametrize("stream_requested, expect_ping", [(True, True), (False, False)])
async def test_base_process_llm_request_pings_while_the_upstream_call_is_still_running(
stream_requested, expect_ping
):
"""The wiring, not the helper: every route funnels through this method, and the
whole time-to-first-token is spent inside the call it wraps."""
async def slow_inner(self, **kwargs):
await asyncio.sleep(0.25)
return _sse_response(['data: {"late": true}\n\n'])
processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-4o", "stream": stream_requested})
with patch.object(litellm, "sse_keepalive_ping_interval_seconds", 0.05):
with patch.object(ProxyBaseLLMRequestProcessing, "_process_llm_request", slow_inner):
response = await processor.base_process_llm_request(
request=MagicMock(spec=Request),
fastapi_response=Response(),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
route_type="acompletion",
proxy_logging_obj=MagicMock(spec=ProxyLogging),
general_settings={},
proxy_config=MagicMock(spec=ProxyConfig),
)
collected = await _drain(response)
assert (collected[0] == TTFT_PING) is expect_ping
assert collected[-1] == (b'data: {"late": true}\n\n' if expect_ping else 'data: {"late": true}\n\n')
def _request_disconnecting_after(delay_seconds):
"""A Request whose ASGI channel delivers one http.disconnect, then goes quiet."""
request = MagicMock(spec=Request)
delivered = {"done": False}
async def receive():
if delivered["done"]:
await asyncio.Event().wait()
await asyncio.sleep(delay_seconds)
delivered["done"] = True
return {"type": "http.disconnect"}
request.receive = receive
return request
@pytest.mark.asyncio
@pytest.mark.parametrize(
"disconnect_after, expect_full_delivery",
[(0.25, False), (999.0, True)],
)
async def test_opening_the_response_early_still_closes_the_upstream_on_disconnect(
disconnect_after, expect_full_delivery
):
"""Once the response is opened early, create_response's own disconnect
monitoring runs while Starlette is already serving, so both read the same ASGI
channel. Whichever observes the disconnect, the upstream LLM stream must close.
"""
upstream_closed = asyncio.Event()
delivered = []
async def upstream():
try:
await asyncio.sleep(0.4)
for chunk in ('data: {"a": 1}\n\n', "data: [DONE]\n\n"):
delivered.append(chunk)
yield chunk
finally:
upstream_closed.set()
request = _request_disconnecting_after(disconnect_after)
async def produce():
await asyncio.sleep(0.15)
return await create_response(
generator=upstream(),
media_type="text/event-stream",
headers={},
request=request,
)
response = await open_sse_before_first_byte(produce(), ping_interval_seconds=0.05)
collected = await _drain(response)
await asyncio.sleep(0.05)
assert collected[0] == TTFT_PING
assert upstream_closed.is_set()
# The control has to actually deliver, or "the upstream closed" proves nothing.
assert (delivered == ['data: {"a": 1}\n\n', "data: [DONE]\n\n"]) is expect_full_delivery
@pytest.mark.asyncio
async def test_a_disconnect_after_the_upstream_answered_still_closes_the_response():
"""The upstream can answer while nobody is draining the relay, e.g. the client
vanished first. Nothing else holds that response, so only this teardown closes
it; cancelling the produce task is not enough because it already finished."""
upstream_closed = asyncio.Event()
body_closed = asyncio.Event()
async def upstream():
try:
yield 'data: {"a": 1}\n\n'
await asyncio.Event().wait()
finally:
upstream_closed.set()
async def body():
try:
yield 'data: {"a": 1}\n\n'
await asyncio.Event().wait()
finally:
body_closed.set()
# Both started, as create_response leaves them: aclose() on a never-started
# generator skips its body, so an unstarted fixture cannot tell cleanup apart
# from no cleanup at all.
upstream_gen, body_gen = upstream(), body()
await upstream_gen.__anext__()
await body_gen.__anext__()
async def produce():
await asyncio.sleep(0.15)
return _UpstreamClosingStreamingResponse(
body_gen, media_type="text/event-stream", upstream_generator=upstream_gen
)
response = await open_sse_before_first_byte(produce(), ping_interval_seconds=0.05)
assert await response.body_iterator.__anext__() == TTFT_PING
await asyncio.sleep(0.25) # the produce task finishes while nothing is pulling
await response.body_iterator.aclose()
await asyncio.sleep(0.05)
assert body_closed.is_set()
assert upstream_closed.is_set()
@pytest.mark.asyncio
async def test_a_late_failure_is_reported_to_the_failure_hook():
"""Once a keepalive is on the wire this can no longer raise, so the caller's
own `except` never runs and the failure would otherwise go unaudited."""
audited = []
async def slow_failure():
await asyncio.sleep(0.2)
raise HTTPException(status_code=500, detail="upstream exploded")
async def record(exc):
audited.append(exc)
response = await open_sse_before_first_byte(
slow_failure(), ping_interval_seconds=0.05, on_late_failure=record
)
collected = await _drain(response)
assert [type(exc).__name__ for exc in audited] == ["HTTPException"]
assert getattr(audited[0], "detail", None) == "upstream exploded"
assert collected[-1] == b"data: [DONE]\n\n"
@pytest.mark.asyncio
async def test_a_failing_audit_hook_never_costs_the_client_its_error_frame():
async def slow_failure():
await asyncio.sleep(0.2)
raise HTTPException(status_code=500, detail="upstream exploded")
async def broken_hook(exc):
raise RuntimeError("the audit backend is down")
response = await open_sse_before_first_byte(
slow_failure(), ping_interval_seconds=0.05, on_late_failure=broken_hook
)
collected = await _drain(response)
error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip())
assert error_frame["error"]["message"] == "upstream exploded"
assert collected[-1] == b"data: [DONE]\n\n"
@pytest.mark.asyncio
async def test_base_process_llm_request_audits_a_failure_that_lands_after_its_keepalive():
"""The helper honouring on_late_failure is not enough: this pins that the shared
funnel actually passes one, which is where the route's own except would have
fired before the response was opened early."""
async def slow_failure(self, **kwargs):
await asyncio.sleep(0.25)
raise HTTPException(status_code=503, detail="upstream exploded")
proxy_logging_obj = MagicMock(spec=ProxyLogging)
# None is what a hook that only audits returns; a bare AsyncMock would hand
# back a MagicMock, which the code correctly reads as a sanitized replacement.
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-4o", "stream": True})
with patch.object(litellm, "sse_keepalive_ping_interval_seconds", 0.05):
with patch.object(ProxyBaseLLMRequestProcessing, "_process_llm_request", slow_failure):
response = await processor.base_process_llm_request(
request=MagicMock(spec=Request),
fastapi_response=Response(),
user_api_key_dict=user_api_key_dict,
route_type="acompletion",
proxy_logging_obj=proxy_logging_obj,
general_settings={},
proxy_config=MagicMock(spec=ProxyConfig),
)
collected = await _drain(response)
proxy_logging_obj.post_call_failure_hook.assert_awaited_once()
call = proxy_logging_obj.post_call_failure_hook.await_args.kwargs
assert call["user_api_key_dict"] is user_api_key_dict
assert call["request_data"] is processor.data
assert getattr(call["original_exception"], "detail", None) == "upstream exploded"
assert collected[0] == TTFT_PING
assert collected[-1] == b"data: [DONE]\n\n"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"deployment_keepalive, expect_ping",
[(0, False), (None, True)],
ids=["operator-hard-disabled-this-deployment", "deployment-says-nothing"],
)
async def test_base_process_llm_request_honours_a_deployment_hard_disable(
deployment_keepalive, expect_ping
):
"""`keepalive_seconds: 0` is documented as a disable a request cannot lift. The
funnel has to hand its router to the gate for that to hold before the upstream
has answered, since no deployment has served the request yet."""
params = {"model": "openai/gpt-4o"}
if deployment_keepalive is not None:
params["keepalive_seconds"] = deployment_keepalive
llm_router = MagicMock()
llm_router.get_model_list = MagicMock(return_value=[{"model_name": "m", "litellm_params": params}])
async def slow_inner(self, **kwargs):
await asyncio.sleep(0.25)
return _sse_response(['data: {"late": true}\n\n'])
processor = ProxyBaseLLMRequestProcessing(data={"model": "m", "stream": True})
with patch.object(litellm, "sse_keepalive_ping_interval_seconds", 0.05):
with patch.object(ProxyBaseLLMRequestProcessing, "_process_llm_request", slow_inner):
response = await processor.base_process_llm_request(
request=MagicMock(spec=Request),
fastapi_response=Response(),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
route_type="acompletion",
proxy_logging_obj=MagicMock(spec=ProxyLogging),
general_settings={},
proxy_config=MagicMock(spec=ProxyConfig),
llm_router=llm_router,
)
collected = await _drain(response)
assert (collected[0] == TTFT_PING) is expect_ping
@pytest.mark.asyncio
async def test_a_hook_returning_a_replacement_decides_what_the_client_sees():
"""post_call_failure_hook exists partly to sanitize client-facing errors.
Serializing the original would leak provider detail a deployment configured away."""
async def slow_failure():
await asyncio.sleep(0.2)
raise HTTPException(status_code=500, detail="upstream said host=10.0.0.7 key=sk-internal")
async def sanitize(exc):
return HTTPException(status_code=502, detail="upstream unavailable")
response = await open_sse_before_first_byte(
slow_failure(), ping_interval_seconds=0.05, on_late_failure=sanitize
)
collected = await _drain(response)
error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip())
assert error_frame["error"]["message"] == "upstream unavailable"
assert "sk-internal" not in collected[-2].decode()
@pytest.mark.asyncio
async def test_a_hook_raising_a_replacement_also_decides_what_the_client_sees():
"""The hook's contract is return *or* raise, and raising is the path a
suppress(Exception) around the call would silently discard."""
async def slow_failure():
await asyncio.sleep(0.2)
raise HTTPException(status_code=500, detail="upstream said host=10.0.0.7 key=sk-internal")
async def sanitize_by_raising(exc):
raise HTTPException(status_code=403, detail="blocked by policy")
response = await open_sse_before_first_byte(
slow_failure(), ping_interval_seconds=0.05, on_late_failure=sanitize_by_raising
)
collected = await _drain(response)
error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip())
assert error_frame["error"]["message"] == "blocked by policy"
assert "sk-internal" not in collected[-2].decode()
@pytest.mark.asyncio
async def test_a_hook_that_returns_nothing_leaves_the_real_error_intact():
async def slow_failure():
await asyncio.sleep(0.2)
raise HTTPException(status_code=429, detail="rate limited")
async def audit_only(exc):
return None
response = await open_sse_before_first_byte(
slow_failure(), ping_interval_seconds=0.05, on_late_failure=audit_only
)
collected = await _drain(response)
error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip())
assert error_frame["error"]["message"] == "rate limited"
assert error_frame["error"]["code"] == "429"
@pytest.mark.asyncio
async def test_a_broken_hook_does_not_replace_the_real_error_with_its_own_bug():
async def slow_failure():
await asyncio.sleep(0.2)
raise HTTPException(status_code=429, detail="rate limited")
async def broken_hook(exc):
raise RuntimeError("the audit backend is down")
response = await open_sse_before_first_byte(
slow_failure(), ping_interval_seconds=0.05, on_late_failure=broken_hook
)
collected = await _drain(response)
error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip())
assert error_frame["error"]["message"] == "rate limited"
assert "audit backend" not in collected[-2].decode()

View file

@ -1037,6 +1037,65 @@ async def test_add_litellm_data_to_request_ignores_forged_client_side_timeout():
assert not updated.get("client_side_timeout")
@pytest.mark.asyncio
async def test_client_side_timeout_marker_never_reaches_the_provider():
"""A proxy request with a caller-supplied timeout gets kwargs["client_side_timeout"]
stamped for the router's cooldown logic. That router-only marker must not ride
into the provider payload: unregistered kwargs are swept into extra_body /
additionalModelRequestFields, so Bedrock rejects the whole call with
`client_side_timeout: Extra inputs are not permitted`."""
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
updated = await add_litellm_data_to_request(
data={
"model": "bedrock/us.anthropic.claude-sonnet-5",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 10,
"timeout": 30,
},
request=request_mock,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated["client_side_timeout"] is True
converse_response = MagicMock()
converse_response.status_code = 200
converse_response.headers = {}
converse_response.json.return_value = {
"output": {"message": {"role": "assistant", "content": [{"text": "ok"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15},
}
converse_response.text = json.dumps(converse_response.json.return_value)
client = AsyncHTTPHandler()
with patch.object(client, "post", return_value=converse_response) as mock_post:
await litellm.acompletion(
**updated,
aws_access_key_id="fake-access-key",
aws_secret_access_key="fake-secret-key",
aws_region_name="us-east-1",
client=client,
)
mock_post.assert_called_once()
assert mock_post.call_args.kwargs["url"].endswith("/converse")
provider_body = json.loads(mock_post.call_args.kwargs["data"])
assert "client_side_timeout" not in json.dumps(provider_body), provider_body
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_allows_client_mock_response_with_admin_opt_in():
request_mock = MagicMock(spec=Request)

View file

@ -102,6 +102,30 @@ class TestStripClientPricingOverrides:
assert data["metadata"] == {"user_session": "keep-me"}
assert data["litellm_metadata"] == {}
def test_metadata_guardrail_information_dropped(self):
# Client-seeded guardrail entries would otherwise be summed into
# response_cost and spend, letting a caller forge (even negative)
# guardrail cost against their own budget.
data = {
"model": "gpt-4",
"metadata": {
"user_session": "keep-me",
"standard_logging_guardrail_information": [
{
"guardrail_name": "forged",
"guardrail_status": "success",
"guardrail_cost": -0.005,
}
],
},
"litellm_metadata": {
"standard_logging_guardrail_information": [{"guardrail_cost": 5.0}],
},
}
_strip_client_pricing_overrides(data)
assert data["metadata"] == {"user_session": "keep-me"}
assert data["litellm_metadata"] == {}
def test_non_pricing_fields_untouched(self):
data = {
"model": "gpt-4",
@ -129,6 +153,7 @@ class TestStripClientPricingOverrides:
def test_metadata_field_set_contains_model_info(self):
assert "model_info" in _CLIENT_PRICING_METADATA_FIELDS
assert "standard_logging_guardrail_information" in _CLIENT_PRICING_METADATA_FIELDS
def test_strip_emits_debug_log_listing_dropped_fields(self, caplog):
# Operators need a paper trail so they can diagnose why a previously

View file

@ -154,7 +154,7 @@ class TestDeleteDeploymentResilience:
# Router has a model ID that's not in DB or config -> should be deleted
mock_router.get_model_ids.return_value = ["db-id-1", "stale-id"]
mock_router.delete_deployment.return_value = True
mock_router._generate_model_id = MagicMock(return_value="config-id-1")
mock_router.generate_model_id = MagicMock(return_value="config-id-1")
with (
patch.object(
@ -182,3 +182,111 @@ class TestDeleteDeploymentResilience:
"the returned set must be what the db + config still want, so a caller can "
f"tell that eviction apart from a deployment that went missing; got {result}"
)
class TestDeleteDeploymentKeepsPluginConfigModels:
"""Regression: _delete_deployment re-reads the raw config and hashes litellm_params to
compute the ids the config wants served. The Router used to derive plugin-bearing
deployment ids from the RESOLVED params (dotted paths swapped for live instances), so
the reconcile computed different ids and evicted every plugin-bearing auto-router one
sync after startup. load_config now pins model_info.id from the raw params before
resolution, so both sides hash the same input and the reconcile needs no resolution."""
@staticmethod
def _write_plugin_module(tmp_path):
(tmp_path / "rig_classifier.py").write_text(
"class _Classifier:\n"
" async def classify(self, context):\n"
" return 'SIMPLE'\n"
"\n"
"class _Narrower:\n"
" async def run(self, context):\n"
" return context\n"
"\n"
"classifier_instance = _Classifier()\n"
"narrower_instance = _Narrower()\n"
)
@staticmethod
def _raw_model_entry():
return {
"model_name": "smart-router",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_default_model": "gpt-4o-mini",
"complexity_router_config": {
"classifier_type": "custom",
"classifier_plugin": "rig_classifier.classifier_instance",
"plugins": ["rig_classifier.narrower_instance"],
"tiers": {"SIMPLE": "gpt-4o-mini"},
},
},
}
@pytest.mark.asyncio
async def test_plugin_bearing_config_model_survives_reconcile_and_stale_ids_still_evict(self, tmp_path):
import copy
from litellm import Router
from litellm.proxy.proxy_server import (
pin_complexity_router_model_id,
resolve_complexity_router_plugins,
)
self._write_plugin_module(tmp_path)
config_file_path = str(tmp_path / "config.yaml")
resolved_entry = copy.deepcopy(self._raw_model_entry())
pin_complexity_router_model_id(resolved_entry)
resolve_complexity_router_plugins(
model_name="smart-router",
complexity_router_config=resolved_entry["litellm_params"]["complexity_router_config"],
config_file_path=config_file_path,
)
router = Router(
model_list=[
{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}},
resolved_entry,
{
"model_name": "stale-model",
"litellm_params": {"model": "gpt-4o"},
"model_info": {"id": "stale-id"},
},
]
)
assert "smart-router" in router.model_names
assert "stale-model" in router.model_names
raw_config = {
"model_list": [
{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}},
self._raw_model_entry(),
]
}
proxy_config = ProxyConfig()
with (
patch.object(proxy_config, "get_config", new_callable=AsyncMock, return_value=raw_config),
patch("litellm.proxy.proxy_server.llm_router", router),
patch("litellm.proxy.proxy_server.user_config_file_path", config_file_path),
patch("litellm.proxy.proxy_server.premium_user", False),
):
result = await proxy_config._delete_deployment(db_models=[])
assert result is not None
assert "smart-router" in router.model_names
assert "stale-model" not in router.model_names
def test_pin_respects_an_explicit_model_id(self):
from litellm.proxy.proxy_server import pin_complexity_router_model_id
entry = self._raw_model_entry()
entry["model_info"] = {"id": "operator-pinned"}
pin_complexity_router_model_id(entry)
assert entry["model_info"]["id"] == "operator-pinned"
def test_pin_is_a_noop_without_a_complexity_router_config(self):
from litellm.proxy.proxy_server import pin_complexity_router_model_id
entry = {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}
pin_complexity_router_model_id(entry)
assert "model_info" not in entry

View file

@ -12,6 +12,7 @@ import httpx
import pytest
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
def _expected_dir() -> Path:
@ -367,3 +368,57 @@ async def test_aresponses_client_header_conflict_is_case_insensitive():
assert [name for name in request_headers if name.lower() == "x-shared"] == ["x-shared"]
assert request_headers["x-shared"] == "from-caller"
@pytest.mark.asyncio
@pytest.mark.parametrize(
("model", "custom_llm_provider"),
[
("openai/responses/gpt-5.6", None),
("responses/gpt-5.6", "openai"),
],
)
async def test_aresponses_strips_responses_routing_prefix_from_openai_model(model, custom_llm_provider):
"""
`responses/` is LiteLLM routing sugar, never part of the provider model id.
Deployments configured as openai/responses/<model> reach this path directly via
/v1/responses and via the /v1/messages adapter (which passes responses/<model>
with custom_llm_provider="openai"), so both shapes must hit OpenAI as <model>.
"""
injected_client = AsyncHTTPHandler()
mock_post = AsyncMock(return_value=MockResponse(_minimal_responses_api_payload("resp_prefix_test", "gpt-5.6"), 200))
injected_client.post = mock_post
await litellm.aresponses(
model=model,
custom_llm_provider=custom_llm_provider,
input="ping",
api_key="sk-test",
client=injected_client,
)
mock_post.assert_called_once()
assert mock_post.call_args.kwargs["url"].endswith("/responses")
assert mock_post.call_args.kwargs["json"]["model"] == "gpt-5.6"
@pytest.mark.asyncio
async def test_aresponses_websocket_strips_responses_routing_prefix_from_openai_model():
from unittest.mock import MagicMock
from litellm.responses.main import _aresponses_websocket
with patch(
"litellm.responses.main.base_llm_http_handler.async_responses_websocket",
new_callable=AsyncMock,
) as mock_ws:
await _aresponses_websocket(
model="openai/responses/gpt-5.6",
websocket=MagicMock(),
api_key="sk-test",
litellm_logging_obj=MagicMock(),
)
mock_ws.assert_awaited_once()
assert mock_ws.call_args.kwargs["model"] == "gpt-5.6"
assert mock_ws.call_args.kwargs["custom_llm_provider"] == "openai"

View file

@ -479,3 +479,85 @@ class TestAutoRouterEmbeddingInputCap:
assert auto_router.routelayer is not None
assert auto_router.routelayer.encoder.max_input_chars == 777
class TestAutoRouterRoutesResponsesApiInput:
"""Responses API requests carry the prompt in `input`, not `messages`, and still have to reach the route layer."""
@pytest.mark.asyncio
async def test_should_route_a_string_input_when_messages_is_none(self):
from semantic_router.schema import RouteChoice
layer: Final = FixedRouteLayer(RouteChoice(name="code-model"))
auto_router: Final = _auto_router(layer)
result: Final = await auto_router.async_pre_routing_hook(
model="my-auto-router",
request_kwargs={
"input": "fix this stack trace",
"litellm_metadata": {"user_api_key_request_route": "/v1/responses"},
},
messages=None,
)
assert result is not None
assert result.model == "code-model"
assert result.messages is None
assert layer.seen_text == "fix this stack trace"
@pytest.mark.asyncio
async def test_should_route_a_list_input_with_instructions_when_messages_is_none(self):
from semantic_router.schema import RouteChoice
layer: Final = FixedRouteLayer(RouteChoice(name="code-model"))
auto_router: Final = _auto_router(layer)
result: Final = await auto_router.async_pre_routing_hook(
model="my-auto-router",
request_kwargs={
"instructions": "You are a coding agent.",
"input": [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "fix this stack trace"}],
}
],
"litellm_metadata": {"user_api_key_request_route": "/v1/responses"},
},
messages=None,
)
assert result is not None
assert result.model == "code-model"
assert layer.seen_text is not None
assert "fix this stack trace" in layer.seen_text
@pytest.mark.asyncio
async def test_should_skip_routing_when_neither_messages_nor_input_is_present(self):
layer: Final = FixedRouteLayer(None)
auto_router: Final = _auto_router(layer)
result: Final = await auto_router.async_pre_routing_hook(
model="my-auto-router",
request_kwargs={"litellm_metadata": {"user_api_key_request_route": "/v1/responses"}},
messages=None,
)
assert result is None
assert layer.seen_text is None
@pytest.mark.asyncio
async def test_should_keep_routing_an_empty_messages_list_to_the_default_model(self):
layer: Final = FixedRouteLayer(None)
auto_router: Final = _auto_router(layer)
result: Final = await auto_router.async_pre_routing_hook(
model="my-auto-router",
request_kwargs={"messages": [], "litellm_metadata": {"user_api_key_request_route": "/v1/chat/completions"}},
messages=[],
)
assert result is not None
assert result.model == "fallback-model"
assert layer.seen_text == ""

View file

@ -4127,6 +4127,315 @@ class TestRoutingPlugins:
assert spy.call_count == 2
class _FixedTierClassifier:
"""Classifier plugin double returning a fixed verdict; records the context it received."""
def __init__(self, verdict):
self.verdict = verdict
self.seen_context = None
async def classify(self, context):
self.seen_context = context
return self.verdict
class _TeamTierClassifier:
async def classify(self, context):
team = context.metadata.get("user_api_key_team_id")
return "REASONING" if team == "team-premium" else "SIMPLE"
class _RaisingClassifier:
async def classify(self, context):
raise RuntimeError("lookup service down")
class _SlowClassifier:
async def classify(self, context):
await asyncio.sleep(5)
return "SIMPLE"
def _plugin_router(mock_router_instance, plugin, **config_overrides):
config = {
"tiers": {
"SIMPLE": "gpt-4o-mini",
"MEDIUM": "gpt-4o",
"COMPLEX": "claude-sonnet-4-20250514",
"REASONING": "o1-preview",
},
"classifier_type": "custom",
"classifier_plugin": plugin,
**config_overrides,
}
return ComplexityRouter(
model_name="test-complexity-router",
litellm_router_instance=mock_router_instance,
complexity_router_config=config,
)
class TestClassifierPluginConfig:
"""Config validation for classifier_type='custom'."""
def test_plugin_classifier_type_requires_plugin(self):
with pytest.raises(ValidationError, match="classifier_plugin is required"):
ComplexityRouterConfig(classifier_type="custom")
def test_classifier_plugin_without_plugin_mode_raises(self):
"""A wired hook that would silently never run is a config error, not a no-op."""
with pytest.raises(ValidationError, match="would never run"):
ComplexityRouterConfig(classifier_plugin=_FixedTierClassifier("SIMPLE"))
def test_plugin_mode_tolerates_stale_llm_config(self):
"""Switching classifier_type llm -> plugin must not force deleting classifier_llm_config,
matching how classifier_type='heuristic' tolerates it."""
config = ComplexityRouterConfig(
classifier_type="custom",
classifier_plugin=_FixedTierClassifier("SIMPLE"),
classifier_llm_config={"model": "haiku-classifier"},
)
assert config.classifier_type == "custom"
def test_plugin_mode_composes_with_adaptive(self):
"""adaptive replaces selection, not classification, so a classifier plugin is allowed
where narrowing `plugins` are rejected (their pools bypass the bandit)."""
config = ComplexityRouterConfig(
classifier_type="custom",
classifier_plugin=_FixedTierClassifier("SIMPLE"),
adaptive=True,
)
assert config.adaptive is True
def test_plugin_mode_composes_with_tier_definitions(self):
config = ComplexityRouterConfig(
classifier_type="custom",
classifier_plugin=_FixedTierClassifier("cheap"),
tiers={"cheap": "gpt-4o-mini", "premium": "o1-preview"},
tier_definitions=[
{"name": "cheap", "description": "routine asks"},
{"name": "premium", "description": "hard asks"},
],
fallback_tier="cheap",
)
assert config.tier_names() == ("cheap", "premium")
def test_tier_definitions_still_reject_heuristic(self):
with pytest.raises(ValidationError, match="heuristic scorer only"):
ComplexityRouterConfig(
classifier_type="heuristic",
tiers={"cheap": "gpt-4o-mini", "premium": "o1-preview"},
tier_definitions=[
{"name": "cheap", "description": "routine asks"},
{"name": "premium", "description": "hard asks"},
],
fallback_tier="cheap",
)
class TestClassifierPlugin:
"""classifier_type='custom': an operator hook decides the tier."""
@pytest.mark.asyncio
async def test_plugin_verdict_decides_tier_without_scorer_or_llm(self, mock_router_instance):
mock_router_instance.acompletion = AsyncMock()
router = _plugin_router(mock_router_instance, _FixedTierClassifier("COMPLEX"))
outcome = await router.aclassify("hello")
assert outcome.cause == "classifier_plugin"
assert outcome.tier == ComplexityTier.COMPLEX
assert outcome.score is None
assert outcome.signals == ("classifier-plugin:COMPLEX",)
mock_router_instance.acompletion.assert_not_called()
@pytest.mark.asyncio
async def test_plugin_verdict_resolves_case_insensitively(self, mock_router_instance):
router = _plugin_router(mock_router_instance, _FixedTierClassifier("reasoning"))
outcome = await router.aclassify("hello")
assert outcome.tier == ComplexityTier.REASONING
assert outcome.cause == "classifier_plugin"
@pytest.mark.asyncio
async def test_plugin_reads_caller_identity_from_request_metadata(self, mock_router_instance):
router = _plugin_router(mock_router_instance, _TeamTierClassifier())
premium = await router.aclassify("hi", request_kwargs={"metadata": {"user_api_key_team_id": "team-premium"}})
basic = await router.aclassify(
"hi", request_kwargs={"litellm_metadata": {"user_api_key_team_id": "team-basic"}}
)
assert premium.tier == ComplexityTier.REASONING
assert basic.tier == ComplexityTier.SIMPLE
@pytest.mark.asyncio
async def test_plugin_context_carries_messages_and_all_tier_models(self, mock_router_instance):
plugin = _FixedTierClassifier("SIMPLE")
router = _plugin_router(mock_router_instance, plugin)
raw = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]
await router.aclassify("hi", messages=[{"role": "user", "content": "hi"}], raw_messages=raw)
assert plugin.seen_context.raw_messages == raw
assert plugin.seen_context.structured_messages == raw
assert plugin.seen_context.candidate_models == [
"gpt-4o-mini",
"gpt-4o",
"claude-sonnet-4-20250514",
"o1-preview",
]
@pytest.mark.asyncio
async def test_plugin_runs_without_messages(self, mock_router_instance):
"""A prompt-only call (no message list) still reaches the plugin with an empty context."""
plugin = _FixedTierClassifier("COMPLEX")
router = _plugin_router(mock_router_instance, plugin)
outcome = await router.aclassify("hello", raw_messages=None)
assert outcome.cause == "classifier_plugin"
assert plugin.seen_context.raw_messages == []
assert plugin.seen_context.structured_messages == []
@pytest.mark.asyncio
async def test_plugin_decline_falls_back_to_heuristic(self, mock_router_instance):
router = _plugin_router(mock_router_instance, _FixedTierClassifier(None))
outcome = await router.aclassify("what is 2+2?")
assert outcome.cause == "heuristic_scorer"
@pytest.mark.asyncio
async def test_plugin_error_falls_back_to_heuristic(self, mock_router_instance):
router = _plugin_router(mock_router_instance, _RaisingClassifier())
outcome = await router.aclassify("what is 2+2?")
assert outcome.cause == "heuristic_scorer"
@pytest.mark.asyncio
async def test_plugin_timeout_falls_back_to_heuristic(self, mock_router_instance):
router = _plugin_router(mock_router_instance, _SlowClassifier(), classifier_plugin_timeout_ms=20)
outcome = await router.aclassify("what is 2+2?")
assert outcome.cause == "heuristic_scorer"
@pytest.mark.asyncio
async def test_plugin_non_string_verdict_falls_back_to_heuristic(self, mock_router_instance):
"""An operator hook returning a non-string must fall back, not raise into the request."""
router = _plugin_router(mock_router_instance, _FixedTierClassifier(42))
outcome = await router.aclassify("what is 2+2?")
assert outcome.cause == "heuristic_scorer"
@pytest.mark.asyncio
async def test_plugin_unknown_tier_falls_back_to_heuristic(self, mock_router_instance):
router = _plugin_router(mock_router_instance, _FixedTierClassifier("galactic"))
outcome = await router.aclassify("what is 2+2?")
assert outcome.cause == "heuristic_scorer"
@pytest.mark.asyncio
async def test_plugin_tier_without_pool_falls_back(self, mock_router_instance):
"""A built-in tier the operator gave no models is a decline, not a later routing error."""
router = ComplexityRouter(
model_name="test-complexity-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {"SIMPLE": "gpt-4o-mini"},
"classifier_type": "custom",
"classifier_plugin": _FixedTierClassifier("COMPLEX"),
},
)
outcome = await router.aclassify("what is 2+2?")
assert outcome.cause == "heuristic_scorer"
@pytest.mark.asyncio
async def test_plugin_failure_with_default_model_fallback(self, mock_router_instance):
router = _plugin_router(
mock_router_instance,
_RaisingClassifier(),
classifier_fallback="default_model",
default_model="gpt-4o-mini",
)
outcome = await router.aclassify("hello")
assert outcome.cause == "default_model_fallback"
@pytest.mark.asyncio
async def test_plugin_with_custom_tiers_routes_defined_name(self, mock_router_instance):
router = _plugin_router(
mock_router_instance,
_FixedTierClassifier("premium"),
tiers={"cheap": "gpt-4o-mini", "premium": "o1-preview"},
tier_definitions=[
{"name": "cheap", "description": "routine asks"},
{"name": "premium", "description": "hard asks"},
],
fallback_tier="cheap",
)
outcome = await router.aclassify("hello")
assert outcome.tier == "premium"
assert outcome.cause == "classifier_plugin"
assert outcome.signals == ("classifier-plugin:premium",)
@pytest.mark.asyncio
async def test_plugin_failure_with_custom_tiers_routes_fallback_tier(self, mock_router_instance):
router = _plugin_router(
mock_router_instance,
_RaisingClassifier(),
tiers={"cheap": "gpt-4o-mini", "premium": "o1-preview"},
tier_definitions=[
{"name": "cheap", "description": "routine asks"},
{"name": "premium", "description": "hard asks"},
],
fallback_tier="cheap",
)
outcome = await router.aclassify("hello")
assert outcome.tier == "cheap"
assert outcome.cause == "classifier_fallback"
assert outcome.signals == ("classifier-fallback:cheap",)
@pytest.mark.asyncio
async def test_hook_records_plugin_cause_without_score(self, mock_router_instance):
router = _plugin_router(mock_router_instance, _TeamTierClassifier())
response = await router.async_pre_routing_hook(
model="test-complexity-router",
request_kwargs={"metadata": {"user_api_key_team_id": "team-premium"}},
messages=[{"role": "user", "content": "prove P != NP"}],
)
decision = response.routing_decision
assert decision["cause"] == "classifier_plugin"
assert decision["tier"] == "REASONING"
assert decision["routed_model"] == "o1-preview"
assert response.model == "o1-preview"
assert "score" not in decision
assert "tier_boundaries" not in decision
@pytest.mark.asyncio
async def test_plugin_composes_with_narrowing_plugins(self, mock_router_instance):
class _BlockO1:
async def run(self, context):
context.candidate_models = [m for m in context.candidate_models if m != "o1-preview"]
return context
router = _plugin_router(
mock_router_instance,
_FixedTierClassifier("REASONING"),
tiers={
"SIMPLE": "gpt-4o-mini",
"MEDIUM": "gpt-4o",
"COMPLEX": "claude-sonnet-4-20250514",
"REASONING": ["o1-preview", "claude-sonnet-4-20250514"],
},
plugins=[_BlockO1()],
)
response = await router.async_pre_routing_hook(
model="test-complexity-router",
request_kwargs={},
messages=[{"role": "user", "content": "prove P != NP"}],
)
assert response.model == "claude-sonnet-4-20250514"
assert response.routing_decision["cause"] == "classifier_plugin"
def test_classifier_plugin_alone_keeps_tier_pinning_enabled(self, mock_router_instance):
"""Narrowing plugins suppress session pinning (a policy verdict can change between turns);
a classifier plugin picks among operator-approved tiers, so pinning must stay on."""
pinning = _plugin_router(mock_router_instance, _FixedTierClassifier("SIMPLE"), session_affinity=True)
suppressed = _plugin_router(
mock_router_instance,
_FixedTierClassifier("SIMPLE"),
session_affinity=True,
plugins=[_DummyPlugin()],
)
assert pinning._uses_tier_pin is True
assert suppressed._uses_tier_pin is False
class TestEscalationKeywords:
"""Test user-triggered escalation: a keyword in the prompt bumps the resolved tier
one step higher so a user can force a stronger model when unhappy with results."""

View file

@ -221,7 +221,7 @@ def test_filter_by_routing_plugin_candidates_narrows_and_raises_when_empty():
def test_json_default_stable_id_is_stable_across_instances():
"""_generate_model_id's json.dumps `default=` fallback must not embed an object's
"""generate_model_id's json.dumps `default=` fallback must not embed an object's
memory address (e.g. plain str() on an object with no custom __repr__ falls back
to object.__repr__'s `<module.Class object at 0x...>`) -- that would make the
deployment id churn on every process restart for any deployment whose
@ -232,7 +232,7 @@ def test_json_default_stable_id_is_stable_across_instances():
assert router._json_default_stable_id(LanguageDetector()) != router._json_default_stable_id(TenantPolicy())
def test_generate_model_id_is_stable_when_litellm_params_contain_a_plugin_instance():
def testgenerate_model_id_is_stable_when_litellm_params_contain_a_plugin_instance():
"""End-to-end: a deployment id built from litellm_params containing a routing
plugin instance (e.g. complexity_router_config.plugins) must be identical across
separate calls, not just non-crashing."""
@ -242,8 +242,8 @@ def test_generate_model_id_is_stable_when_litellm_params_contain_a_plugin_instan
"complexity_router_config": {"plugins": [LanguageDetector()]},
}
id1 = router._generate_model_id("smart-router", litellm_params)
id2 = router._generate_model_id(
id1 = router.generate_model_id("smart-router", litellm_params)
id2 = router.generate_model_id(
"smart-router",
{
"model": "auto_router/complexity_router",

View file

@ -2389,6 +2389,114 @@ def test_stream_chunk_builder_text_completion_combines_text_and_usage():
assert response.usage.total_tokens == response.usage.prompt_tokens + response.usage.completion_tokens
def test_completion_forwards_store_and_prompt_cache_key_to_openai():
"""
Regression test for https://github.com/BerriAI/litellm/issues/33184
store and prompt_cache_key are documented OpenAI chat completion params that
were accepted as supported but silently dropped before the provider request
was built, because they were not named parameters of completion() and
get_optional_params() the way safety_identifier is.
"""
from openai import OpenAI
client = OpenAI(api_key="fake-api-key")
with patch.object(client.chat.completions.with_raw_response, "create") as mock_client:
try:
litellm.completion(
model="openai/gpt-4o",
messages=[{"role": "user", "content": "Hello"}],
store=False,
prompt_cache_key="test-cache-key",
client=client,
)
except Exception as e:
print(e)
mock_client.assert_called_once()
request_body = mock_client.call_args.kwargs
assert request_body["store"] is False
assert request_body["prompt_cache_key"] == "test-cache-key"
@pytest.mark.asyncio
async def test_acompletion_forwards_store_and_prompt_cache_key_to_openai():
"""
Async variant of the store/prompt_cache_key forwarding regression test for
https://github.com/BerriAI/litellm/issues/33184
"""
from openai import AsyncOpenAI
client = AsyncOpenAI(api_key="fake-api-key")
with patch.object(client.chat.completions.with_raw_response, "create") as mock_client:
try:
await litellm.acompletion(
model="openai/gpt-4o",
messages=[{"role": "user", "content": "Hello"}],
store=False,
prompt_cache_key="test-cache-key",
client=client,
)
except Exception as e:
print(e)
mock_client.assert_called_once()
request_body = mock_client.call_args.kwargs
assert request_body["store"] is False
assert request_body["prompt_cache_key"] == "test-cache-key"
def test_completion_omits_store_and_prompt_cache_key_when_not_passed():
"""
When store and prompt_cache_key are not passed, they must not appear in the
outbound request body (guards against always forwarding None defaults).
"""
from openai import OpenAI
client = OpenAI(api_key="fake-api-key")
with patch.object(client.chat.completions.with_raw_response, "create") as mock_client:
try:
litellm.completion(
model="openai/gpt-4o",
messages=[{"role": "user", "content": "Hello"}],
client=client,
)
except Exception as e:
print(e)
mock_client.assert_called_once()
request_body = mock_client.call_args.kwargs
assert "store" not in request_body
assert "prompt_cache_key" not in request_body
def test_completion_forwards_store_and_prompt_cache_key_to_mcp_gateway():
"""
Regression test for the MCP gateway early-return in completion(): store and
prompt_cache_key are named params, so they no longer travel via **kwargs and
must be forwarded explicitly like safety_identifier and service_tier.
"""
with patch(
"litellm.responses.mcp.chat_completions_handler.acompletion_with_mcp"
) as mock_mcp:
result = litellm.completion(
model="openai/gpt-4o",
messages=[{"role": "user", "content": "Hello"}],
tools=[{"type": "mcp", "server_url": "litellm_proxy"}],
store=False,
prompt_cache_key="test-cache-key",
)
result.close()
mock_mcp.assert_called_once()
call_kwargs = mock_mcp.call_args.kwargs
assert call_kwargs["store"] is False
assert call_kwargs["prompt_cache_key"] == "test-cache-key"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"aws_credential_kwargs",

View file

@ -524,6 +524,126 @@ async def test_async_router_acreate_file_uses_deployment_custom_llm_provider():
assert mock_acreate_file.call_args.kwargs["custom_llm_provider"] == "azure"
@pytest.mark.asyncio
async def test_async_router_acreate_file_forwards_target_model_names_to_litellm_proxy():
import json
from io import BytesIO
from unittest.mock import MagicMock, patch
jsonl_file = BytesIO(
json.dumps({"body": {"model": "chained-batch", "messages": [{"role": "user", "content": "hi"}]}}).encode(
"utf-8"
)
)
jsonl_file.name = "test.jsonl"
router = litellm.Router(
model_list=[
{
"model_name": "chained-batch",
"litellm_params": {
"model": "litellm_proxy/gpt-4.1-batch",
"api_base": "http://localhost:4001/v1",
"api_key": "sk-proxy-b",
},
},
],
)
with patch("litellm.acreate_file", return_value=MagicMock()) as mock_acreate_file:
await router.acreate_file(
model="chained-batch",
purpose="batch",
file=jsonl_file,
)
assert mock_acreate_file.call_count == 1
call_kwargs = mock_acreate_file.call_args.kwargs
assert call_kwargs["custom_llm_provider"] == "litellm_proxy"
assert call_kwargs["extra_body"] == {"target_model_names": "gpt-4.1-batch"}
uploaded_line = json.loads(call_kwargs["file"].read().decode("utf-8").split("\n")[0])
assert uploaded_line["body"]["model"] == "gpt-4.1-batch"
@pytest.mark.asyncio
async def test_async_router_acreate_file_does_not_inject_target_model_names_for_other_providers():
from unittest.mock import MagicMock, patch
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4.1-batch",
"litellm_params": {"model": "gpt-4.1"},
},
],
)
with patch("litellm.acreate_file", return_value=MagicMock()) as mock_acreate_file:
await router.acreate_file(
model="gpt-4.1-batch",
purpose="batch",
file=MagicMock(),
)
assert mock_acreate_file.call_count == 1
assert mock_acreate_file.call_args.kwargs.get("extra_body") is None
@pytest.mark.asyncio
async def test_async_router_acreate_file_litellm_proxy_sends_target_model_names_in_multipart_form():
import json
from io import BytesIO
import httpx
import respx
jsonl_file = BytesIO(
json.dumps({"body": {"model": "chained-batch", "messages": [{"role": "user", "content": "hi"}]}}).encode(
"utf-8"
)
)
jsonl_file.name = "test.jsonl"
router = litellm.Router(
model_list=[
{
"model_name": "chained-batch",
"litellm_params": {
"model": "litellm_proxy/gpt-4.1-batch",
"api_base": "http://localhost:4001/v1",
"api_key": "sk-proxy-b",
},
},
],
)
file_object_json = {
"id": "file-abc123",
"object": "file",
"bytes": 100,
"created_at": 1700000000,
"filename": "test.jsonl",
"purpose": "batch",
"status": "processed",
}
with respx.mock(assert_all_called=True) as respx_mock:
create_route = respx_mock.post("http://localhost:4001/v1/files").mock(
return_value=httpx.Response(200, json=file_object_json)
)
response = await router.acreate_file(
model="chained-batch",
purpose="batch",
file=jsonl_file,
)
assert response.id == "file-abc123"
request_body = create_route.calls.last.request.content
assert b'name="target_model_names"' in request_body
assert b"gpt-4.1-batch" in request_body
assert b'name="purpose"' in request_body
@pytest.mark.asyncio
async def test_async_router_afile_content_uses_deployment_custom_llm_provider():
"""

View file

@ -869,6 +869,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"container",
"image_edit",
"embedding",
"guardrail",
"image_generation",
"video_generation",
"moderation",
@ -976,6 +977,10 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"type": "string",
},
},
"guardrail_cost_per_unit": {
"type": "object",
"additionalProperties": {"type": "number"},
},
"search_context_cost_per_query": {
"type": "object",
"properties": {
@ -4797,6 +4802,24 @@ def test_bedrock_batch_params_never_reach_the_provider():
)
def test_client_side_timeout_marker_never_reaches_the_provider():
"""The proxy stamps kwargs["client_side_timeout"] = True whenever a request carries
a caller-supplied timeout (body timeout / request_timeout / stream_timeout or the
x-litellm-timeout headers) so the router can skip cooldowns on the resulting 408s.
The marker is only meaningful to the router, so it must be filtered out of the
provider params: swept into extra_body / additionalModelRequestFields it turns every
timed-out request into a provider 400 (`client_side_timeout: Extra inputs are not
permitted`)."""
kwargs = {"a_real_provider_specific_param": 1, "client_side_timeout": True}
non_default = get_non_default_completion_params(kwargs)
assert non_default == {"a_real_provider_specific_param": 1}, (
"client_side_timeout leaked into the provider params: "
f"{sorted(set(non_default) - {'a_real_provider_specific_param'})}"
)
def test_rust_flag_not_forwarded_as_provider_param():
forwarded = get_non_default_completion_params({"rust": True, "temperature": 0.5})
assert "rust" not in forwarded

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 22895
"limit": 22894
},
"LIT002": {
"limit": 26889
"limit": 26888
},
"LIT003": {
"limit": 269
@ -27,7 +27,7 @@
"limit": 0
},
"LIT010": {
"limit": 16661
"limit": 16700
},
"LIT011": {
"limit": 5590

View file

@ -16,7 +16,7 @@
},
"src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx": {
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
@ -161,9 +161,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 2
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -401,12 +398,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-nested-ternary": {
"count": 5
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -415,12 +406,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-nested-ternary": {
"count": 5
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -430,11 +415,6 @@
"count": 1
}
},
"src/app/(dashboard)/guardrails/_components/llm_judge/LLMJudgeFields.tsx": {
"no-restricted-imports": {
"count": 2
}
},
"src/app/(dashboard)/guardrails/_components/pii_components.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -623,7 +603,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 4
@ -669,7 +649,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
@ -685,7 +665,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 2
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/OpenAPIFormSection.tsx": {
@ -736,7 +716,7 @@
"count": 3
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
@ -760,7 +740,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/static-components": {
"count": 4
@ -803,7 +783,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/immutability": {
"count": 1
@ -976,7 +956,7 @@
},
"src/app/(dashboard)/playground/components/compareUI/components/ModelSelector.tsx": {
"no-restricted-imports": {
"count": 2
"count": 1
}
},
"src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx": {
@ -1346,9 +1326,6 @@
},
"src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx": {
"no-restricted-imports": {
"count": 2
},
"react-hooks/set-state-in-effect": {
"count": 1
}
},
@ -1357,7 +1334,7 @@
"count": 2
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/static-components": {
"count": 1
@ -1409,7 +1386,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 3
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
@ -1479,9 +1456,6 @@
"src/app/(dashboard)/users/_components/user_edit_view.test.tsx": {
"no-nested-ternary": {
"count": 1
},
"react/display-name": {
"count": 1
}
},
"src/app/(dashboard)/users/_components/user_edit_view.tsx": {
@ -1511,7 +1485,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
@ -1519,7 +1493,7 @@
},
"src/app/(dashboard)/vector-stores/_components/CreateVectorStore.tsx": {
"no-restricted-imports": {
"count": 3
"count": 2
}
},
"src/app/(dashboard)/vector-stores/_components/S3VectorsConfig.tsx": {
@ -1547,9 +1521,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -1676,9 +1647,6 @@
}
},
"src/components/SCIM.tsx": {
"no-restricted-imports": {
"count": 2
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -1690,7 +1658,7 @@
},
"src/components/SSOModals.tsx": {
"no-restricted-imports": {
"count": 2
"count": 1
}
},
"src/components/Settings/AdminSettings/HashicorpVault/EditHashicorpVaultModal.tsx": {
@ -1734,11 +1702,6 @@
"count": 1
}
},
"src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx": {
"no-restricted-imports": {
"count": 2
}
},
"src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.test.tsx": {
"max-nested-callbacks": {
"count": 1
@ -1789,7 +1752,7 @@
},
"src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx": {
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
@ -1816,7 +1779,7 @@
"count": 2
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"prefer-const": {
"count": 2
@ -1874,7 +1837,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 4
"count": 3
}
},
"src/components/add_model/ClassificationMethodConfig.tsx": {
@ -1925,7 +1888,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 4
"count": 3
},
"prefer-const": {
"count": 2
@ -1960,7 +1923,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 2
@ -1989,7 +1952,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 3
"count": 2
}
},
"src/components/add_model/model_connection_test.tsx": {
@ -2010,10 +1973,10 @@
"count": 1
},
"no-nested-ternary": {
"count": 5
"count": 3
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/immutability": {
"count": 3
@ -2053,10 +2016,7 @@
"count": 1
},
"no-nested-ternary": {
"count": 4
},
"no-restricted-imports": {
"count": 2
"count": 1
}
},
"src/components/bulk_create_users_button.tsx": {
@ -2120,7 +2080,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"no-restricted-syntax": {
"count": 3
@ -2131,7 +2091,7 @@
},
"src/components/common_components/AccessGroupSelector.tsx": {
"no-restricted-imports": {
"count": 2
"count": 1
}
},
"src/components/common_components/DeleteResourceModal.tsx": {
@ -2149,7 +2109,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 2
"count": 1
}
},
"src/components/common_components/MetadataKeyValueFields.test.tsx": {
@ -2380,7 +2340,7 @@
},
"src/components/mcp_tools/MCPToolArgumentsForm.tsx": {
"no-nested-ternary": {
"count": 5
"count": 1
},
"no-restricted-imports": {
"count": 1
@ -2398,7 +2358,7 @@
},
"src/components/model_add/CredentialModal.tsx": {
"no-restricted-imports": {
"count": 3
"count": 2
}
},
"src/components/model_add/reuse_credentials.tsx": {
@ -2490,7 +2450,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 2
"count": 1
}
},
"src/components/organisms/RegenerateKeyModal.tsx": {
@ -2500,7 +2460,7 @@
},
"src/components/organisms/create_key_button.test.tsx": {
"@typescript-eslint/no-require-imports": {
"count": 2
"count": 1
},
"react/display-name": {
"count": 8
@ -2517,7 +2477,7 @@
"count": 1
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"prefer-const": {
"count": 2
@ -2647,9 +2607,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 3
}
@ -2749,14 +2706,6 @@
"count": 1
}
},
"src/components/shared/usage_date_picker.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/tag_management/types.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -2772,7 +2721,7 @@
},
"src/components/team/LoggingSettings.tsx": {
"no-restricted-imports": {
"count": 2
"count": 1
}
},
"src/components/team/TeamInfo.tsx": {
@ -2783,7 +2732,7 @@
"count": 3
},
"no-restricted-imports": {
"count": 2
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
@ -2828,7 +2777,7 @@
"count": 2
},
"no-restricted-imports": {
"count": 2
"count": 1
}
},
"src/components/templates/key_info_view.tsx": {

View file

@ -0,0 +1,6 @@
<?xml version="1.0" encoding="UTF-8"?>
<svg width="64" height="73" viewBox="0 0 64 73" xmlns="http://www.w3.org/2000/svg">
<g id="Group-copy">
<path id="Path" fill="#123678" fill-rule="evenodd" stroke="none" d="M 13.482285 60.694962 L 0.998384 52.884399 L 0.998384 19.502914 L 31.527868 2.001205 L 61.317604 19.532024 L 61.317604 54.64489 L 31.054855 71.68927 L 20.548372 65.115807 L 20.548372 51.041328 L 20.548372 49.119896 L 14.851504 45.555508 L 14.851504 27.453159 L 31.346497 17.99712 L 47.464485 27.482262 L 47.464485 46.451157 L 34.703495 53.638138 L 34.703495 45.998573 C 38.52874 44.52552 41.274452 40.739189 41.274452 36.270489 C 41.274452 30.510658 36.712814 25.88438 31.158138 25.88438 C 25.603172 25.88438 21.041817 30.510658 21.041817 36.270489 C 21.041817 40.739189 23.787249 44.52552 27.612494 45.998573 L 27.612494 60.473576 L 31.261133 62.756348 L 53.635483 50.15464 L 53.635483 23.924595 L 31.477489 10.884869 L 8.680504 23.953705 L 8.680504 48.628967 L 13.482285 51.633297 L 13.482285 60.694962 Z M 31.158138 31.498383 C 33.671822 31.498383 35.660439 33.664162 35.660439 36.270489 C 35.660439 38.876804 33.671822 41.042587 31.158138 41.042587 C 28.644447 41.042587 26.655558 38.876804 26.655558 36.270489 C 26.655558 33.664162 28.644447 31.498383 31.158138 31.498383 Z" />
</g>
</svg>

After

Width:  |  Height:  |  Size: 1.3 KiB

View file

@ -1,147 +1,179 @@
"use client";
import { BotIcon, InfoIcon, LayersIcon, ServerIcon } from "lucide-react";
import type { UseFormReturn } from "react-hook-form";
import { z } from "zod/v4";
import { useAgents } from "@/app/(dashboard)/hooks/agents/useAgents";
import { useMCPServers } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers";
import { ModelSelect } from "@/components/ModelSelect/ModelSelect";
import type { FormInstance } from "antd";
import { Form, Input, Select, Space, Tabs } from "antd";
import { BotIcon, InfoIcon, LayersIcon, ServerIcon } from "lucide-react";
import { FieldGroup } from "@/components/shared/form/field";
import { FormField } from "@/components/shared/form/FormField";
import { Input } from "@/components/ui/input";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
import { Textarea } from "@/components/ui/textarea";
const { TextArea } = Input;
export const accessGroupFormSchema = z.object({
name: z.string().min(1, "Please enter the access group name"),
description: z.string(),
modelIds: z.array(z.string()),
mcpServerIds: z.array(z.string()),
agentIds: z.array(z.string()),
});
export interface AccessGroupFormValues {
name: string;
description: string;
modelIds: string[];
mcpServerIds: string[];
agentIds: string[];
export type AccessGroupFormValues = z.output<typeof accessGroupFormSchema>;
export const GENERAL_TAB = "general";
export const MODELS_TAB = "models";
export const MCP_SERVERS_TAB = "mcp-servers";
export const AGENTS_TAB = "agents";
interface MultiSelectOption {
value: string;
label: string;
}
interface MultiSelectProps {
id: string;
value: string[];
onChange: (value: string[]) => void;
options: MultiSelectOption[];
placeholder: string;
"aria-invalid": true | undefined;
"aria-describedby": string | undefined;
}
const MultiSelect = ({
id,
value,
onChange,
options,
placeholder,
"aria-invalid": ariaInvalid,
"aria-describedby": ariaDescribedBy,
}: MultiSelectProps) => (
<Select multiple items={options} value={value} onValueChange={onChange}>
<SelectTrigger id={id} aria-invalid={ariaInvalid} aria-describedby={ariaDescribedBy} className="w-full">
<SelectValue placeholder={placeholder}>
{(selected: string[]) =>
selected.length === 0
? placeholder
: options
.filter((option) => selected.includes(option.value))
.map((option) => option.label)
.join(", ")
}
</SelectValue>
</SelectTrigger>
<SelectContent>
{options.map((option) => (
<SelectItem key={option.value} value={option.value} title={option.label}>
{option.label}
</SelectItem>
))}
</SelectContent>
</Select>
);
interface AccessGroupBaseFormProps {
form: FormInstance<AccessGroupFormValues>;
form: UseFormReturn<AccessGroupFormValues>;
isNameDisabled?: boolean;
activeTab: string;
onTabChange: (tab: string) => void;
}
export function AccessGroupBaseForm({ form, isNameDisabled = false }: AccessGroupBaseFormProps) {
export function AccessGroupBaseForm({
form,
isNameDisabled = false,
activeTab,
onTabChange,
}: AccessGroupBaseFormProps) {
const { data: agentsData } = useAgents();
const { data: mcpServersData } = useMCPServers();
const agents = agentsData?.agents ?? [];
const mcpServers = mcpServersData ?? [];
const items = [
{
key: "1",
label: (
<Space align="center" size={4}>
<InfoIcon size={16} />
General Info
</Space>
),
children: (
<div style={{ paddingTop: 16 }}>
<Form.Item
name="name"
label="Group Name"
rules={[
{
required: true,
message: "Please enter the access group name",
},
]}
>
<Input placeholder="e.g. Engineering Team" disabled={isNameDisabled} />
</Form.Item>
<Form.Item name="description" label="Description">
<TextArea rows={4} placeholder="Describe the purpose of this access group..." />
</Form.Item>
</div>
),
},
{
key: "2",
label: (
<Space align="center" size={4}>
<LayersIcon size={16} />
Models
</Space>
),
children: (
<div style={{ paddingTop: 16 }}>
<Form.Item name="modelIds" label="Allowed Models">
<ModelSelect
context="global"
value={form.getFieldValue("modelIds") ?? []}
onChange={(values) => form.setFieldsValue({ modelIds: values })}
style={{ width: "100%" }}
/>
</Form.Item>
</div>
),
},
{
key: "3",
label: (
<Space align="center" size={4}>
<ServerIcon size={16} />
MCP Servers
</Space>
),
children: (
<div style={{ paddingTop: 16 }}>
<Form.Item name="mcpServerIds" label="Allowed MCP Servers">
<Select
mode="multiple"
placeholder="Select MCP servers"
style={{ width: "100%" }}
optionFilterProp="label"
allowClear
options={mcpServers.map((server) => ({
label: server.server_name ?? server.server_id,
value: server.server_id,
}))}
/>
</Form.Item>
</div>
),
},
{
key: "4",
label: (
<Space align="center" size={4}>
<BotIcon size={16} />
Agents
</Space>
),
children: (
<div style={{ paddingTop: 16 }}>
<Form.Item name="agentIds" label="Allowed Agents">
<Select
mode="multiple"
placeholder="Select agents"
style={{ width: "100%" }}
optionFilterProp="label"
allowClear
options={agents.map((agent) => ({
label: agent.agent_name,
value: agent.agent_id,
}))}
/>
</Form.Item>
</div>
),
},
];
const mcpServerOptions = (mcpServersData ?? []).map((server) => ({
value: server.server_id,
label: server.server_name ?? server.server_id,
}));
const agentOptions = (agentsData?.agents ?? []).map((agent) => ({
value: agent.agent_id,
label: agent.agent_name,
}));
return (
<Form
form={form}
layout="vertical"
name="access_group_form"
initialValues={{
modelIds: [],
mcpServerIds: [],
agentIds: [],
}}
>
<Tabs defaultActiveKey="1" items={items} />
</Form>
<Tabs value={activeTab} onValueChange={onTabChange}>
<TabsList className="w-full">
<TabsTrigger value={GENERAL_TAB}>
<InfoIcon size={16} />
General Info
</TabsTrigger>
<TabsTrigger value={MODELS_TAB}>
<LayersIcon size={16} />
Models
</TabsTrigger>
<TabsTrigger value={MCP_SERVERS_TAB}>
<ServerIcon size={16} />
MCP Servers
</TabsTrigger>
<TabsTrigger value={AGENTS_TAB}>
<BotIcon size={16} />
Agents
</TabsTrigger>
</TabsList>
<TabsContent value={GENERAL_TAB} className="pt-4">
<FieldGroup>
<FormField control={form.control} name="name" label="Group Name">
{({ ref, ...field }) => (
<Input {...field} ref={ref} placeholder="e.g. Engineering Team" disabled={isNameDisabled} />
)}
</FormField>
<FormField control={form.control} name="description" label="Description">
{({ ref, ...field }) => (
<Textarea {...field} ref={ref} rows={4} placeholder="Describe the purpose of this access group..." />
)}
</FormField>
</FieldGroup>
</TabsContent>
<TabsContent value={MODELS_TAB} className="pt-4">
<FormField control={form.control} name="modelIds" label="Allowed Models">
{(field) => <ModelSelect context="global" value={field.value} onChange={field.onChange} />}
</FormField>
</TabsContent>
<TabsContent value={MCP_SERVERS_TAB} className="pt-4">
<FormField control={form.control} name="mcpServerIds" label="Allowed MCP Servers">
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
<MultiSelect
id={id}
value={value}
onChange={onChange}
options={mcpServerOptions}
placeholder="Select MCP servers"
aria-invalid={ariaInvalid}
aria-describedby={ariaDescribedBy}
/>
)}
</FormField>
</TabsContent>
<TabsContent value={AGENTS_TAB} className="pt-4">
<FormField control={form.control} name="agentIds" label="Allowed Agents">
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
<MultiSelect
id={id}
value={value}
onChange={onChange}
options={agentOptions}
placeholder="Select agents"
aria-invalid={ariaInvalid}
aria-describedby={ariaDescribedBy}
/>
)}
</FormField>
</TabsContent>
</Tabs>
);
}

View file

@ -0,0 +1,175 @@
import { describe, it, expect, vi, beforeEach } from "vitest";
import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event";
import { renderWithProviders, screen, waitFor } from "../../../../../../tests/test-utils";
import { AccessGroupEditModal } from "./AccessGroupEditModal";
import { AccessGroupResponse } from "@/app/(dashboard)/hooks/accessGroups/useAccessGroups";
const mutate = vi.fn();
vi.mock("@/app/(dashboard)/hooks/accessGroups/useEditAccessGroup", () => ({
useEditAccessGroup: () => ({ mutate, isPending: false }),
}));
vi.mock("@/app/(dashboard)/hooks/agents/useAgents", () => ({
useAgents: () => ({ data: { agents: [{ agent_id: "agent-1", agent_name: "Support Bot" }] } }),
}));
vi.mock("@/app/(dashboard)/hooks/mcpServers/useMCPServers", () => ({
useMCPServers: () => ({ data: [{ server_id: "srv-1", server_name: "Files" }] }),
}));
vi.mock("@/components/ModelSelect/ModelSelect", () => ({
ModelSelect: ({ value, onChange }: { value: string[]; onChange: (next: string[]) => void }) => (
<button type="button" aria-label="model-select" onClick={() => onChange([...(value ?? []), "gpt-4"])}>
{(value ?? []).join(",")}
</button>
),
}));
vi.mock("@/lib/toast", () => ({
toast: { success: vi.fn(), fromError: vi.fn(), error: vi.fn() },
}));
const setup = () => userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
type User = ReturnType<typeof setup>;
const accessGroup: AccessGroupResponse = {
access_group_id: "ag-1",
access_group_name: "Engineering",
description: "Engineers",
access_model_names: ["gpt-4"],
access_mcp_server_ids: ["srv-1"],
access_agent_ids: ["agent-1"],
assigned_team_ids: [],
assigned_key_ids: [],
created_at: "2024-01-01T00:00:00Z",
created_by: "user-1",
updated_at: "2024-01-02T00:00:00Z",
updated_by: "user-1",
};
const renderModal = (data: AccessGroupResponse = accessGroup) =>
renderWithProviders(<AccessGroupEditModal visible accessGroup={data} onCancel={vi.fn()} />);
const save = async (user: User) => user.click(screen.getByRole("button", { name: "Save Changes" }));
const variables = () => mutate.mock.calls.at(-1)?.[0] as { accessGroupId: string; params: Record<string, unknown> };
describe("AccessGroupEditModal submit payload", () => {
beforeEach(() => {
vi.clearAllMocks();
});
it("sends exactly the antd payload for an untouched save", async () => {
const user = setup();
renderModal();
await screen.findByDisplayValue("Engineering");
await save(user);
await waitFor(() => expect(mutate).toHaveBeenCalled());
expect(mutate).toHaveBeenCalledTimes(1);
expect(variables().accessGroupId).toBe("ag-1");
expect(variables().params).toStrictEqual({
access_group_name: "Engineering",
description: "Engineers",
access_model_names: undefined,
access_mcp_server_ids: undefined,
access_agent_ids: undefined,
});
});
it("sends a tab's field only once that tab has been visited", async () => {
const user = setup();
renderModal();
await screen.findByDisplayValue("Engineering");
await user.click(screen.getByRole("tab", { name: /MCP Servers/ }));
await user.click(screen.getByRole("tab", { name: /Agents/ }));
await user.click(screen.getByRole("tab", { name: /General Info/ }));
await save(user);
await waitFor(() => expect(mutate).toHaveBeenCalled());
expect(variables().params).toStrictEqual({
access_group_name: "Engineering",
description: "Engineers",
access_model_names: undefined,
access_mcp_server_ids: ["srv-1"],
access_agent_ids: ["agent-1"],
});
});
it("coerces a null description to an empty string", async () => {
const user = setup();
renderModal({ ...accessGroup, description: null });
await screen.findByDisplayValue("Engineering");
await save(user);
await waitFor(() => expect(mutate).toHaveBeenCalled());
expect(variables().params.description).toBe("");
});
it("does not trim surrounding whitespace from the group name", async () => {
const user = setup();
renderModal();
const nameInput = await screen.findByDisplayValue("Engineering");
await user.clear(nameInput);
await user.type(nameInput, " Padded ");
await save(user);
await waitFor(() => expect(mutate).toHaveBeenCalled());
expect(variables().params.access_group_name).toBe(" Padded ");
});
it("never sends server-only fields from the loaded record", async () => {
const user = setup();
renderModal();
await screen.findByDisplayValue("Engineering");
await save(user);
await waitFor(() => expect(mutate).toHaveBeenCalled());
expect(variables().params).not.toHaveProperty("access_group_id");
expect(variables().params).not.toHaveProperty("created_at");
expect(variables().params).not.toHaveProperty("assigned_team_ids");
expect(variables().params).not.toHaveProperty("assigned_key_ids");
});
it("does not submit when the group name is cleared", async () => {
const user = setup();
renderModal();
const nameInput = await screen.findByDisplayValue("Engineering");
await user.clear(nameInput);
await save(user);
expect(await screen.findByText("Please enter the access group name")).toBeInTheDocument();
expect(mutate).not.toHaveBeenCalled();
});
it("does not save when Enter is pressed in the name field", async () => {
const user = setup();
renderModal();
const nameInput = await screen.findByDisplayValue("Engineering");
await user.type(nameInput, "{Enter}");
expect(mutate).not.toHaveBeenCalled();
});
it("sends models chosen on the Models tab", async () => {
const user = setup();
renderModal({ ...accessGroup, access_model_names: [] });
await screen.findByDisplayValue("Engineering");
await user.click(screen.getByRole("tab", { name: /Models/ }));
await user.click(await screen.findByLabelText("model-select"));
await save(user);
await waitFor(() => expect(mutate).toHaveBeenCalled());
expect(variables().params.access_model_names).toStrictEqual(["gpt-4"]);
});
});

View file

@ -1,10 +1,24 @@
import React, { useEffect } from "react";
import { Modal, Form } from "antd";
"use client";
import React, { useState } from "react";
import { Modal } from "antd";
import { toast } from "@/lib/toast";
import { AccessGroupBaseForm, AccessGroupFormValues } from "./AccessGroupBaseForm";
import { useZodForm } from "@/lib/forms/useZodForm";
import { Button } from "@/components/ui/button";
import { useEditAccessGroup, AccessGroupUpdateParams } from "@/app/(dashboard)/hooks/accessGroups/useEditAccessGroup";
import { AccessGroupResponse } from "@/app/(dashboard)/hooks/accessGroups/useAccessGroups";
import {
AccessGroupBaseForm,
accessGroupFormSchema,
AGENTS_TAB,
GENERAL_TAB,
MCP_SERVERS_TAB,
MODELS_TAB,
type AccessGroupFormValues,
} from "./AccessGroupBaseForm";
interface AccessGroupEditModalProps {
visible: boolean;
accessGroup: AccessGroupResponse;
@ -12,62 +26,74 @@ interface AccessGroupEditModalProps {
onSuccess?: () => void;
}
export function AccessGroupEditModal({ visible, accessGroup, onCancel, onSuccess }: AccessGroupEditModalProps) {
const [form] = Form.useForm<AccessGroupFormValues>();
const toFormValues = (accessGroup: AccessGroupResponse): AccessGroupFormValues => ({
name: accessGroup.access_group_name,
description: accessGroup.description ?? "",
modelIds: accessGroup.access_model_names ?? [],
mcpServerIds: accessGroup.access_mcp_server_ids ?? [],
agentIds: accessGroup.access_agent_ids ?? [],
});
function AccessGroupEditForm({ accessGroup, onCancel, onSuccess }: Omit<AccessGroupEditModalProps, "visible">) {
const form = useZodForm(accessGroupFormSchema, { defaultValues: toFormValues(accessGroup) });
const editMutation = useEditAccessGroup();
const [activeTab, setActiveTab] = useState(GENERAL_TAB);
const [visitedTabs, setVisitedTabs] = useState<ReadonlySet<string>>(new Set([GENERAL_TAB]));
// Populate the form with initial values whenever the modal opens or the data changes
useEffect(() => {
if (visible && accessGroup) {
form.setFieldsValue({
name: accessGroup.access_group_name,
description: accessGroup.description ?? "",
modelIds: accessGroup.access_model_names ?? [],
mcpServerIds: accessGroup.access_mcp_server_ids ?? [],
agentIds: accessGroup.access_agent_ids ?? [],
});
}
}, [visible, accessGroup, form]);
const handleOk = () => {
form
.validateFields()
.then((values) => {
const params: AccessGroupUpdateParams = {
access_group_name: values.name,
description: values.description,
access_model_names: values.modelIds,
access_mcp_server_ids: values.mcpServerIds,
access_agent_ids: values.agentIds,
};
editMutation.mutate(
{ accessGroupId: accessGroup.access_group_id, params },
{
onSuccess: () => {
toast.success("Access group updated successfully");
onSuccess?.();
onCancel();
},
},
);
})
.catch((info) => {});
const handleTabChange = (tab: string) => {
setActiveTab(tab);
setVisitedTabs((previous) => new Set([...previous, tab]));
};
const handleOk = form.handleSubmit(
(values) => {
const params: AccessGroupUpdateParams = {
access_group_name: values.name,
description: values.description,
access_model_names: visitedTabs.has(MODELS_TAB) ? values.modelIds : undefined,
access_mcp_server_ids: visitedTabs.has(MCP_SERVERS_TAB) ? values.mcpServerIds : undefined,
access_agent_ids: visitedTabs.has(AGENTS_TAB) ? values.agentIds : undefined,
};
editMutation.mutate(
{ accessGroupId: accessGroup.access_group_id, params },
{
onSuccess: () => {
toast.success("Access group updated successfully");
onSuccess?.();
onCancel();
},
},
);
},
() => setActiveTab(GENERAL_TAB),
);
return (
<Modal
title="Edit Access Group"
open={visible}
onOk={handleOk}
onCancel={onCancel}
width={700}
okText="Save Changes"
cancelText="Cancel"
confirmLoading={editMutation.isPending}
destroyOnHidden
>
<AccessGroupBaseForm form={form} />
<form onSubmit={(event) => event.preventDefault()}>
<AccessGroupBaseForm form={form} activeTab={activeTab} onTabChange={handleTabChange} />
<div className="mt-6 flex justify-end gap-2">
<Button type="button" variant="outline" onClick={onCancel} disabled={editMutation.isPending}>
Cancel
</Button>
<Button type="button" onClick={() => void handleOk()} disabled={editMutation.isPending}>
Save Changes
</Button>
</div>
</form>
);
}
export function AccessGroupEditModal({ visible, accessGroup, onCancel, onSuccess }: AccessGroupEditModalProps) {
return (
<Modal title="Edit Access Group" open={visible} onCancel={onCancel} width={700} footer={null} destroyOnHidden>
<AccessGroupEditForm
key={accessGroup.access_group_id}
accessGroup={accessGroup}
onCancel={onCancel}
onSuccess={onSuccess}
/>
</Modal>
);
}

View file

@ -1,4 +1,4 @@
import { render, screen, waitFor } from "@testing-library/react";
import { render, screen, waitFor, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import AdminPanel from "./AdminPanel";
@ -323,3 +323,73 @@ describe("AdminPanel", () => {
});
});
});
describe("AdminPanel add allowed IP form", () => {
beforeEach(async () => {
vi.clearAllMocks();
mockUseAuthorized.mockReturnValue({
premiumUser: true,
accessToken: "test-token",
userId: "user-1",
});
mockGetSSOSettings.mockResolvedValue({ values: {} });
mockGetAllowedIPs.mockResolvedValue(["10.0.0.1"]);
mockAddAllowedIP.mockResolvedValue({});
const user = userEvent.setup();
render(<AdminPanel />);
await user.click(screen.getByRole("tab", { name: /security settings/i }));
await user.click(screen.getByRole("button", { name: /allowed ips/i }));
const manageDialog = await screen.findByRole("dialog", { name: /manage allowed ip addresses/i });
await user.click(within(manageDialog).getByRole("button", { name: /add ip address/i }));
await screen.findByPlaceholderText("Enter IP address");
});
const ipField = () => screen.getByPlaceholderText("Enter IP address") as HTMLInputElement;
const submitAddIP = async (user: ReturnType<typeof userEvent.setup>) => {
const addIpForm = ipField().form as HTMLFormElement;
await user.click(within(addIpForm).getByText("Add IP Address"));
};
it("sends the access token and the typed IP address", async () => {
const user = userEvent.setup();
await user.type(ipField(), "192.168.1.50");
await submitAddIP(user);
await waitFor(() => {
expect(mockAddAllowedIP).toHaveBeenCalledWith("test-token", "192.168.1.50");
});
expect(mockAddAllowedIP).toHaveBeenCalledTimes(1);
});
it("blocks the submit and shows the required message when no IP is typed", async () => {
const user = userEvent.setup();
await submitAddIP(user);
expect(await screen.findByText("Please enter an IP address")).toBeInTheDocument();
expect(mockAddAllowedIP).not.toHaveBeenCalled();
});
it("submits on Enter from the IP field", async () => {
const user = userEvent.setup();
await user.type(ipField(), "172.16.0.9{Enter}");
await waitFor(() => {
expect(mockAddAllowedIP).toHaveBeenCalledWith("test-token", "172.16.0.9");
});
});
it("refreshes the allowed IP list after a successful add", async () => {
const user = userEvent.setup();
mockGetAllowedIPs.mockResolvedValue(["10.0.0.1", "192.168.1.50"]);
await user.type(ipField(), "192.168.1.50");
await submitAddIP(user);
expect(await screen.findByText("192.168.1.50")).toBeInTheDocument();
});
});

View file

@ -3,18 +3,12 @@
* Use this to avoid sharing master key with others
*/
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import {
Button,
Callout,
Card,
Table,
TableBody,
TableCell,
TableHead,
TableHeaderCell,
TableRow,
} from "@tremor/react";
import { Alert, Button as Button2, Form, Input, Modal, Space, Tabs, Typography } from "antd";
import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert";
import { Button } from "@/components/ui/button";
import { Card } from "@/components/ui/card";
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
import { Alert as AntdAlert, Modal, Space, Tabs, Typography } from "antd";
import { Info } from "lucide-react";
import React, { useEffect, useState } from "react";
import NewBadge from "@/components/common_components/NewBadge";
import { useBaseUrl } from "@/components/constants";
@ -28,17 +22,50 @@ import UserBannerSettings from "@/components/Settings/AdminSettings/UserBannerSe
import HashicorpVault from "@/components/Settings/AdminSettings/HashicorpVault/HashicorpVault";
import PluginSettings from "@/components/Settings/AdminSettings/PluginSettings/PluginSettings";
import SSOModals from "@/components/SSOModals";
import {
emptySSOSettingsFormValues,
useSSOSettingsForm,
type SSOSettingsFormValues,
} from "@/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm";
import UIAccessControlForm from "@/components/UIAccessControlForm";
import { z } from "zod/v4";
import { FieldGroup } from "@/components/shared/form/field";
import { FormField } from "@/components/shared/form/FormField";
import { Input } from "@/components/ui/input";
import { useZodForm } from "@/lib/forms/useZodForm";
const { Title, Paragraph, Text } = Typography;
const allowedIPSchema = z.object({
ip: z.string().min(1, "Please enter an IP address"),
});
type AllowedIPFormValues = z.infer<typeof allowedIPSchema>;
const AddAllowedIPForm = ({ onSubmit }: { onSubmit: (values: AllowedIPFormValues) => Promise<void> }) => {
const form = useZodForm(allowedIPSchema, { defaultValues: { ip: "" } });
return (
<form onSubmit={form.handleSubmit(onSubmit)}>
<FieldGroup>
<FormField control={form.control} name="ip">
{({ ref, ...field }) => <Input ref={ref} placeholder="Enter IP address" {...field} />}
</FormField>
<div>
<Button type="submit">Add IP Address</Button>
</div>
</FieldGroup>
</form>
);
};
interface AdminPanelProps {
proxySettings?: any;
}
const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
const { premiumUser, accessToken, userId: userID } = useAuthorized();
const [form] = Form.useForm();
const form = useSSOSettingsForm("admin-panel");
const [isAddSSOModalVisible, setIsAddSSOModalVisible] = useState(false);
const [isInstructionsModalVisible, setIsInstructionsModalVisible] = useState(false);
const [isAllowedIPModalVisible, setIsAllowedIPModalVisible] = useState(false);
@ -141,7 +168,7 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
const handleAddSSOOk = () => {
setIsAddSSOModalVisible(false);
form.resetFields();
form.reset(emptySSOSettingsFormValues);
if (accessToken && premiumUser) {
checkSSOConfiguration();
}
@ -149,10 +176,10 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
const handleAddSSOCancel = () => {
setIsAddSSOModalVisible(false);
form.resetFields();
form.reset(emptySSOSettingsFormValues);
};
const handleShowInstructions = (formValues: Record<string, any>) => {
const handleShowInstructions = (formValues: SSOSettingsFormValues) => {
setIsAddSSOModalVisible(false);
setIsInstructionsModalVisible(true);
};
@ -194,9 +221,9 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
label: "Security Settings",
children: (
<>
<Card>
<Card className="block p-6">
<Title level={4}> ✨ Security Settings</Title>
<Alert
<AntdAlert
message="SSO Configuration Deprecated"
description="Editing SSO Settings on this page is deprecated and will be removed in a future version. Please use the SSO Settings tab for SSO configuration."
type="warning"
@ -264,19 +291,19 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
]}
>
<Table>
<TableHead>
<TableHeader>
<TableRow>
<TableHeaderCell>IP Address</TableHeaderCell>
<TableHeaderCell className="text-right">Action</TableHeaderCell>
<TableHead>IP Address</TableHead>
<TableHead className="text-right">Action</TableHead>
</TableRow>
</TableHead>
</TableHeader>
<TableBody>
{allowedIPs.map((ip, index) => (
<TableRow key={index}>
<TableCell>{ip}</TableCell>
<TableCell className="text-right">
{ip !== all_ip_address_allowed && (
<Button onClick={() => handleDeleteIP(ip)} color="red" size="xs">
<Button onClick={() => handleDeleteIP(ip)} variant="destructive" size="sm">
Delete
</Button>
)}
@ -293,14 +320,7 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
onCancel={() => setIsAddIPModalVisible(false)}
footer={null}
>
<Form onFinish={handleAddIP}>
<Form.Item name="ip" rules={[{ required: true, message: "Please enter an IP address" }]}>
<Input placeholder="Enter IP address" />
</Form.Item>
<Form.Item>
<Button2 htmlType="submit">Add IP Address</Button2>
</Form.Item>
</Form>
<AddAllowedIPForm onSubmit={handleAddIP} />
</Modal>
<Modal
@ -338,12 +358,16 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
/>
</Modal>
</div>
<Callout title="Login without SSO" color="teal">
If you need to login without sso, you can access{" "}
<a href={nonSssoUrl} target="_blank" rel="noopener noreferrer">
<b>{nonSssoUrl}</b>{" "}
</a>
</Callout>
<Alert variant="info">
<Info />
<AlertTitle>Login without SSO</AlertTitle>
<AlertDescription>
If you need to login without sso, you can access{" "}
<a href={nonSssoUrl} target="_blank" rel="noopener noreferrer">
<b>{nonSssoUrl}</b>{" "}
</a>
</AlertDescription>
</Alert>
</>
),
},

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