diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 7ac52925459..a7ec31f2ffd 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -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 diff --git a/ci_cd/generate_model_prices_schema.py b/ci_cd/generate_model_prices_schema.py index 1b60f986ca4..153fbc0fdc2 100644 --- a/ci_cd/generate_model_prices_schema.py +++ b/ci_cd/generate_model_prices_schema.py @@ -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).", diff --git a/litellm/__init__.py b/litellm/__init__.py index ae0fee11aeb..1ecb04b6e54 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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": diff --git a/litellm/batches/main.py b/litellm/batches/main.py index a5df9b78601..2aa7b527c57 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -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, diff --git a/litellm/caching/_embedding_router.py b/litellm/caching/_embedding_router.py index 1073b34ef25..8dfcddf158a 100644 --- a/litellm/caching/_embedding_router.py +++ b/litellm/caching/_embedding_router.py @@ -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 diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index f0fb91b987f..6b68ae98111 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -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() diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index 8f8323550f3..8270c655d82 100644 --- a/litellm/caching/qdrant_semantic_cache.py +++ b/litellm/caching/qdrant_semantic_cache.py @@ -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}, ) diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index 604d6395ea1..d91260f4d9c 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -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"] diff --git a/litellm/caching/valkey_semantic_cache.py b/litellm/caching/valkey_semantic_cache.py index aa10d91fc66..737d212a89d 100644 --- a/litellm/caching/valkey_semantic_cache.py +++ b/litellm/caching/valkey_semantic_cache.py @@ -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 ( diff --git a/litellm/constants.py b/litellm/constants.py index e73fba1cc9f..39a49e55f0d 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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({}) diff --git a/litellm/files/main.py b/litellm/files/main.py index 9a64c78552b..294c62f3d80 100644 --- a/litellm/files/main.py +++ b/litellm/files/main.py @@ -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 diff --git a/litellm/files/types.py b/litellm/files/types.py index 8cadd69f024..b4ec9996f37 100644 --- a/litellm/files/types.py +++ b/litellm/files/types.py @@ -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): diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index edb4d56a5b7..10e681b816e 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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), diff --git a/litellm/litellm_core_utils/llm_cost_calc/guardrail_cost.py b/litellm/litellm_core_utils/llm_cost_calc/guardrail_cost.py new file mode 100644 index 00000000000..4645a8c3074 --- /dev/null +++ b/litellm/litellm_core_utils/llm_cost_calc/guardrail_cost.py @@ -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 diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 9d6ad8b6e39..f73c4942a1c 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -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") diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index 1660f56378f..30b5df1e4ee 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -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: diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/advisor.py b/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/advisor.py index dfae7b4f4cf..701211049db 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/advisor.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/interceptors/advisor.py @@ -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], diff --git a/litellm/llms/azure_ai/chat/transformation.py b/litellm/llms/azure_ai/chat/transformation.py index 8545d646035..bc8ea31ea8c 100644 --- a/litellm/llms/azure_ai/chat/transformation.py +++ b/litellm/llms/azure_ai/chat/transformation.py @@ -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: diff --git a/litellm/llms/base_llm/vector_store/transformation.py b/litellm/llms/base_llm/vector_store/transformation.py index 8083d2485ba..02a51a8bace 100644 --- a/litellm/llms/base_llm/vector_store/transformation.py +++ b/litellm/llms/base_llm/vector_store/transformation.py @@ -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 diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index d10628903dd..7998a967344 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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: diff --git a/litellm/llms/fireworks_ai/common_utils.py b/litellm/llms/fireworks_ai/common_utils.py index e07e7a26f9e..8e35cfebc5b 100644 --- a/litellm/llms/fireworks_ai/common_utils.py +++ b/litellm/llms/fireworks_ai/common_utils.py @@ -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}" diff --git a/litellm/llms/valkey/__init__.py b/litellm/llms/valkey/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/valkey/common_utils.py b/litellm/llms/valkey/common_utils.py new file mode 100644 index 00000000000..9691450f3e0 --- /dev/null +++ b/litellm/llms/valkey/common_utils.py @@ -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) diff --git a/litellm/llms/valkey/vector_stores/__init__.py b/litellm/llms/valkey/vector_stores/__init__.py new file mode 100644 index 00000000000..c826607a800 --- /dev/null +++ b/litellm/llms/valkey/vector_stores/__init__.py @@ -0,0 +1,3 @@ +from litellm.llms.valkey.vector_stores.transformation import ValkeyVectorStoreConfig + +__all__ = ("ValkeyVectorStoreConfig",) diff --git a/litellm/llms/valkey/vector_stores/transformation.py b/litellm/llms/valkey/vector_stores/transformation.py new file mode 100644 index 00000000000..3cbfca0f1a9 --- /dev/null +++ b/litellm/llms/valkey/vector_stores/transformation.py @@ -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) diff --git a/litellm/main.py b/litellm/main.py index 2a8ed6c87b6..cc27da830d8 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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, } diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 78b53cefc53..409022016b0 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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", diff --git a/litellm/proxy/_experimental/out/assets/logos/valkey.svg b/litellm/proxy/_experimental/out/assets/logos/valkey.svg new file mode 100644 index 00000000000..0e97e680df4 --- /dev/null +++ b/litellm/proxy/_experimental/out/assets/logos/valkey.svg @@ -0,0 +1,6 @@ + + + + + + diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index a7b10afb0ab..8e57327b31b 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 = [] diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index adae59a1174..a0b69ecb0bf 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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, diff --git a/litellm/proxy/common_utils/sse_keepalive.py b/litellm/proxy/common_utils/sse_keepalive.py index e5183ac29d4..51ace009f54 100644 --- a/litellm/proxy/common_utils/sse_keepalive.py +++ b/litellm/proxy/common_utils/sse_keepalive.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 93cbb989e23..c70a2ee8a74 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -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]: diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index ca89c7587ba..9d0d84dc2b1 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -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, } }, diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 5dc92d82bda..99d0c94d11b 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index ef6e590ba26..2ec5c34958c 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -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 diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index 3ae871b476e..ffca858c0ce 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -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) diff --git a/litellm/proxy/management_endpoints/team_callback_endpoints.py b/litellm/proxy/management_endpoints/team_callback_endpoints.py index 834d4e8b73b..472e25bbc28 100644 --- a/litellm/proxy/management_endpoints/team_callback_endpoints.py +++ b/litellm/proxy/management_endpoints/team_callback_endpoints.py @@ -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"], diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 211c3742343..0df0aaa1bcd 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5d4a306a73e..d8fe7fce78f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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"], ) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index e0af363b1a5..d09a30a7e3a 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -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, diff --git a/litellm/router.py b/litellm/router.py index 8b9c4b0db1a..c5881960c80 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 ``, 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"] = {} diff --git a/litellm/router_strategy/auto_router/auto_router.py b/litellm/router_strategy/auto_router/auto_router.py index 20c0ece46b6..c77745a498d 100644 --- a/litellm/router_strategy/auto_router/auto_router.py +++ b/litellm/router_strategy/auto_router/auto_router.py @@ -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( diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index c829e6d9442..d16063b9bd4 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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: diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 20329de34e2..6d43199c948 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -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: diff --git a/litellm/types/management_endpoints/auto_router_endpoints.py b/litellm/types/management_endpoints/auto_router_endpoints.py index 61601dd31eb..9461297feca 100644 --- a/litellm/types/management_endpoints/auto_router_endpoints.py +++ b/litellm/types/management_endpoints/auto_router_endpoints.py @@ -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): diff --git a/litellm/types/router.py b/litellm/types/router.py index f3f9276e6ba..7d1dd1358d5 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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.""" diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 13831799c7f..d44d4cca6c4 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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): diff --git a/litellm/utils.py b/litellm/utils.py index 68f4278c87a..1c880ee9521 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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 diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 78b53cefc53..409022016b0 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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", diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index 4c54822736c..cd02fde595f 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -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", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 0712e8e383d..ec0b1c27344 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -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", diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 4a6ff6af902..6882479a344 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -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 diff --git a/tests/documentation_tests/test_readme_providers.py b/tests/documentation_tests/test_readme_providers.py index f9de25bc85b..d3b4e22180b 100644 --- a/tests/documentation_tests/test_readme_providers.py +++ b/tests/documentation_tests/test_readme_providers.py @@ -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 } diff --git a/tests/local_testing/test_router.py b/tests/local_testing/test_router.py index 2ee1aee9710..af1a052e86d 100644 --- a/tests/local_testing/test_router.py +++ b/tests/local_testing/test_router.py @@ -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 ) diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index c883890f5f6..c3db9e67f9c 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -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) diff --git a/tests/test_litellm/caching/test_caching.py b/tests/test_litellm/caching/test_caching.py index b65e8773c85..955b0e531bc 100644 --- a/tests/test_litellm/caching/test_caching.py +++ b/tests/test_litellm/caching/test_caching.py @@ -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"} diff --git a/tests/test_litellm/caching/test_embedding_router.py b/tests/test_litellm/caching/test_embedding_router.py index 550095a112a..9ebe669d32d 100644 --- a/tests/test_litellm/caching/test_embedding_router.py +++ b/tests/test_litellm/caching/test_embedding_router.py @@ -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 diff --git a/tests/test_litellm/caching/test_qdrant_semantic_cache.py b/tests/test_litellm/caching/test_qdrant_semantic_cache.py index 67d4e2d9892..852bed4a9df 100644 --- a/tests/test_litellm/caching/test_qdrant_semantic_cache.py +++ b/tests/test_litellm/caching/test_qdrant_semantic_cache.py @@ -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 diff --git a/tests/test_litellm/caching/test_redis_semantic_cache.py b/tests/test_litellm/caching/test_redis_semantic_cache.py index 1d3129d6467..9fd333cf87c 100644 --- a/tests/test_litellm/caching/test_redis_semantic_cache.py +++ b/tests/test_litellm/caching/test_redis_semantic_cache.py @@ -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() diff --git a/tests/test_litellm/caching/test_valkey_semantic_cache.py b/tests/test_litellm/caching/test_valkey_semantic_cache.py index d2df0a98e12..acf5a914e5c 100644 --- a/tests/test_litellm/caching/test_valkey_semantic_cache.py +++ b/tests/test_litellm/caching/test_valkey_semantic_cache.py @@ -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( diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_guardrail_cost.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_guardrail_cost.py new file mode 100644 index 00000000000..cf36a2b9b25 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_guardrail_cost.py @@ -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} diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 1402e056b72..1826f56d667 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -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(): diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 946f19b7658..0d54680fa81 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -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 diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py index 2eb8e077320..bd02c61752e 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py @@ -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) diff --git a/tests/test_litellm/llms/anthropic/messages/test_advisor_orchestration.py b/tests/test_litellm/llms/anthropic/messages/test_advisor_orchestration.py index 3d35e93167f..da5b5ac3867 100644 --- a/tests/test_litellm/llms/anthropic/messages/test_advisor_orchestration.py +++ b/tests/test_litellm/llms/anthropic/messages/test_advisor_orchestration.py @@ -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 diff --git a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py index a541ab2b3c6..900372f3e54 100644 --- a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py +++ b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py @@ -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" + } diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index fddd8d09dfc..2dda8bf722a 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -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" diff --git a/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_common_utils.py b/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_common_utils.py index 4af395baf41..b52c910d5a6 100644 --- a/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_common_utils.py +++ b/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_common_utils.py @@ -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): diff --git a/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py b/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py index 35c0a63573f..93c518599d6 100644 --- a/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py +++ b/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py @@ -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"]) diff --git a/tests/test_litellm/llms/valkey/vector_stores/test_valkey_transformation.py b/tests/test_litellm/llms/valkey/vector_stores/test_valkey_transformation.py new file mode 100644 index 00000000000..a2ee2c2bdb1 --- /dev/null +++ b/tests/test_litellm/llms/valkey/vector_stores/test_valkey_transformation.py @@ -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) diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 0bfb10320f7..21511c74154 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -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, + ) diff --git a/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py b/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py index 9cca9bbfe12..89ae74920fe 100644 --- a/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py +++ b/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 53921e7e74a..e3516b6eda7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -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"] diff --git a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py index be92b6fc6c4..c4a0a62ef97 100644 --- a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py @@ -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." ) diff --git a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py index f63c08a2c39..ff143bd055f 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py @@ -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) diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index bca8210baa6..50c93ed5275 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py index 049d6ddb183..77149457e82 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py index 7ed123f6cdf..3061da336f6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py @@ -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": []}} diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py index 21e25d30b82..c2610d88927 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py @@ -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"] diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index eefdefcee80..b1b934b0949 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -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()) diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 17dd486763d..f31f67c317a 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 9ddd74a46a8..355c6d27eb2 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -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() diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 0171e83be08..b1071150f3b 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -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) diff --git a/tests/test_litellm/proxy/test_pricing_field_strip.py b/tests/test_litellm/proxy/test_pricing_field_strip.py index 25377a6d209..bbdddd1cd8c 100644 --- a/tests/test_litellm/proxy/test_pricing_field_strip.py +++ b/tests/test_litellm/proxy/test_pricing_field_strip.py @@ -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 diff --git a/tests/test_litellm/proxy/test_update_llm_router_resilience.py b/tests/test_litellm/proxy/test_update_llm_router_resilience.py index a7dc9c1783e..d6ebfde1091 100644 --- a/tests/test_litellm/proxy/test_update_llm_router_resilience.py +++ b/tests/test_litellm/proxy/test_update_llm_router_resilience.py @@ -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 diff --git a/tests/test_litellm/responses/test_responses_api_request_body.py b/tests/test_litellm/responses/test_responses_api_request_body.py index e7f87f1e326..ea5eb3afe9a 100644 --- a/tests/test_litellm/responses/test_responses_api_request_body.py +++ b/tests/test_litellm/responses/test_responses_api_request_body.py @@ -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/ reach this path directly via + /v1/responses and via the /v1/messages adapter (which passes responses/ + with custom_llm_provider="openai"), so both shapes must hit OpenAI as . + """ + 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" diff --git a/tests/test_litellm/router_strategy/test_auto_router.py b/tests/test_litellm/router_strategy/test_auto_router.py index b73485e6019..c71a6b0e27f 100644 --- a/tests/test_litellm/router_strategy/test_auto_router.py +++ b/tests/test_litellm/router_strategy/test_auto_router.py @@ -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 == "" diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 7034ebc9bd7..e1e8d9553b3 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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.""" diff --git a/tests/test_litellm/router_strategy/test_router_routing_plugins.py b/tests/test_litellm/router_strategy/test_router_routing_plugins.py index 78ed71f5ffd..293af36080a 100644 --- a/tests/test_litellm/router_strategy/test_router_routing_plugins.py +++ b/tests/test_litellm/router_strategy/test_router_routing_plugins.py @@ -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 ``) -- 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", diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index 58373df024c..68b8d1c62b5 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -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", diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index b3c348a1221..49ed236356c 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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(): """ diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 2b0b8b6ab20..afdfdf170ac 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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 diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 7a099462256..f8e481dc142 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -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 diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 2436cb0e3b7..bc23c2433d3 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -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": { diff --git a/ui/litellm-dashboard/public/assets/logos/valkey.svg b/ui/litellm-dashboard/public/assets/logos/valkey.svg new file mode 100644 index 00000000000..0e97e680df4 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/valkey.svg @@ -0,0 +1,6 @@ + + + + + + diff --git a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupBaseForm.tsx b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupBaseForm.tsx index 72b89e34301..c9ba082f7a6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupBaseForm.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupBaseForm.tsx @@ -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; + +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) => ( + +); + interface AccessGroupBaseFormProps { - form: FormInstance; + form: UseFormReturn; 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: ( - - - General Info - - ), - children: ( -
- - - - -