mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_pr35110_itpm_otpm
# Conflicts: # type-discipline-budget.json
This commit is contained in:
commit
c435c25da2
259 changed files with 21880 additions and 7178 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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).",
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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({})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
78
litellm/litellm_core_utils/llm_cost_calc/guardrail_cost.py
Normal file
78
litellm/litellm_core_utils/llm_cost_calc/guardrail_cost.py
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
import math
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.utils import CostBreakdown
|
||||
|
||||
BEDROCK_GUARDRAIL_PRICING_KEY: Final = "bedrock/guardrails"
|
||||
|
||||
|
||||
class GuardrailPricing(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
guardrail_cost_per_unit: Mapping[str, float]
|
||||
|
||||
|
||||
class GuardrailCostEntry(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
guardrail_cost: float | None = None
|
||||
|
||||
|
||||
GuardrailInformationShape = tuple[GuardrailCostEntry, ...] | GuardrailCostEntry | None
|
||||
|
||||
_GUARDRAIL_INFORMATION_ADAPTER: Final[TypeAdapter[GuardrailInformationShape]] = TypeAdapter(GuardrailInformationShape)
|
||||
|
||||
|
||||
def _bedrock_guardrail_pricing(aws_region_name: str | None) -> GuardrailPricing | None:
|
||||
regional_key: Final = f"bedrock/{aws_region_name}/guardrails" if aws_region_name else None
|
||||
for key in (regional_key, BEDROCK_GUARDRAIL_PRICING_KEY):
|
||||
if key is None or key not in litellm.model_cost:
|
||||
continue
|
||||
try:
|
||||
return GuardrailPricing.model_validate(litellm.model_cost[key])
|
||||
except ValidationError as e:
|
||||
verbose_logger.warning("Ignoring malformed guardrail pricing entry %s: %s", key, e)
|
||||
return None
|
||||
|
||||
|
||||
def bedrock_guardrail_cost(usage_units: Mapping[str, int], aws_region_name: str | None) -> float:
|
||||
pricing: Final = _bedrock_guardrail_pricing(aws_region_name)
|
||||
if pricing is None:
|
||||
return 0.0
|
||||
return sum(units * pricing.guardrail_cost_per_unit.get(counter, 0.0) for counter, units in usage_units.items())
|
||||
|
||||
|
||||
def _billable_entry_cost(entry: GuardrailCostEntry) -> float:
|
||||
cost: Final = entry.guardrail_cost
|
||||
if cost is None or not math.isfinite(cost) or cost <= 0.0:
|
||||
return 0.0
|
||||
return cost
|
||||
|
||||
|
||||
def guardrail_information_cost(guardrail_information: object) -> float:
|
||||
try:
|
||||
parsed: Final = _GUARDRAIL_INFORMATION_ADAPTER.validate_python(guardrail_information)
|
||||
except ValidationError:
|
||||
return 0.0
|
||||
if parsed is None:
|
||||
return 0.0
|
||||
if isinstance(parsed, GuardrailCostEntry):
|
||||
return _billable_entry_cost(parsed)
|
||||
return sum(_billable_entry_cost(entry) for entry in parsed)
|
||||
|
||||
|
||||
def cost_breakdown_with_guardrail(cost_breakdown: CostBreakdown | None, guardrail_cost: float) -> CostBreakdown | None:
|
||||
if guardrail_cost <= 0.0:
|
||||
return cost_breakdown
|
||||
existing: Final[CostBreakdown] = cost_breakdown if cost_breakdown is not None else CostBreakdown()
|
||||
merged: Final[CostBreakdown] = {
|
||||
**existing,
|
||||
"guardrail_cost": guardrail_cost,
|
||||
"total_cost": existing.get("total_cost", 0.0) + guardrail_cost,
|
||||
}
|
||||
return merged
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
0
litellm/llms/valkey/__init__.py
Normal file
0
litellm/llms/valkey/__init__.py
Normal file
18
litellm/llms/valkey/common_utils.py
Normal file
18
litellm/llms/valkey/common_utils.py
Normal file
|
|
@ -0,0 +1,18 @@
|
|||
"""Shared helpers for Valkey integrations (semantic cache, vector stores)."""
|
||||
|
||||
import struct
|
||||
from collections.abc import Sequence
|
||||
from typing import Final
|
||||
from urllib.parse import quote
|
||||
|
||||
|
||||
def build_valkey_url(host: str, port: str, password: str | None = None, ssl: bool = False) -> str:
|
||||
"""Deliberately reads no environment: callers of the vector store control the
|
||||
host, so an env-sourced password would be sent to a caller-chosen server."""
|
||||
credentials: Final = f":{quote(password, safe='')}@" if password else ""
|
||||
scheme: Final = "rediss" if ssl else "redis"
|
||||
return f"{scheme}://{credentials}{host}:{port}"
|
||||
|
||||
|
||||
def pack_vector(embedding: Sequence[float]) -> bytes:
|
||||
return struct.pack(f"<{len(embedding)}f", *embedding)
|
||||
3
litellm/llms/valkey/vector_stores/__init__.py
Normal file
3
litellm/llms/valkey/vector_stores/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from litellm.llms.valkey.vector_stores.transformation import ValkeyVectorStoreConfig
|
||||
|
||||
__all__ = ("ValkeyVectorStoreConfig",)
|
||||
299
litellm/llms/valkey/vector_stores/transformation.py
Normal file
299
litellm/llms/valkey/vector_stores/transformation.py
Normal file
|
|
@ -0,0 +1,299 @@
|
|||
"""
|
||||
Valkey vector store provider.
|
||||
|
||||
Valkey's vector search (the valkey-search module) speaks RESP only, no HTTP
|
||||
API, so this config extends BaseDirectVectorStoreConfig and executes the
|
||||
FT.SEARCH KNN query itself via redis-py instead of shaping an httpx request.
|
||||
Documents are HASHes indexed by an FT index named after the vector_store_id.
|
||||
"""
|
||||
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, NoReturn
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.vector_store.transformation import BaseDirectVectorStoreConfig
|
||||
from litellm.llms.valkey.common_utils import build_valkey_url, pack_vector
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.types.vector_stores import (
|
||||
VectorStoreCreateOptionalRequestParams,
|
||||
VectorStoreResultContent,
|
||||
VectorStoreSearchOptionalRequestParams,
|
||||
VectorStoreSearchResponse,
|
||||
VectorStoreSearchResult,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from redis import Redis
|
||||
from redis.asyncio import Redis as AsyncRedis
|
||||
from redis.commands.search.document import Document
|
||||
from redis.commands.search.query import Query
|
||||
from redis.commands.search.result import Result
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
DEFAULT_VALKEY_PORT: Final = 6379
|
||||
DEFAULT_SOCKET_CONNECT_TIMEOUT_SECONDS: Final = 5.0
|
||||
DEFAULT_SOCKET_TIMEOUT_SECONDS: Final = 30.0
|
||||
DEFAULT_MAX_NUM_RESULTS: Final = 10
|
||||
MIN_MAX_NUM_RESULTS: Final = 1
|
||||
MAX_MAX_NUM_RESULTS: Final = 50
|
||||
DEFAULT_EMBEDDING_FIELD_NAME: Final = "embedding"
|
||||
DEFAULT_TEXT_FIELD_NAME: Final = "text"
|
||||
DISTANCE_FIELD_NAME: Final = "vector_distance"
|
||||
|
||||
_EMPTY_EMBEDDING_CONFIG: Final = MappingProxyType({})
|
||||
_REDIS_INSTALL_HINT: Final = (
|
||||
"The Valkey vector store requires the 'redis' package. Run 'pip install redis' to install it."
|
||||
)
|
||||
_SEARCH_ONLY_MESSAGE: Final = "Valkey vector store is search-only; create indexes with FT.CREATE directly"
|
||||
|
||||
|
||||
def _import_sync_redis() -> "type[Redis]":
|
||||
try:
|
||||
from redis import Redis as SyncRedisClient
|
||||
except ImportError as e:
|
||||
raise ValueError(_REDIS_INSTALL_HINT) from e
|
||||
return SyncRedisClient
|
||||
|
||||
|
||||
def _import_async_redis() -> "type[AsyncRedis]":
|
||||
try:
|
||||
from redis.asyncio import Redis as AsyncRedisClient
|
||||
except ImportError as e:
|
||||
raise ValueError(_REDIS_INSTALL_HINT) from e
|
||||
return AsyncRedisClient
|
||||
|
||||
|
||||
def _import_query() -> "type[Query]":
|
||||
try:
|
||||
from redis.commands.search.query import Query as RedisQuery
|
||||
except ImportError as e:
|
||||
raise ValueError(_REDIS_INSTALL_HINT) from e
|
||||
return RedisQuery
|
||||
|
||||
|
||||
class _ValkeySearchParams(BaseModel):
|
||||
"""Typed view over the vector store's litellm_params; unrelated keys are ignored."""
|
||||
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
litellm_embedding_model: str | None = None
|
||||
litellm_embedding_config: Mapping[str, object] | None = None
|
||||
valkey_host: str | None = None
|
||||
valkey_port: int | None = None
|
||||
valkey_password: str | None = None
|
||||
valkey_ssl: bool | None = None
|
||||
valkey_text_field: str | None = None
|
||||
valkey_embedding_field: str | None = None
|
||||
|
||||
@property
|
||||
def text_field(self) -> str:
|
||||
return self.valkey_text_field or DEFAULT_TEXT_FIELD_NAME
|
||||
|
||||
@property
|
||||
def embedding_field(self) -> str:
|
||||
return self.valkey_embedding_field or DEFAULT_EMBEDDING_FIELD_NAME
|
||||
|
||||
def require_embedding_model(self) -> str:
|
||||
if not self.litellm_embedding_model:
|
||||
raise ValueError(
|
||||
"litellm_embedding_model is required in litellm_params for the Valkey vector store. "
|
||||
"Example: litellm_params['litellm_embedding_model'] = 'openai/text-embedding-3-small'"
|
||||
)
|
||||
return self.litellm_embedding_model
|
||||
|
||||
def connection_url(self) -> str:
|
||||
if not self.valkey_host:
|
||||
raise ValueError(
|
||||
"valkey_host is required in litellm_params for the Valkey vector store. "
|
||||
"Set it on the vector store's litellm_params, e.g. valkey_host: my-valkey.example.com"
|
||||
)
|
||||
return build_valkey_url(
|
||||
host=self.valkey_host,
|
||||
port=str(self.valkey_port or DEFAULT_VALKEY_PORT),
|
||||
password=self.valkey_password,
|
||||
ssl=bool(self.valkey_ssl),
|
||||
)
|
||||
|
||||
|
||||
class ValkeyVectorStoreConfig(BaseDirectVectorStoreConfig):
|
||||
def __init__(
|
||||
self,
|
||||
sync_client: "Redis | None" = None,
|
||||
async_client: "AsyncRedis | None" = None,
|
||||
embedding_fn: Callable[..., EmbeddingResponse] | None = None,
|
||||
aembedding_fn: Callable[..., Awaitable[EmbeddingResponse]] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.sync_client = sync_client
|
||||
self.async_client = async_client
|
||||
self.embedding_fn = embedding_fn if embedding_fn is not None else litellm.embedding
|
||||
self.aembedding_fn = aembedding_fn if aembedding_fn is not None else litellm.aembedding
|
||||
|
||||
@staticmethod
|
||||
def _query_text(query: str | Sequence[str]) -> str:
|
||||
if isinstance(query, str):
|
||||
return query
|
||||
if not query:
|
||||
raise ValueError("query must not be empty")
|
||||
return " ".join(query)
|
||||
|
||||
@staticmethod
|
||||
def _socket_timeouts(timeout: float | httpx.Timeout | None) -> tuple[float, float]:
|
||||
if isinstance(timeout, httpx.Timeout):
|
||||
return (
|
||||
timeout.connect or DEFAULT_SOCKET_CONNECT_TIMEOUT_SECONDS,
|
||||
timeout.read or DEFAULT_SOCKET_TIMEOUT_SECONDS,
|
||||
)
|
||||
if timeout is not None:
|
||||
return (min(float(timeout), DEFAULT_SOCKET_CONNECT_TIMEOUT_SECONDS), float(timeout))
|
||||
return (DEFAULT_SOCKET_CONNECT_TIMEOUT_SECONDS, DEFAULT_SOCKET_TIMEOUT_SECONDS)
|
||||
|
||||
@staticmethod
|
||||
def _knn_limit(vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams) -> int:
|
||||
requested: Final = vector_store_search_optional_params.get("max_num_results")
|
||||
if requested is None:
|
||||
return DEFAULT_MAX_NUM_RESULTS
|
||||
if not MIN_MAX_NUM_RESULTS <= requested <= MAX_MAX_NUM_RESULTS:
|
||||
raise ValueError(
|
||||
f"max_num_results must be between {MIN_MAX_NUM_RESULTS} and {MAX_MAX_NUM_RESULTS}, got {requested}"
|
||||
)
|
||||
return requested
|
||||
|
||||
@classmethod
|
||||
def _knn_query(
|
||||
cls,
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
embedding_field: str,
|
||||
text_field: str,
|
||||
) -> "Query":
|
||||
if vector_store_search_optional_params.get("filters") is not None:
|
||||
raise ValueError("Valkey vector store does not support the filters parameter yet")
|
||||
k: Final = cls._knn_limit(vector_store_search_optional_params)
|
||||
query_cls: Final = _import_query()
|
||||
knn_expr: Final = f"*=>[KNN {k} @{embedding_field} $vec AS {DISTANCE_FIELD_NAME}]"
|
||||
# valkey-search rejects SORTBY on the KNN distance alias, so results are
|
||||
# re-ordered client-side in _to_response instead.
|
||||
return query_cls(knn_expr).return_fields(text_field, DISTANCE_FIELD_NAME).paging(0, k).dialect(2)
|
||||
|
||||
@staticmethod
|
||||
def _to_result(doc: "Document", text_field: str) -> VectorStoreSearchResult:
|
||||
content: Final = [ # mutable-ok: VectorStoreSearchResult declares a list of content parts
|
||||
VectorStoreResultContent(text=str(getattr(doc, text_field, "")), type="text")
|
||||
]
|
||||
return VectorStoreSearchResult(
|
||||
score=1.0 - float(getattr(doc, DISTANCE_FIELD_NAME)),
|
||||
content=content,
|
||||
file_id=getattr(doc, "id", None),
|
||||
filename=getattr(doc, "id", None),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _to_response(cls, search_result: "Result", query_text: str, text_field: str) -> VectorStoreSearchResponse:
|
||||
docs: Final = getattr(search_result, "docs", None) or ()
|
||||
data: Final = sorted(
|
||||
(cls._to_result(doc, text_field) for doc in docs),
|
||||
key=lambda result: result.get("score") or 0.0,
|
||||
reverse=True,
|
||||
)
|
||||
return VectorStoreSearchResponse(
|
||||
object="vector_store.search_results.page",
|
||||
search_query=query_text,
|
||||
data=data,
|
||||
)
|
||||
|
||||
def execute_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
litellm_params: Mapping[str, object],
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
params: Final = _ValkeySearchParams.model_validate(litellm_params)
|
||||
query_text: Final = self._query_text(query)
|
||||
knn: Final = self._knn_query(
|
||||
vector_store_search_optional_params,
|
||||
embedding_field=params.embedding_field,
|
||||
text_field=params.text_field,
|
||||
)
|
||||
embedding_response: Final = self.embedding_fn(
|
||||
model=params.require_embedding_model(),
|
||||
input=[query_text], # mutable-ok: litellm.embedding's input contract is a list
|
||||
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
|
||||
)
|
||||
vec_params: Final = {"vec": pack_vector(embedding_response.data[0]["embedding"])} # mutable-ok: redis-py API
|
||||
|
||||
if self.sync_client is not None:
|
||||
raw: Final = self.sync_client.ft(vector_store_id).search(knn, query_params=vec_params)
|
||||
return self._to_response(raw, query_text, params.text_field)
|
||||
|
||||
connect_timeout, op_timeout = self._socket_timeouts(timeout)
|
||||
client: Final = _import_sync_redis().from_url(
|
||||
params.connection_url(),
|
||||
socket_connect_timeout=connect_timeout,
|
||||
socket_timeout=op_timeout,
|
||||
)
|
||||
try:
|
||||
raw_result: Final = client.ft(vector_store_id).search(knn, query_params=vec_params)
|
||||
return self._to_response(raw_result, query_text, params.text_field)
|
||||
finally:
|
||||
client.close()
|
||||
|
||||
async def aexecute_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
litellm_params: Mapping[str, object],
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
params: Final = _ValkeySearchParams.model_validate(litellm_params)
|
||||
query_text: Final = self._query_text(query)
|
||||
knn: Final = self._knn_query(
|
||||
vector_store_search_optional_params,
|
||||
embedding_field=params.embedding_field,
|
||||
text_field=params.text_field,
|
||||
)
|
||||
embedding_response: Final = await self.aembedding_fn(
|
||||
model=params.require_embedding_model(),
|
||||
input=[query_text], # mutable-ok: litellm.embedding's input contract is a list
|
||||
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
|
||||
)
|
||||
vec_params: Final = {"vec": pack_vector(embedding_response.data[0]["embedding"])} # mutable-ok: redis-py API
|
||||
|
||||
if self.async_client is not None:
|
||||
raw: Final = await self.async_client.ft(vector_store_id).search( # pyright: ignore[reportGeneralTypeIssues] # types-redis 4.6 stubs shadow redis 5.3.1 and type the async client's ft() as the sync Search, so search() returns a non-awaitable Result; it is a coroutine at runtime
|
||||
knn, query_params=vec_params
|
||||
)
|
||||
return self._to_response(raw, query_text, params.text_field)
|
||||
|
||||
connect_timeout, op_timeout = self._socket_timeouts(timeout)
|
||||
client: Final = _import_async_redis().from_url(
|
||||
params.connection_url(),
|
||||
socket_connect_timeout=connect_timeout,
|
||||
socket_timeout=op_timeout,
|
||||
)
|
||||
try:
|
||||
raw_result: Final = await client.ft(vector_store_id).search( # pyright: ignore[reportGeneralTypeIssues] # types-redis 4.6 stubs shadow redis 5.3.1 and type the async client's ft() as the sync Search, so search() returns a non-awaitable Result; it is a coroutine at runtime
|
||||
knn, query_params=vec_params
|
||||
)
|
||||
return self._to_response(raw_result, query_text, params.text_field)
|
||||
finally:
|
||||
await client.aclose()
|
||||
|
||||
def transform_create_vector_store_request(
|
||||
self,
|
||||
vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams,
|
||||
api_base: str,
|
||||
) -> NoReturn:
|
||||
raise NotImplementedError(_SEARCH_ONLY_MESSAGE)
|
||||
|
||||
def transform_create_vector_store_response(self, response: httpx.Response) -> NoReturn:
|
||||
raise NotImplementedError(_SEARCH_ONLY_MESSAGE)
|
||||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
6
litellm/proxy/_experimental/out/assets/logos/valkey.svg
Normal file
6
litellm/proxy/_experimental/out/assets/logos/valkey.svg
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<svg width="64" height="73" viewBox="0 0 64 73" xmlns="http://www.w3.org/2000/svg">
|
||||
<g id="Group-copy">
|
||||
<path id="Path" fill="#123678" fill-rule="evenodd" stroke="none" d="M 13.482285 60.694962 L 0.998384 52.884399 L 0.998384 19.502914 L 31.527868 2.001205 L 61.317604 19.532024 L 61.317604 54.64489 L 31.054855 71.68927 L 20.548372 65.115807 L 20.548372 51.041328 L 20.548372 49.119896 L 14.851504 45.555508 L 14.851504 27.453159 L 31.346497 17.99712 L 47.464485 27.482262 L 47.464485 46.451157 L 34.703495 53.638138 L 34.703495 45.998573 C 38.52874 44.52552 41.274452 40.739189 41.274452 36.270489 C 41.274452 30.510658 36.712814 25.88438 31.158138 25.88438 C 25.603172 25.88438 21.041817 30.510658 21.041817 36.270489 C 21.041817 40.739189 23.787249 44.52552 27.612494 45.998573 L 27.612494 60.473576 L 31.261133 62.756348 L 53.635483 50.15464 L 53.635483 23.924595 L 31.477489 10.884869 L 8.680504 23.953705 L 8.680504 48.628967 L 13.482285 51.633297 L 13.482285 60.694962 Z M 31.158138 31.498383 C 33.671822 31.498383 35.660439 33.664162 35.660439 36.270489 C 35.660439 38.876804 33.671822 41.042587 31.158138 41.042587 C 28.644447 41.042587 26.655558 38.876804 26.655558 36.270489 C 26.655558 33.664162 28.644447 31.498383 31.158138 31.498383 Z" />
|
||||
</g>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 1.3 KiB |
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -204,6 +204,7 @@ from litellm.types.utils import (
|
|||
CustomPricingLiteLLMParams,
|
||||
GenericBudgetConfigType,
|
||||
LiteLLMBatch,
|
||||
LlmProviders,
|
||||
ModelInfo,
|
||||
ModelResponseStream,
|
||||
StandardLoggingPayload,
|
||||
|
|
@ -3193,7 +3194,7 @@ class Router:
|
|||
function_name=function_name,
|
||||
)
|
||||
model_group: Final = kwargs.get(metadata_variable_name, {}).get("model_group")
|
||||
_model_id: Final = self._generate_model_id(model_group=model_group, litellm_params=dynamic_litellm_params)
|
||||
_model_id: Final = self.generate_model_id(model_group=model_group, litellm_params=dynamic_litellm_params)
|
||||
original_model_id: Final = model_info.get("id")
|
||||
model_info["id"] = _model_id
|
||||
model_info["original_model_id"] = original_model_id
|
||||
|
|
@ -5087,6 +5088,13 @@ class Router:
|
|||
)
|
||||
|
||||
kwargs_copy["file"] = file
|
||||
if custom_llm_provider == LlmProviders.LITELLM_PROXY.value:
|
||||
kwargs_copy["extra_body"] = MappingProxyType(
|
||||
{
|
||||
**(kwargs_copy.get("extra_body") or MappingProxyType({})),
|
||||
"target_model_names": stripped_model,
|
||||
}
|
||||
)
|
||||
if (
|
||||
"gcs_bucket_name" in data
|
||||
): # TODO: Remove this once we have a better way to handle GCS bucket name: Problem is that we need to pass the gcs_bucket_name to the router for the create_file call but it doesn't show up there
|
||||
|
|
@ -7570,13 +7578,14 @@ class Router:
|
|||
|
||||
@staticmethod
|
||||
def _json_default_stable_id(value: object) -> str:
|
||||
"""json.dumps default= for _generate_model_id: plain str() on an arbitrary
|
||||
"""json.dumps default= for generate_model_id: plain str() on an arbitrary
|
||||
object (e.g. a RoutingPlugin instance) falls back to object.__repr__'s
|
||||
`<module.Class object at 0x...>`, so the hash -- and deployment id -- would
|
||||
change every restart. Use the class name instead, stable across restarts."""
|
||||
return f"{type(value).__module__}.{type(value).__qualname__}"
|
||||
|
||||
def _generate_model_id(self, model_group: str, litellm_params: dict):
|
||||
@staticmethod
|
||||
def generate_model_id(model_group: str, litellm_params: dict) -> str: # mutable-ok: hashed read-only
|
||||
"""
|
||||
Helper function to consistently generate the same id for a deployment
|
||||
|
||||
|
|
@ -7591,14 +7600,14 @@ class Router:
|
|||
if isinstance(k, str):
|
||||
parts.append(k)
|
||||
elif isinstance(k, dict):
|
||||
parts.append(json.dumps(k, default=self._json_default_stable_id))
|
||||
parts.append(json.dumps(k, default=Router._json_default_stable_id))
|
||||
else:
|
||||
parts.append(str(k))
|
||||
|
||||
if isinstance(v, str):
|
||||
parts.append(v)
|
||||
elif isinstance(v, dict):
|
||||
parts.append(json.dumps(v, default=self._json_default_stable_id))
|
||||
parts.append(json.dumps(v, default=Router._json_default_stable_id))
|
||||
else:
|
||||
parts.append(str(v))
|
||||
|
||||
|
|
@ -8192,7 +8201,7 @@ class Router:
|
|||
|
||||
# check if model info has id
|
||||
if "id" not in _model_info:
|
||||
_id = self._generate_model_id(_model_name, _litellm_params)
|
||||
_id = self.generate_model_id(_model_name, _litellm_params)
|
||||
_model_info["id"] = _id
|
||||
|
||||
if _litellm_params.get("organization", None) is not None and isinstance(
|
||||
|
|
@ -9750,7 +9759,7 @@ class Router:
|
|||
if model_id is None:
|
||||
model_name = model.get("model_name", "")
|
||||
litellm_params = model.get("litellm_params", {})
|
||||
model_id = self._generate_model_id(model_name, litellm_params)
|
||||
model_id = self.generate_model_id(model_name, litellm_params)
|
||||
# Update the model_info in the original list
|
||||
if "model_info" not in model:
|
||||
model["model_info"] = {}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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": []}}
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import httpx
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
|
||||
def _expected_dir() -> Path:
|
||||
|
|
@ -367,3 +368,57 @@ async def test_aresponses_client_header_conflict_is_case_insensitive():
|
|||
|
||||
assert [name for name in request_headers if name.lower() == "x-shared"] == ["x-shared"]
|
||||
assert request_headers["x-shared"] == "from-caller"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("model", "custom_llm_provider"),
|
||||
[
|
||||
("openai/responses/gpt-5.6", None),
|
||||
("responses/gpt-5.6", "openai"),
|
||||
],
|
||||
)
|
||||
async def test_aresponses_strips_responses_routing_prefix_from_openai_model(model, custom_llm_provider):
|
||||
"""
|
||||
`responses/` is LiteLLM routing sugar, never part of the provider model id.
|
||||
Deployments configured as openai/responses/<model> reach this path directly via
|
||||
/v1/responses and via the /v1/messages adapter (which passes responses/<model>
|
||||
with custom_llm_provider="openai"), so both shapes must hit OpenAI as <model>.
|
||||
"""
|
||||
injected_client = AsyncHTTPHandler()
|
||||
mock_post = AsyncMock(return_value=MockResponse(_minimal_responses_api_payload("resp_prefix_test", "gpt-5.6"), 200))
|
||||
injected_client.post = mock_post
|
||||
|
||||
await litellm.aresponses(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
input="ping",
|
||||
api_key="sk-test",
|
||||
client=injected_client,
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
assert mock_post.call_args.kwargs["url"].endswith("/responses")
|
||||
assert mock_post.call_args.kwargs["json"]["model"] == "gpt-5.6"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aresponses_websocket_strips_responses_routing_prefix_from_openai_model():
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.responses.main import _aresponses_websocket
|
||||
|
||||
with patch(
|
||||
"litellm.responses.main.base_llm_http_handler.async_responses_websocket",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_ws:
|
||||
await _aresponses_websocket(
|
||||
model="openai/responses/gpt-5.6",
|
||||
websocket=MagicMock(),
|
||||
api_key="sk-test",
|
||||
litellm_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
mock_ws.assert_awaited_once()
|
||||
assert mock_ws.call_args.kwargs["model"] == "gpt-5.6"
|
||||
assert mock_ws.call_args.kwargs["custom_llm_provider"] == "openai"
|
||||
|
|
|
|||
|
|
@ -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 == ""
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -221,7 +221,7 @@ def test_filter_by_routing_plugin_candidates_narrows_and_raises_when_empty():
|
|||
|
||||
|
||||
def test_json_default_stable_id_is_stable_across_instances():
|
||||
"""_generate_model_id's json.dumps `default=` fallback must not embed an object's
|
||||
"""generate_model_id's json.dumps `default=` fallback must not embed an object's
|
||||
memory address (e.g. plain str() on an object with no custom __repr__ falls back
|
||||
to object.__repr__'s `<module.Class object at 0x...>`) -- that would make the
|
||||
deployment id churn on every process restart for any deployment whose
|
||||
|
|
@ -232,7 +232,7 @@ def test_json_default_stable_id_is_stable_across_instances():
|
|||
assert router._json_default_stable_id(LanguageDetector()) != router._json_default_stable_id(TenantPolicy())
|
||||
|
||||
|
||||
def test_generate_model_id_is_stable_when_litellm_params_contain_a_plugin_instance():
|
||||
def testgenerate_model_id_is_stable_when_litellm_params_contain_a_plugin_instance():
|
||||
"""End-to-end: a deployment id built from litellm_params containing a routing
|
||||
plugin instance (e.g. complexity_router_config.plugins) must be identical across
|
||||
separate calls, not just non-crashing."""
|
||||
|
|
@ -242,8 +242,8 @@ def test_generate_model_id_is_stable_when_litellm_params_contain_a_plugin_instan
|
|||
"complexity_router_config": {"plugins": [LanguageDetector()]},
|
||||
}
|
||||
|
||||
id1 = router._generate_model_id("smart-router", litellm_params)
|
||||
id2 = router._generate_model_id(
|
||||
id1 = router.generate_model_id("smart-router", litellm_params)
|
||||
id2 = router.generate_model_id(
|
||||
"smart-router",
|
||||
{
|
||||
"model": "auto_router/complexity_router",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
6
ui/litellm-dashboard/public/assets/logos/valkey.svg
Normal file
6
ui/litellm-dashboard/public/assets/logos/valkey.svg
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<svg width="64" height="73" viewBox="0 0 64 73" xmlns="http://www.w3.org/2000/svg">
|
||||
<g id="Group-copy">
|
||||
<path id="Path" fill="#123678" fill-rule="evenodd" stroke="none" d="M 13.482285 60.694962 L 0.998384 52.884399 L 0.998384 19.502914 L 31.527868 2.001205 L 61.317604 19.532024 L 61.317604 54.64489 L 31.054855 71.68927 L 20.548372 65.115807 L 20.548372 51.041328 L 20.548372 49.119896 L 14.851504 45.555508 L 14.851504 27.453159 L 31.346497 17.99712 L 47.464485 27.482262 L 47.464485 46.451157 L 34.703495 53.638138 L 34.703495 45.998573 C 38.52874 44.52552 41.274452 40.739189 41.274452 36.270489 C 41.274452 30.510658 36.712814 25.88438 31.158138 25.88438 C 25.603172 25.88438 21.041817 30.510658 21.041817 36.270489 C 21.041817 40.739189 23.787249 44.52552 27.612494 45.998573 L 27.612494 60.473576 L 31.261133 62.756348 L 53.635483 50.15464 L 53.635483 23.924595 L 31.477489 10.884869 L 8.680504 23.953705 L 8.680504 48.628967 L 13.482285 51.633297 L 13.482285 60.694962 Z M 31.158138 31.498383 C 33.671822 31.498383 35.660439 33.664162 35.660439 36.270489 C 35.660439 38.876804 33.671822 41.042587 31.158138 41.042587 C 28.644447 41.042587 26.655558 38.876804 26.655558 36.270489 C 26.655558 33.664162 28.644447 31.498383 31.158138 31.498383 Z" />
|
||||
</g>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 1.3 KiB |
|
|
@ -1,147 +1,179 @@
|
|||
"use client";
|
||||
|
||||
import { BotIcon, InfoIcon, LayersIcon, ServerIcon } from "lucide-react";
|
||||
import type { UseFormReturn } from "react-hook-form";
|
||||
import { z } from "zod/v4";
|
||||
|
||||
import { useAgents } from "@/app/(dashboard)/hooks/agents/useAgents";
|
||||
import { useMCPServers } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers";
|
||||
import { ModelSelect } from "@/components/ModelSelect/ModelSelect";
|
||||
import type { FormInstance } from "antd";
|
||||
import { Form, Input, Select, Space, Tabs } from "antd";
|
||||
import { BotIcon, InfoIcon, LayersIcon, ServerIcon } from "lucide-react";
|
||||
import { FieldGroup } from "@/components/shared/form/field";
|
||||
import { FormField } from "@/components/shared/form/FormField";
|
||||
import { Input } from "@/components/ui/input";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
|
||||
import { Textarea } from "@/components/ui/textarea";
|
||||
|
||||
const { TextArea } = Input;
|
||||
export const accessGroupFormSchema = z.object({
|
||||
name: z.string().min(1, "Please enter the access group name"),
|
||||
description: z.string(),
|
||||
modelIds: z.array(z.string()),
|
||||
mcpServerIds: z.array(z.string()),
|
||||
agentIds: z.array(z.string()),
|
||||
});
|
||||
|
||||
export interface AccessGroupFormValues {
|
||||
name: string;
|
||||
description: string;
|
||||
modelIds: string[];
|
||||
mcpServerIds: string[];
|
||||
agentIds: string[];
|
||||
export type AccessGroupFormValues = z.output<typeof accessGroupFormSchema>;
|
||||
|
||||
export const GENERAL_TAB = "general";
|
||||
export const MODELS_TAB = "models";
|
||||
export const MCP_SERVERS_TAB = "mcp-servers";
|
||||
export const AGENTS_TAB = "agents";
|
||||
|
||||
interface MultiSelectOption {
|
||||
value: string;
|
||||
label: string;
|
||||
}
|
||||
|
||||
interface MultiSelectProps {
|
||||
id: string;
|
||||
value: string[];
|
||||
onChange: (value: string[]) => void;
|
||||
options: MultiSelectOption[];
|
||||
placeholder: string;
|
||||
"aria-invalid": true | undefined;
|
||||
"aria-describedby": string | undefined;
|
||||
}
|
||||
|
||||
const MultiSelect = ({
|
||||
id,
|
||||
value,
|
||||
onChange,
|
||||
options,
|
||||
placeholder,
|
||||
"aria-invalid": ariaInvalid,
|
||||
"aria-describedby": ariaDescribedBy,
|
||||
}: MultiSelectProps) => (
|
||||
<Select multiple items={options} value={value} onValueChange={onChange}>
|
||||
<SelectTrigger id={id} aria-invalid={ariaInvalid} aria-describedby={ariaDescribedBy} className="w-full">
|
||||
<SelectValue placeholder={placeholder}>
|
||||
{(selected: string[]) =>
|
||||
selected.length === 0
|
||||
? placeholder
|
||||
: options
|
||||
.filter((option) => selected.includes(option.value))
|
||||
.map((option) => option.label)
|
||||
.join(", ")
|
||||
}
|
||||
</SelectValue>
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{options.map((option) => (
|
||||
<SelectItem key={option.value} value={option.value} title={option.label}>
|
||||
{option.label}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
);
|
||||
|
||||
interface AccessGroupBaseFormProps {
|
||||
form: FormInstance<AccessGroupFormValues>;
|
||||
form: UseFormReturn<AccessGroupFormValues>;
|
||||
isNameDisabled?: boolean;
|
||||
activeTab: string;
|
||||
onTabChange: (tab: string) => void;
|
||||
}
|
||||
|
||||
export function AccessGroupBaseForm({ form, isNameDisabled = false }: AccessGroupBaseFormProps) {
|
||||
export function AccessGroupBaseForm({
|
||||
form,
|
||||
isNameDisabled = false,
|
||||
activeTab,
|
||||
onTabChange,
|
||||
}: AccessGroupBaseFormProps) {
|
||||
const { data: agentsData } = useAgents();
|
||||
const { data: mcpServersData } = useMCPServers();
|
||||
|
||||
const agents = agentsData?.agents ?? [];
|
||||
const mcpServers = mcpServersData ?? [];
|
||||
const items = [
|
||||
{
|
||||
key: "1",
|
||||
label: (
|
||||
<Space align="center" size={4}>
|
||||
<InfoIcon size={16} />
|
||||
General Info
|
||||
</Space>
|
||||
),
|
||||
children: (
|
||||
<div style={{ paddingTop: 16 }}>
|
||||
<Form.Item
|
||||
name="name"
|
||||
label="Group Name"
|
||||
rules={[
|
||||
{
|
||||
required: true,
|
||||
message: "Please enter the access group name",
|
||||
},
|
||||
]}
|
||||
>
|
||||
<Input placeholder="e.g. Engineering Team" disabled={isNameDisabled} />
|
||||
</Form.Item>
|
||||
<Form.Item name="description" label="Description">
|
||||
<TextArea rows={4} placeholder="Describe the purpose of this access group..." />
|
||||
</Form.Item>
|
||||
</div>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "2",
|
||||
label: (
|
||||
<Space align="center" size={4}>
|
||||
<LayersIcon size={16} />
|
||||
Models
|
||||
</Space>
|
||||
),
|
||||
children: (
|
||||
<div style={{ paddingTop: 16 }}>
|
||||
<Form.Item name="modelIds" label="Allowed Models">
|
||||
<ModelSelect
|
||||
context="global"
|
||||
value={form.getFieldValue("modelIds") ?? []}
|
||||
onChange={(values) => form.setFieldsValue({ modelIds: values })}
|
||||
style={{ width: "100%" }}
|
||||
/>
|
||||
</Form.Item>
|
||||
</div>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "3",
|
||||
label: (
|
||||
<Space align="center" size={4}>
|
||||
<ServerIcon size={16} />
|
||||
MCP Servers
|
||||
</Space>
|
||||
),
|
||||
children: (
|
||||
<div style={{ paddingTop: 16 }}>
|
||||
<Form.Item name="mcpServerIds" label="Allowed MCP Servers">
|
||||
<Select
|
||||
mode="multiple"
|
||||
placeholder="Select MCP servers"
|
||||
style={{ width: "100%" }}
|
||||
optionFilterProp="label"
|
||||
allowClear
|
||||
options={mcpServers.map((server) => ({
|
||||
label: server.server_name ?? server.server_id,
|
||||
value: server.server_id,
|
||||
}))}
|
||||
/>
|
||||
</Form.Item>
|
||||
</div>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "4",
|
||||
label: (
|
||||
<Space align="center" size={4}>
|
||||
<BotIcon size={16} />
|
||||
Agents
|
||||
</Space>
|
||||
),
|
||||
children: (
|
||||
<div style={{ paddingTop: 16 }}>
|
||||
<Form.Item name="agentIds" label="Allowed Agents">
|
||||
<Select
|
||||
mode="multiple"
|
||||
placeholder="Select agents"
|
||||
style={{ width: "100%" }}
|
||||
optionFilterProp="label"
|
||||
allowClear
|
||||
options={agents.map((agent) => ({
|
||||
label: agent.agent_name,
|
||||
value: agent.agent_id,
|
||||
}))}
|
||||
/>
|
||||
</Form.Item>
|
||||
</div>
|
||||
),
|
||||
},
|
||||
];
|
||||
const mcpServerOptions = (mcpServersData ?? []).map((server) => ({
|
||||
value: server.server_id,
|
||||
label: server.server_name ?? server.server_id,
|
||||
}));
|
||||
const agentOptions = (agentsData?.agents ?? []).map((agent) => ({
|
||||
value: agent.agent_id,
|
||||
label: agent.agent_name,
|
||||
}));
|
||||
|
||||
return (
|
||||
<Form
|
||||
form={form}
|
||||
layout="vertical"
|
||||
name="access_group_form"
|
||||
initialValues={{
|
||||
modelIds: [],
|
||||
mcpServerIds: [],
|
||||
agentIds: [],
|
||||
}}
|
||||
>
|
||||
<Tabs defaultActiveKey="1" items={items} />
|
||||
</Form>
|
||||
<Tabs value={activeTab} onValueChange={onTabChange}>
|
||||
<TabsList className="w-full">
|
||||
<TabsTrigger value={GENERAL_TAB}>
|
||||
<InfoIcon size={16} />
|
||||
General Info
|
||||
</TabsTrigger>
|
||||
<TabsTrigger value={MODELS_TAB}>
|
||||
<LayersIcon size={16} />
|
||||
Models
|
||||
</TabsTrigger>
|
||||
<TabsTrigger value={MCP_SERVERS_TAB}>
|
||||
<ServerIcon size={16} />
|
||||
MCP Servers
|
||||
</TabsTrigger>
|
||||
<TabsTrigger value={AGENTS_TAB}>
|
||||
<BotIcon size={16} />
|
||||
Agents
|
||||
</TabsTrigger>
|
||||
</TabsList>
|
||||
|
||||
<TabsContent value={GENERAL_TAB} className="pt-4">
|
||||
<FieldGroup>
|
||||
<FormField control={form.control} name="name" label="Group Name">
|
||||
{({ ref, ...field }) => (
|
||||
<Input {...field} ref={ref} placeholder="e.g. Engineering Team" disabled={isNameDisabled} />
|
||||
)}
|
||||
</FormField>
|
||||
<FormField control={form.control} name="description" label="Description">
|
||||
{({ ref, ...field }) => (
|
||||
<Textarea {...field} ref={ref} rows={4} placeholder="Describe the purpose of this access group..." />
|
||||
)}
|
||||
</FormField>
|
||||
</FieldGroup>
|
||||
</TabsContent>
|
||||
|
||||
<TabsContent value={MODELS_TAB} className="pt-4">
|
||||
<FormField control={form.control} name="modelIds" label="Allowed Models">
|
||||
{(field) => <ModelSelect context="global" value={field.value} onChange={field.onChange} />}
|
||||
</FormField>
|
||||
</TabsContent>
|
||||
|
||||
<TabsContent value={MCP_SERVERS_TAB} className="pt-4">
|
||||
<FormField control={form.control} name="mcpServerIds" label="Allowed MCP Servers">
|
||||
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
|
||||
<MultiSelect
|
||||
id={id}
|
||||
value={value}
|
||||
onChange={onChange}
|
||||
options={mcpServerOptions}
|
||||
placeholder="Select MCP servers"
|
||||
aria-invalid={ariaInvalid}
|
||||
aria-describedby={ariaDescribedBy}
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
</TabsContent>
|
||||
|
||||
<TabsContent value={AGENTS_TAB} className="pt-4">
|
||||
<FormField control={form.control} name="agentIds" label="Allowed Agents">
|
||||
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
|
||||
<MultiSelect
|
||||
id={id}
|
||||
value={value}
|
||||
onChange={onChange}
|
||||
options={agentOptions}
|
||||
placeholder="Select agents"
|
||||
aria-invalid={ariaInvalid}
|
||||
aria-describedby={ariaDescribedBy}
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
</TabsContent>
|
||||
</Tabs>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,175 @@
|
|||
import { describe, it, expect, vi, beforeEach } from "vitest";
|
||||
import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event";
|
||||
import { renderWithProviders, screen, waitFor } from "../../../../../../tests/test-utils";
|
||||
import { AccessGroupEditModal } from "./AccessGroupEditModal";
|
||||
import { AccessGroupResponse } from "@/app/(dashboard)/hooks/accessGroups/useAccessGroups";
|
||||
|
||||
const mutate = vi.fn();
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/accessGroups/useEditAccessGroup", () => ({
|
||||
useEditAccessGroup: () => ({ mutate, isPending: false }),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/agents/useAgents", () => ({
|
||||
useAgents: () => ({ data: { agents: [{ agent_id: "agent-1", agent_name: "Support Bot" }] } }),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/mcpServers/useMCPServers", () => ({
|
||||
useMCPServers: () => ({ data: [{ server_id: "srv-1", server_name: "Files" }] }),
|
||||
}));
|
||||
|
||||
vi.mock("@/components/ModelSelect/ModelSelect", () => ({
|
||||
ModelSelect: ({ value, onChange }: { value: string[]; onChange: (next: string[]) => void }) => (
|
||||
<button type="button" aria-label="model-select" onClick={() => onChange([...(value ?? []), "gpt-4"])}>
|
||||
{(value ?? []).join(",")}
|
||||
</button>
|
||||
),
|
||||
}));
|
||||
|
||||
vi.mock("@/lib/toast", () => ({
|
||||
toast: { success: vi.fn(), fromError: vi.fn(), error: vi.fn() },
|
||||
}));
|
||||
|
||||
const setup = () => userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
|
||||
type User = ReturnType<typeof setup>;
|
||||
|
||||
const accessGroup: AccessGroupResponse = {
|
||||
access_group_id: "ag-1",
|
||||
access_group_name: "Engineering",
|
||||
description: "Engineers",
|
||||
access_model_names: ["gpt-4"],
|
||||
access_mcp_server_ids: ["srv-1"],
|
||||
access_agent_ids: ["agent-1"],
|
||||
assigned_team_ids: [],
|
||||
assigned_key_ids: [],
|
||||
created_at: "2024-01-01T00:00:00Z",
|
||||
created_by: "user-1",
|
||||
updated_at: "2024-01-02T00:00:00Z",
|
||||
updated_by: "user-1",
|
||||
};
|
||||
|
||||
const renderModal = (data: AccessGroupResponse = accessGroup) =>
|
||||
renderWithProviders(<AccessGroupEditModal visible accessGroup={data} onCancel={vi.fn()} />);
|
||||
|
||||
const save = async (user: User) => user.click(screen.getByRole("button", { name: "Save Changes" }));
|
||||
|
||||
const variables = () => mutate.mock.calls.at(-1)?.[0] as { accessGroupId: string; params: Record<string, unknown> };
|
||||
|
||||
describe("AccessGroupEditModal submit payload", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("sends exactly the antd payload for an untouched save", async () => {
|
||||
const user = setup();
|
||||
renderModal();
|
||||
await screen.findByDisplayValue("Engineering");
|
||||
|
||||
await save(user);
|
||||
|
||||
await waitFor(() => expect(mutate).toHaveBeenCalled());
|
||||
|
||||
expect(mutate).toHaveBeenCalledTimes(1);
|
||||
expect(variables().accessGroupId).toBe("ag-1");
|
||||
expect(variables().params).toStrictEqual({
|
||||
access_group_name: "Engineering",
|
||||
description: "Engineers",
|
||||
access_model_names: undefined,
|
||||
access_mcp_server_ids: undefined,
|
||||
access_agent_ids: undefined,
|
||||
});
|
||||
});
|
||||
|
||||
it("sends a tab's field only once that tab has been visited", async () => {
|
||||
const user = setup();
|
||||
renderModal();
|
||||
await screen.findByDisplayValue("Engineering");
|
||||
|
||||
await user.click(screen.getByRole("tab", { name: /MCP Servers/ }));
|
||||
await user.click(screen.getByRole("tab", { name: /Agents/ }));
|
||||
await user.click(screen.getByRole("tab", { name: /General Info/ }));
|
||||
await save(user);
|
||||
|
||||
await waitFor(() => expect(mutate).toHaveBeenCalled());
|
||||
expect(variables().params).toStrictEqual({
|
||||
access_group_name: "Engineering",
|
||||
description: "Engineers",
|
||||
access_model_names: undefined,
|
||||
access_mcp_server_ids: ["srv-1"],
|
||||
access_agent_ids: ["agent-1"],
|
||||
});
|
||||
});
|
||||
|
||||
it("coerces a null description to an empty string", async () => {
|
||||
const user = setup();
|
||||
renderModal({ ...accessGroup, description: null });
|
||||
await screen.findByDisplayValue("Engineering");
|
||||
|
||||
await save(user);
|
||||
|
||||
await waitFor(() => expect(mutate).toHaveBeenCalled());
|
||||
expect(variables().params.description).toBe("");
|
||||
});
|
||||
|
||||
it("does not trim surrounding whitespace from the group name", async () => {
|
||||
const user = setup();
|
||||
renderModal();
|
||||
const nameInput = await screen.findByDisplayValue("Engineering");
|
||||
|
||||
await user.clear(nameInput);
|
||||
await user.type(nameInput, " Padded ");
|
||||
await save(user);
|
||||
|
||||
await waitFor(() => expect(mutate).toHaveBeenCalled());
|
||||
expect(variables().params.access_group_name).toBe(" Padded ");
|
||||
});
|
||||
|
||||
it("never sends server-only fields from the loaded record", async () => {
|
||||
const user = setup();
|
||||
renderModal();
|
||||
await screen.findByDisplayValue("Engineering");
|
||||
|
||||
await save(user);
|
||||
|
||||
await waitFor(() => expect(mutate).toHaveBeenCalled());
|
||||
expect(variables().params).not.toHaveProperty("access_group_id");
|
||||
expect(variables().params).not.toHaveProperty("created_at");
|
||||
expect(variables().params).not.toHaveProperty("assigned_team_ids");
|
||||
expect(variables().params).not.toHaveProperty("assigned_key_ids");
|
||||
});
|
||||
|
||||
it("does not submit when the group name is cleared", async () => {
|
||||
const user = setup();
|
||||
renderModal();
|
||||
const nameInput = await screen.findByDisplayValue("Engineering");
|
||||
|
||||
await user.clear(nameInput);
|
||||
await save(user);
|
||||
|
||||
expect(await screen.findByText("Please enter the access group name")).toBeInTheDocument();
|
||||
expect(mutate).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("does not save when Enter is pressed in the name field", async () => {
|
||||
const user = setup();
|
||||
renderModal();
|
||||
const nameInput = await screen.findByDisplayValue("Engineering");
|
||||
|
||||
await user.type(nameInput, "{Enter}");
|
||||
|
||||
expect(mutate).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("sends models chosen on the Models tab", async () => {
|
||||
const user = setup();
|
||||
renderModal({ ...accessGroup, access_model_names: [] });
|
||||
await screen.findByDisplayValue("Engineering");
|
||||
|
||||
await user.click(screen.getByRole("tab", { name: /Models/ }));
|
||||
await user.click(await screen.findByLabelText("model-select"));
|
||||
await save(user);
|
||||
|
||||
await waitFor(() => expect(mutate).toHaveBeenCalled());
|
||||
expect(variables().params.access_model_names).toStrictEqual(["gpt-4"]);
|
||||
});
|
||||
});
|
||||
|
|
@ -1,10 +1,24 @@
|
|||
import React, { useEffect } from "react";
|
||||
import { Modal, Form } from "antd";
|
||||
"use client";
|
||||
|
||||
import React, { useState } from "react";
|
||||
import { Modal } from "antd";
|
||||
|
||||
import { toast } from "@/lib/toast";
|
||||
import { AccessGroupBaseForm, AccessGroupFormValues } from "./AccessGroupBaseForm";
|
||||
import { useZodForm } from "@/lib/forms/useZodForm";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { useEditAccessGroup, AccessGroupUpdateParams } from "@/app/(dashboard)/hooks/accessGroups/useEditAccessGroup";
|
||||
import { AccessGroupResponse } from "@/app/(dashboard)/hooks/accessGroups/useAccessGroups";
|
||||
|
||||
import {
|
||||
AccessGroupBaseForm,
|
||||
accessGroupFormSchema,
|
||||
AGENTS_TAB,
|
||||
GENERAL_TAB,
|
||||
MCP_SERVERS_TAB,
|
||||
MODELS_TAB,
|
||||
type AccessGroupFormValues,
|
||||
} from "./AccessGroupBaseForm";
|
||||
|
||||
interface AccessGroupEditModalProps {
|
||||
visible: boolean;
|
||||
accessGroup: AccessGroupResponse;
|
||||
|
|
@ -12,62 +26,74 @@ interface AccessGroupEditModalProps {
|
|||
onSuccess?: () => void;
|
||||
}
|
||||
|
||||
export function AccessGroupEditModal({ visible, accessGroup, onCancel, onSuccess }: AccessGroupEditModalProps) {
|
||||
const [form] = Form.useForm<AccessGroupFormValues>();
|
||||
const toFormValues = (accessGroup: AccessGroupResponse): AccessGroupFormValues => ({
|
||||
name: accessGroup.access_group_name,
|
||||
description: accessGroup.description ?? "",
|
||||
modelIds: accessGroup.access_model_names ?? [],
|
||||
mcpServerIds: accessGroup.access_mcp_server_ids ?? [],
|
||||
agentIds: accessGroup.access_agent_ids ?? [],
|
||||
});
|
||||
|
||||
function AccessGroupEditForm({ accessGroup, onCancel, onSuccess }: Omit<AccessGroupEditModalProps, "visible">) {
|
||||
const form = useZodForm(accessGroupFormSchema, { defaultValues: toFormValues(accessGroup) });
|
||||
const editMutation = useEditAccessGroup();
|
||||
const [activeTab, setActiveTab] = useState(GENERAL_TAB);
|
||||
const [visitedTabs, setVisitedTabs] = useState<ReadonlySet<string>>(new Set([GENERAL_TAB]));
|
||||
|
||||
// Populate the form with initial values whenever the modal opens or the data changes
|
||||
useEffect(() => {
|
||||
if (visible && accessGroup) {
|
||||
form.setFieldsValue({
|
||||
name: accessGroup.access_group_name,
|
||||
description: accessGroup.description ?? "",
|
||||
modelIds: accessGroup.access_model_names ?? [],
|
||||
mcpServerIds: accessGroup.access_mcp_server_ids ?? [],
|
||||
agentIds: accessGroup.access_agent_ids ?? [],
|
||||
});
|
||||
}
|
||||
}, [visible, accessGroup, form]);
|
||||
|
||||
const handleOk = () => {
|
||||
form
|
||||
.validateFields()
|
||||
.then((values) => {
|
||||
const params: AccessGroupUpdateParams = {
|
||||
access_group_name: values.name,
|
||||
description: values.description,
|
||||
access_model_names: values.modelIds,
|
||||
access_mcp_server_ids: values.mcpServerIds,
|
||||
access_agent_ids: values.agentIds,
|
||||
};
|
||||
|
||||
editMutation.mutate(
|
||||
{ accessGroupId: accessGroup.access_group_id, params },
|
||||
{
|
||||
onSuccess: () => {
|
||||
toast.success("Access group updated successfully");
|
||||
onSuccess?.();
|
||||
onCancel();
|
||||
},
|
||||
},
|
||||
);
|
||||
})
|
||||
.catch((info) => {});
|
||||
const handleTabChange = (tab: string) => {
|
||||
setActiveTab(tab);
|
||||
setVisitedTabs((previous) => new Set([...previous, tab]));
|
||||
};
|
||||
|
||||
const handleOk = form.handleSubmit(
|
||||
(values) => {
|
||||
const params: AccessGroupUpdateParams = {
|
||||
access_group_name: values.name,
|
||||
description: values.description,
|
||||
access_model_names: visitedTabs.has(MODELS_TAB) ? values.modelIds : undefined,
|
||||
access_mcp_server_ids: visitedTabs.has(MCP_SERVERS_TAB) ? values.mcpServerIds : undefined,
|
||||
access_agent_ids: visitedTabs.has(AGENTS_TAB) ? values.agentIds : undefined,
|
||||
};
|
||||
|
||||
editMutation.mutate(
|
||||
{ accessGroupId: accessGroup.access_group_id, params },
|
||||
{
|
||||
onSuccess: () => {
|
||||
toast.success("Access group updated successfully");
|
||||
onSuccess?.();
|
||||
onCancel();
|
||||
},
|
||||
},
|
||||
);
|
||||
},
|
||||
() => setActiveTab(GENERAL_TAB),
|
||||
);
|
||||
|
||||
return (
|
||||
<Modal
|
||||
title="Edit Access Group"
|
||||
open={visible}
|
||||
onOk={handleOk}
|
||||
onCancel={onCancel}
|
||||
width={700}
|
||||
okText="Save Changes"
|
||||
cancelText="Cancel"
|
||||
confirmLoading={editMutation.isPending}
|
||||
destroyOnHidden
|
||||
>
|
||||
<AccessGroupBaseForm form={form} />
|
||||
<form onSubmit={(event) => event.preventDefault()}>
|
||||
<AccessGroupBaseForm form={form} activeTab={activeTab} onTabChange={handleTabChange} />
|
||||
|
||||
<div className="mt-6 flex justify-end gap-2">
|
||||
<Button type="button" variant="outline" onClick={onCancel} disabled={editMutation.isPending}>
|
||||
Cancel
|
||||
</Button>
|
||||
<Button type="button" onClick={() => void handleOk()} disabled={editMutation.isPending}>
|
||||
Save Changes
|
||||
</Button>
|
||||
</div>
|
||||
</form>
|
||||
);
|
||||
}
|
||||
|
||||
export function AccessGroupEditModal({ visible, accessGroup, onCancel, onSuccess }: AccessGroupEditModalProps) {
|
||||
return (
|
||||
<Modal title="Edit Access Group" open={visible} onCancel={onCancel} width={700} footer={null} destroyOnHidden>
|
||||
<AccessGroupEditForm
|
||||
key={accessGroup.access_group_id}
|
||||
accessGroup={accessGroup}
|
||||
onCancel={onCancel}
|
||||
onSuccess={onSuccess}
|
||||
/>
|
||||
</Modal>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
import { render, screen, waitFor } from "@testing-library/react";
|
||||
import { render, screen, waitFor, within } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import AdminPanel from "./AdminPanel";
|
||||
|
|
@ -323,3 +323,73 @@ describe("AdminPanel", () => {
|
|||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("AdminPanel add allowed IP form", () => {
|
||||
beforeEach(async () => {
|
||||
vi.clearAllMocks();
|
||||
mockUseAuthorized.mockReturnValue({
|
||||
premiumUser: true,
|
||||
accessToken: "test-token",
|
||||
userId: "user-1",
|
||||
});
|
||||
mockGetSSOSettings.mockResolvedValue({ values: {} });
|
||||
mockGetAllowedIPs.mockResolvedValue(["10.0.0.1"]);
|
||||
mockAddAllowedIP.mockResolvedValue({});
|
||||
|
||||
const user = userEvent.setup();
|
||||
render(<AdminPanel />);
|
||||
await user.click(screen.getByRole("tab", { name: /security settings/i }));
|
||||
await user.click(screen.getByRole("button", { name: /allowed ips/i }));
|
||||
const manageDialog = await screen.findByRole("dialog", { name: /manage allowed ip addresses/i });
|
||||
await user.click(within(manageDialog).getByRole("button", { name: /add ip address/i }));
|
||||
await screen.findByPlaceholderText("Enter IP address");
|
||||
});
|
||||
|
||||
const ipField = () => screen.getByPlaceholderText("Enter IP address") as HTMLInputElement;
|
||||
|
||||
const submitAddIP = async (user: ReturnType<typeof userEvent.setup>) => {
|
||||
const addIpForm = ipField().form as HTMLFormElement;
|
||||
await user.click(within(addIpForm).getByText("Add IP Address"));
|
||||
};
|
||||
|
||||
it("sends the access token and the typed IP address", async () => {
|
||||
const user = userEvent.setup();
|
||||
|
||||
await user.type(ipField(), "192.168.1.50");
|
||||
await submitAddIP(user);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockAddAllowedIP).toHaveBeenCalledWith("test-token", "192.168.1.50");
|
||||
});
|
||||
expect(mockAddAllowedIP).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("blocks the submit and shows the required message when no IP is typed", async () => {
|
||||
const user = userEvent.setup();
|
||||
|
||||
await submitAddIP(user);
|
||||
|
||||
expect(await screen.findByText("Please enter an IP address")).toBeInTheDocument();
|
||||
expect(mockAddAllowedIP).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("submits on Enter from the IP field", async () => {
|
||||
const user = userEvent.setup();
|
||||
|
||||
await user.type(ipField(), "172.16.0.9{Enter}");
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockAddAllowedIP).toHaveBeenCalledWith("test-token", "172.16.0.9");
|
||||
});
|
||||
});
|
||||
|
||||
it("refreshes the allowed IP list after a successful add", async () => {
|
||||
const user = userEvent.setup();
|
||||
mockGetAllowedIPs.mockResolvedValue(["10.0.0.1", "192.168.1.50"]);
|
||||
|
||||
await user.type(ipField(), "192.168.1.50");
|
||||
await submitAddIP(user);
|
||||
|
||||
expect(await screen.findByText("192.168.1.50")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -3,18 +3,12 @@
|
|||
* Use this to avoid sharing master key with others
|
||||
*/
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import {
|
||||
Button,
|
||||
Callout,
|
||||
Card,
|
||||
Table,
|
||||
TableBody,
|
||||
TableCell,
|
||||
TableHead,
|
||||
TableHeaderCell,
|
||||
TableRow,
|
||||
} from "@tremor/react";
|
||||
import { Alert, Button as Button2, Form, Input, Modal, Space, Tabs, Typography } from "antd";
|
||||
import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Card } from "@/components/ui/card";
|
||||
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
|
||||
import { Alert as AntdAlert, Modal, Space, Tabs, Typography } from "antd";
|
||||
import { Info } from "lucide-react";
|
||||
import React, { useEffect, useState } from "react";
|
||||
import NewBadge from "@/components/common_components/NewBadge";
|
||||
import { useBaseUrl } from "@/components/constants";
|
||||
|
|
@ -28,17 +22,50 @@ import UserBannerSettings from "@/components/Settings/AdminSettings/UserBannerSe
|
|||
import HashicorpVault from "@/components/Settings/AdminSettings/HashicorpVault/HashicorpVault";
|
||||
import PluginSettings from "@/components/Settings/AdminSettings/PluginSettings/PluginSettings";
|
||||
import SSOModals from "@/components/SSOModals";
|
||||
import {
|
||||
emptySSOSettingsFormValues,
|
||||
useSSOSettingsForm,
|
||||
type SSOSettingsFormValues,
|
||||
} from "@/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm";
|
||||
import UIAccessControlForm from "@/components/UIAccessControlForm";
|
||||
import { z } from "zod/v4";
|
||||
import { FieldGroup } from "@/components/shared/form/field";
|
||||
import { FormField } from "@/components/shared/form/FormField";
|
||||
import { Input } from "@/components/ui/input";
|
||||
import { useZodForm } from "@/lib/forms/useZodForm";
|
||||
|
||||
const { Title, Paragraph, Text } = Typography;
|
||||
|
||||
const allowedIPSchema = z.object({
|
||||
ip: z.string().min(1, "Please enter an IP address"),
|
||||
});
|
||||
|
||||
type AllowedIPFormValues = z.infer<typeof allowedIPSchema>;
|
||||
|
||||
const AddAllowedIPForm = ({ onSubmit }: { onSubmit: (values: AllowedIPFormValues) => Promise<void> }) => {
|
||||
const form = useZodForm(allowedIPSchema, { defaultValues: { ip: "" } });
|
||||
|
||||
return (
|
||||
<form onSubmit={form.handleSubmit(onSubmit)}>
|
||||
<FieldGroup>
|
||||
<FormField control={form.control} name="ip">
|
||||
{({ ref, ...field }) => <Input ref={ref} placeholder="Enter IP address" {...field} />}
|
||||
</FormField>
|
||||
<div>
|
||||
<Button type="submit">Add IP Address</Button>
|
||||
</div>
|
||||
</FieldGroup>
|
||||
</form>
|
||||
);
|
||||
};
|
||||
|
||||
interface AdminPanelProps {
|
||||
proxySettings?: any;
|
||||
}
|
||||
|
||||
const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
|
||||
const { premiumUser, accessToken, userId: userID } = useAuthorized();
|
||||
const [form] = Form.useForm();
|
||||
const form = useSSOSettingsForm("admin-panel");
|
||||
const [isAddSSOModalVisible, setIsAddSSOModalVisible] = useState(false);
|
||||
const [isInstructionsModalVisible, setIsInstructionsModalVisible] = useState(false);
|
||||
const [isAllowedIPModalVisible, setIsAllowedIPModalVisible] = useState(false);
|
||||
|
|
@ -141,7 +168,7 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
|
|||
|
||||
const handleAddSSOOk = () => {
|
||||
setIsAddSSOModalVisible(false);
|
||||
form.resetFields();
|
||||
form.reset(emptySSOSettingsFormValues);
|
||||
if (accessToken && premiumUser) {
|
||||
checkSSOConfiguration();
|
||||
}
|
||||
|
|
@ -149,10 +176,10 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
|
|||
|
||||
const handleAddSSOCancel = () => {
|
||||
setIsAddSSOModalVisible(false);
|
||||
form.resetFields();
|
||||
form.reset(emptySSOSettingsFormValues);
|
||||
};
|
||||
|
||||
const handleShowInstructions = (formValues: Record<string, any>) => {
|
||||
const handleShowInstructions = (formValues: SSOSettingsFormValues) => {
|
||||
setIsAddSSOModalVisible(false);
|
||||
setIsInstructionsModalVisible(true);
|
||||
};
|
||||
|
|
@ -194,9 +221,9 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
|
|||
label: "Security Settings",
|
||||
children: (
|
||||
<>
|
||||
<Card>
|
||||
<Card className="block p-6">
|
||||
<Title level={4}> ✨ Security Settings</Title>
|
||||
<Alert
|
||||
<AntdAlert
|
||||
message="SSO Configuration Deprecated"
|
||||
description="Editing SSO Settings on this page is deprecated and will be removed in a future version. Please use the SSO Settings tab for SSO configuration."
|
||||
type="warning"
|
||||
|
|
@ -264,19 +291,19 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
|
|||
]}
|
||||
>
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableHeader>
|
||||
<TableRow>
|
||||
<TableHeaderCell>IP Address</TableHeaderCell>
|
||||
<TableHeaderCell className="text-right">Action</TableHeaderCell>
|
||||
<TableHead>IP Address</TableHead>
|
||||
<TableHead className="text-right">Action</TableHead>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
</TableHeader>
|
||||
<TableBody>
|
||||
{allowedIPs.map((ip, index) => (
|
||||
<TableRow key={index}>
|
||||
<TableCell>{ip}</TableCell>
|
||||
<TableCell className="text-right">
|
||||
{ip !== all_ip_address_allowed && (
|
||||
<Button onClick={() => handleDeleteIP(ip)} color="red" size="xs">
|
||||
<Button onClick={() => handleDeleteIP(ip)} variant="destructive" size="sm">
|
||||
Delete
|
||||
</Button>
|
||||
)}
|
||||
|
|
@ -293,14 +320,7 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
|
|||
onCancel={() => setIsAddIPModalVisible(false)}
|
||||
footer={null}
|
||||
>
|
||||
<Form onFinish={handleAddIP}>
|
||||
<Form.Item name="ip" rules={[{ required: true, message: "Please enter an IP address" }]}>
|
||||
<Input placeholder="Enter IP address" />
|
||||
</Form.Item>
|
||||
<Form.Item>
|
||||
<Button2 htmlType="submit">Add IP Address</Button2>
|
||||
</Form.Item>
|
||||
</Form>
|
||||
<AddAllowedIPForm onSubmit={handleAddIP} />
|
||||
</Modal>
|
||||
|
||||
<Modal
|
||||
|
|
@ -338,12 +358,16 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
|
|||
/>
|
||||
</Modal>
|
||||
</div>
|
||||
<Callout title="Login without SSO" color="teal">
|
||||
If you need to login without sso, you can access{" "}
|
||||
<a href={nonSssoUrl} target="_blank" rel="noopener noreferrer">
|
||||
<b>{nonSssoUrl}</b>{" "}
|
||||
</a>
|
||||
</Callout>
|
||||
<Alert variant="info">
|
||||
<Info />
|
||||
<AlertTitle>Login without SSO</AlertTitle>
|
||||
<AlertDescription>
|
||||
If you need to login without sso, you can access{" "}
|
||||
<a href={nonSssoUrl} target="_blank" rel="noopener noreferrer">
|
||||
<b>{nonSssoUrl}</b>{" "}
|
||||
</a>
|
||||
</AlertDescription>
|
||||
</Alert>
|
||||
</>
|
||||
),
|
||||
},
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue