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

This commit is contained in:
Yuneng Jiang 2026-05-04 18:22:47 -07:00
commit e35cd5af76
No known key found for this signature in database
135 changed files with 11315 additions and 1524 deletions

View file

@ -11,6 +11,10 @@ from typing import Literal
import litellm
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import (
is_text_content_call_type,
iter_message_text,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm._logging import verbose_proxy_logger
from fastapi import HTTPException
@ -73,10 +77,9 @@ class _ENTERPRISE_BannedKeywords(CustomLogger):
- check if user id part of blocked list
"""
self.print_verbose("Inside Banned Keyword List Pre-Call Hook")
if call_type == "completion" and "messages" in data:
for m in data["messages"]:
if "content" in m and isinstance(m["content"], str):
self.test_violation(test_str=m["content"])
if is_text_content_call_type(call_type):
for text in iter_message_text(data):
self.test_violation(test_str=text)
except HTTPException as e:
raise e
@ -93,11 +96,16 @@ class _ENTERPRISE_BannedKeywords(CustomLogger):
user_api_key_dict: UserAPIKeyAuth,
response,
):
if isinstance(response, litellm.ModelResponse) and isinstance(
response.choices[0], litellm.utils.Choices
):
for word in self.banned_keywords_list:
self.test_violation(test_str=response.choices[0].message.content or "")
if not isinstance(response, litellm.ModelResponse):
return
for choice in response.choices:
if not isinstance(choice, litellm.utils.Choices):
continue
message = getattr(choice, "message", None)
content = getattr(message, "content", None)
if isinstance(content, str):
self.test_violation(test_str=content)
async def async_post_call_streaming_hook(
self,

View file

@ -12,6 +12,7 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import iter_message_text
from litellm.types.utils import CallTypesLiteral
@ -94,11 +95,9 @@ class _ENTERPRISE_GoogleTextModeration(CustomLogger):
- Calls Google's Text Moderation API
- Rejects request if it fails safety check
"""
if "messages" in data and isinstance(data["messages"], list):
text = ""
for m in data["messages"]: # assume messages is a list
if "content" in m and isinstance(m["content"], str):
text += m["content"]
# Covers multimodal list content + Responses-API input.
text = "".join(iter_message_text(data))
if text:
document = self.language_document(content=text, type_=self.document_type)
request = self.moderate_text_request(

View file

@ -19,6 +19,7 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import iter_message_text
from litellm.types.utils import CallTypesLiteral
@ -37,11 +38,8 @@ class _ENTERPRISE_OpenAI_Moderation(CustomLogger):
user_api_key_dict: UserAPIKeyAuth,
call_type: CallTypesLiteral,
):
text = ""
if "messages" in data and isinstance(data["messages"], list):
for m in data["messages"]: # assume messages is a list
if "content" in m and isinstance(m["content"], str):
text += m["content"]
# Covers multimodal list content + Responses-API input.
text = "".join(iter_message_text(data))
from litellm.proxy.proxy_server import llm_router

View file

@ -18,6 +18,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import walk_user_text
GUARDRAIL_NAME = "hide_secrets"
@ -473,23 +474,19 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
if await self.should_run_check(user_api_key_dict) is False:
return
if "messages" in data and isinstance(data["messages"], list):
for message in data["messages"]:
if "content" in message and isinstance(message["content"], str):
detected_secrets = self.scan_message_for_secrets(message["content"])
# Covers multimodal list content + Responses-API input.
def _redact_message_text(text: str) -> str:
detected_secrets = self.scan_message_for_secrets(text)
for secret in detected_secrets:
text = text.replace(secret["value"], "[REDACTED]")
if detected_secrets:
secret_types = [secret["type"] for secret in detected_secrets]
verbose_proxy_logger.warning(
f"Detected and redacted secrets in message: {secret_types}"
)
return text
for secret in detected_secrets:
message["content"] = message["content"].replace(
secret["value"], "[REDACTED]"
)
if len(detected_secrets) > 0:
secret_types = [secret["type"] for secret in detected_secrets]
verbose_proxy_logger.warning(
f"Detected and redacted secrets in message: {secret_types}"
)
else:
verbose_proxy_logger.debug("No secrets detected on input.")
walk_user_text(data, _redact_message_text)
if "prompt" in data:
if isinstance(data["prompt"], str):
@ -504,11 +501,15 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
f"Detected and redacted secrets in prompt: {secret_types}"
)
elif isinstance(data["prompt"], list):
for item in data["prompt"]:
# Index back into the list — assigning to ``item`` would only
# rebind the loop variable and leave ``data["prompt"]``
# carrying the unredacted secret.
for idx, item in enumerate(data["prompt"]):
if isinstance(item, str):
detected_secrets = self.scan_message_for_secrets(item)
for secret in detected_secrets:
item = item.replace(secret["value"], "[REDACTED]")
data["prompt"][idx] = item
if len(detected_secrets) > 0:
secret_types = [
secret["type"] for secret in detected_secrets
@ -517,31 +518,6 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
f"Detected and redacted secrets in prompt: {secret_types}"
)
if "input" in data:
if isinstance(data["input"], str):
detected_secrets = self.scan_message_for_secrets(data["input"])
for secret in detected_secrets:
data["input"] = data["input"].replace(secret["value"], "[REDACTED]")
if len(detected_secrets) > 0:
secret_types = [secret["type"] for secret in detected_secrets]
verbose_proxy_logger.warning(
f"Detected and redacted secrets in input: {secret_types}"
)
elif isinstance(data["input"], list):
_input_in_request = data["input"]
for idx, item in enumerate(_input_in_request):
if isinstance(item, str):
detected_secrets = self.scan_message_for_secrets(item)
for secret in detected_secrets:
_input_in_request[idx] = item.replace(
secret["value"], "[REDACTED]"
)
if len(detected_secrets) > 0:
secret_types = [
secret["type"] for secret in detected_secrets
]
verbose_proxy_logger.warning(
f"Detected and redacted secrets in input: {secret_types}"
)
verbose_proxy_logger.debug("Data after redacting input %s", data)
# ``data["input"]`` (Responses API and embeddings/moderation) is
# already covered by ``walk_user_text`` above.
return

View file

@ -166,7 +166,7 @@ langfuse_default_tags: Optional[List[str]] = None
langsmith_batch_size: Optional[int] = None
prometheus_initialize_budget_metrics: Optional[bool] = False
prometheus_latency_buckets: Optional[List[float]] = None
require_auth_for_metrics_endpoint: Optional[bool] = False
require_auth_for_metrics_endpoint: Optional[bool] = True
argilla_batch_size: Optional[int] = None
datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged payload.
gcs_pub_sub_use_v1: Optional[bool] = (
@ -280,6 +280,7 @@ ssl_security_level: Optional[str] = None
ssl_certificate: Optional[str] = None
user_url_validation: bool = True
user_url_allowed_hosts: List[str] = []
provider_url_destination_allowed_hosts: List[str] = []
ssl_ecdh_curve: Optional[str] = (
None # Set to 'X25519' to disable PQC and improve performance
)

View file

@ -72,7 +72,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": null,
"effort-2025-11-24": null,
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
"fine-grained-tool-streaming-2025-05-14": null,
@ -103,7 +103,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": null,
"effort-2025-11-24": null,
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
"fine-grained-tool-streaming-2025-05-14": null,

View file

@ -11,17 +11,23 @@ Has 4 methods:
import ast
import asyncio
import json
from typing import Any, cast
import os
from typing import Any, Dict, cast
import litellm
from litellm._logging import print_verbose
from litellm.constants import QDRANT_SCALAR_QUANTILE, QDRANT_VECTOR_SIZE
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_str_from_messages,
)
from litellm.types.utils import EmbeddingResponse
from .base_cache import BaseCache
class QdrantSemanticCache(BaseCache):
CACHE_KEY_FIELD_NAME = "litellm_cache_key"
def __init__( # noqa: PLR0915
self,
qdrant_api_base=None,
@ -33,8 +39,6 @@ class QdrantSemanticCache(BaseCache):
host_type=None,
vector_size=None,
):
import os
from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
get_async_httpx_client,
@ -115,7 +119,9 @@ class QdrantSemanticCache(BaseCache):
print_verbose(
f"Collection already exists.\nCollection details:{self.collection_info}"
)
self._ensure_cache_key_payload_index()
else:
quantization_params: Dict[str, Any]
if quantization_config is None or quantization_config == "binary":
quantization_params = {
"binary": {
@ -156,6 +162,7 @@ class QdrantSemanticCache(BaseCache):
print_verbose(
f"New collection created.\nCollection details:{self.collection_info}"
)
self._ensure_cache_key_payload_index()
else:
raise Exception("Error while creating new collection")
@ -170,15 +177,94 @@ class QdrantSemanticCache(BaseCache):
cached_response = ast.literal_eval(cached_response)
return cached_response
def _get_qdrant_cache_key_filter(self, key: str) -> dict:
return {
"must": [
{
"key": self.CACHE_KEY_FIELD_NAME,
"match": {"value": str(key)},
}
]
}
def _add_cache_key_filter_to_search_data(self, data: dict, key: str) -> None:
data["filter"] = self._get_qdrant_cache_key_filter(key)
def _ensure_cache_key_payload_index(self) -> None:
try:
response = self.sync_client.put(
url=f"{self.qdrant_api_base}/collections/{self.collection_name}/index",
headers=self.headers,
json={
"field_name": self.CACHE_KEY_FIELD_NAME,
"field_schema": "keyword",
},
)
if response.status_code not in (200, 201):
print_verbose(
"Qdrant semantic-cache could not create cache-key payload index: "
f"{response.text}"
)
except Exception as exc:
print_verbose(
"Qdrant semantic-cache could not create cache-key payload index: "
f"{str(exc)}"
)
def _payload_matches_cache_key(self, payload: dict, key: str) -> bool:
# Pre-isolation points stored only prompt + response with no cache-key
# payload field. Reassigning them to a caller's key would risk
# cross-scope hits, so they're treated as misses and re-populated on
# the next set_cache.
cached_key = payload.get(self.CACHE_KEY_FIELD_NAME)
return cached_key is not None and str(cached_key) == str(key)
async def _get_async_embedding(self, prompt: str, **kwargs) -> Any:
llm_model_list = None
llm_router = None
try:
from litellm.proxy.proxy_server import (
llm_model_list as proxy_llm_model_list,
llm_router as proxy_llm_router,
)
llm_model_list = proxy_llm_model_list
llm_router = proxy_llm_router
except ImportError:
pass
router_model_names = (
[m["model_name"] for m in llm_model_list]
if llm_model_list is not None
else []
)
if llm_router is not None and self.embedding_model in router_model_names:
user_api_key = kwargs.get("metadata", {}).get("user_api_key", "")
return await llm_router.aembedding(
model=self.embedding_model,
input=prompt,
cache={"no-store": True, "no-cache": True},
metadata={
"user_api_key": user_api_key,
"semantic-cache-embedding": True,
"trace_id": kwargs.get("metadata", {}).get("trace_id", None),
},
)
return await litellm.aembedding(
model=self.embedding_model,
input=prompt,
cache={"no-store": True, "no-cache": True},
)
def set_cache(self, key, value, **kwargs):
print_verbose(f"qdrant semantic-cache set_cache, kwargs: {kwargs}")
from litellm._uuid import uuid
# get the prompt
messages = kwargs["messages"]
prompt = ""
for message in messages:
prompt += message["content"]
prompt = get_str_from_messages(messages)
# create an embedding for prompt
embedding_response = cast(
@ -202,6 +288,7 @@ class QdrantSemanticCache(BaseCache):
"id": str(uuid.uuid4()),
"vector": embedding,
"payload": {
self.CACHE_KEY_FIELD_NAME: str(key),
"text": prompt,
"response": value,
},
@ -220,9 +307,7 @@ class QdrantSemanticCache(BaseCache):
# get the messages
messages = kwargs["messages"]
prompt = ""
for message in messages:
prompt += message["content"]
prompt = get_str_from_messages(messages)
# convert to embedding
embedding_response = cast(
@ -249,6 +334,7 @@ class QdrantSemanticCache(BaseCache):
"limit": 1,
"with_payload": True,
}
self._add_cache_key_filter_to_search_data(data=data, key=key)
search_response = self.sync_client.post(
url=f"{self.qdrant_api_base}/collections/{self.collection_name}/points/search",
@ -258,21 +344,33 @@ class QdrantSemanticCache(BaseCache):
results = search_response.json()["result"]
if results is None:
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
return None
if isinstance(results, list):
if len(results) == 0:
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
return None
similarity = results[0]["score"]
cached_prompt = results[0]["payload"]["text"]
payload = results[0]["payload"]
if not self._payload_matches_cache_key(payload=payload, key=key):
print_verbose("Qdrant semantic-cache hit did not match cache key scope")
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
return None
cached_prompt = payload["text"]
# check similarity, if more than self.similarity_threshold, return results
print_verbose(
f"semantic cache: similarity threshold: {self.similarity_threshold}, similarity: {similarity}, prompt: {prompt}, closest_cached_prompt: {cached_prompt}"
)
# update kwargs["metadata"] with similarity, don't rewrite the original metadata
kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity
if similarity >= self.similarity_threshold:
# cache hit !
cached_value = results[0]["payload"]["response"]
cached_value = payload["response"]
print_verbose(
f"got a cache hit, similarity: {similarity}, Current prompt: {prompt}, cached_prompt: {cached_prompt}"
)
@ -285,40 +383,12 @@ class QdrantSemanticCache(BaseCache):
async def async_set_cache(self, key, value, **kwargs):
from litellm._uuid import uuid
from litellm.proxy.proxy_server import llm_model_list, llm_router
print_verbose(f"async qdrant semantic-cache set_cache, kwargs: {kwargs}")
# get the prompt
messages = kwargs["messages"]
prompt = ""
for message in messages:
prompt += message["content"]
# create an embedding for prompt
router_model_names = (
[m["model_name"] for m in llm_model_list]
if llm_model_list is not None
else []
)
if llm_router is not None and self.embedding_model in router_model_names:
user_api_key = kwargs.get("metadata", {}).get("user_api_key", "")
embedding_response = await llm_router.aembedding(
model=self.embedding_model,
input=prompt,
cache={"no-store": True, "no-cache": True},
metadata={
"user_api_key": user_api_key,
"semantic-cache-embedding": True,
"trace_id": kwargs.get("metadata", {}).get("trace_id", None),
},
)
else:
# convert to embedding
embedding_response = await litellm.aembedding(
model=self.embedding_model,
input=prompt,
cache={"no-store": True, "no-cache": True},
)
prompt = get_str_from_messages(messages)
embedding_response = await self._get_async_embedding(prompt, **kwargs)
# get the embedding
embedding = embedding_response["data"][0]["embedding"]
@ -332,6 +402,7 @@ class QdrantSemanticCache(BaseCache):
"id": str(uuid.uuid4()),
"vector": embedding,
"payload": {
self.CACHE_KEY_FIELD_NAME: str(key),
"text": prompt,
"response": value,
},
@ -348,38 +419,12 @@ class QdrantSemanticCache(BaseCache):
async def async_get_cache(self, key, **kwargs):
print_verbose(f"async qdrant semantic-cache get_cache, kwargs: {kwargs}")
from litellm.proxy.proxy_server import llm_model_list, llm_router
# get the messages
messages = kwargs["messages"]
prompt = ""
for message in messages:
prompt += message["content"]
prompt = get_str_from_messages(messages)
router_model_names = (
[m["model_name"] for m in llm_model_list]
if llm_model_list is not None
else []
)
if llm_router is not None and self.embedding_model in router_model_names:
user_api_key = kwargs.get("metadata", {}).get("user_api_key", "")
embedding_response = await llm_router.aembedding(
model=self.embedding_model,
input=prompt,
cache={"no-store": True, "no-cache": True},
metadata={
"user_api_key": user_api_key,
"semantic-cache-embedding": True,
"trace_id": kwargs.get("metadata", {}).get("trace_id", None),
},
)
else:
# convert to embedding
embedding_response = await litellm.aembedding(
model=self.embedding_model,
input=prompt,
cache={"no-store": True, "no-cache": True},
)
embedding_response = await self._get_async_embedding(prompt, **kwargs)
# get the embedding
embedding = embedding_response["data"][0]["embedding"]
@ -396,6 +441,7 @@ class QdrantSemanticCache(BaseCache):
"limit": 1,
"with_payload": True,
}
self._add_cache_key_filter_to_search_data(data=data, key=key)
search_response = await self.async_client.post(
url=f"{self.qdrant_api_base}/collections/{self.collection_name}/points/search",
@ -414,7 +460,13 @@ class QdrantSemanticCache(BaseCache):
return None
similarity = results[0]["score"]
cached_prompt = results[0]["payload"]["text"]
payload = results[0]["payload"]
if not self._payload_matches_cache_key(payload=payload, key=key):
print_verbose("Qdrant semantic-cache hit did not match cache key scope")
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
return None
cached_prompt = payload["text"]
# check similarity, if more than self.similarity_threshold, return results
print_verbose(
@ -426,7 +478,7 @@ class QdrantSemanticCache(BaseCache):
if similarity >= self.similarity_threshold:
# cache hit !
cached_value = results[0]["payload"]["response"]
cached_value = payload["response"]
print_verbose(
f"got a cache hit, similarity: {similarity}, Current prompt: {prompt}, cached_prompt: {cached_prompt}"
)

View file

@ -35,6 +35,7 @@ class RedisSemanticCache(BaseCache):
"""
DEFAULT_REDIS_INDEX_NAME: str = "litellm_semantic_cache_index"
CACHE_KEY_FIELD_NAME: str = "litellm_cache_key"
def __init__(
self,
@ -66,8 +67,8 @@ class RedisSemanticCache(BaseCache):
Exception: If similarity_threshold is not provided or required Redis
connection information is missing
"""
from redisvl.extensions.llmcache import SemanticCache
from redisvl.utils.vectorize import CustomTextVectorizer
from redisvl.extensions.llmcache import SemanticCache # type: ignore[import-not-found, import-untyped]
from redisvl.utils.vectorize import CustomTextVectorizer # type: ignore[import-not-found, import-untyped]
if index_name is None:
index_name = self.DEFAULT_REDIS_INDEX_NAME
@ -109,14 +110,94 @@ class RedisSemanticCache(BaseCache):
# Initialize the Redis vectorizer and cache
cache_vectorizer = CustomTextVectorizer(self._get_embedding)
self.llmcache = SemanticCache(
name=index_name,
self.llmcache = self._init_semantic_cache(
semantic_cache_cls=SemanticCache,
index_name=index_name,
redis_url=redis_url,
vectorizer=cache_vectorizer,
distance_threshold=self.distance_threshold,
overwrite=False,
cache_vectorizer=cache_vectorizer,
)
@classmethod
def _cache_key_filterable_field(cls) -> Dict[str, str]:
return {
"name": cls.CACHE_KEY_FIELD_NAME,
"type": "tag",
}
def _init_semantic_cache(
self,
semantic_cache_cls: Any,
index_name: str,
redis_url: str,
cache_vectorizer: Any,
) -> Any:
def _is_schema_mismatch(exc: ValueError) -> bool:
error_message = str(exc).lower()
return any(
phrase in error_message
for phrase in ("schema does not match", "index schema")
)
try:
return semantic_cache_cls(
name=index_name,
redis_url=redis_url,
vectorizer=cache_vectorizer,
distance_threshold=self.distance_threshold,
filterable_fields=[self._cache_key_filterable_field()],
overwrite=False,
)
except ValueError as exc:
if not _is_schema_mismatch(exc):
raise
isolated_index_name = f"{index_name}_isolated"
print_verbose(
"Redis semantic-cache existing index schema is not isolated; "
f"using isolated index - {isolated_index_name}"
)
try:
return semantic_cache_cls(
name=isolated_index_name,
redis_url=redis_url,
vectorizer=cache_vectorizer,
distance_threshold=self.distance_threshold,
filterable_fields=[self._cache_key_filterable_field()],
overwrite=False,
)
except ValueError as isolated_exc:
if not _is_schema_mismatch(isolated_exc):
raise
print_verbose(
"Redis semantic-cache isolated index schema is stale; "
f"recreating isolated index - {isolated_index_name}"
)
return semantic_cache_cls(
name=isolated_index_name,
redis_url=redis_url,
vectorizer=cache_vectorizer,
distance_threshold=self.distance_threshold,
filterable_fields=[self._cache_key_filterable_field()],
overwrite=True,
)
def _get_cache_filters(self, key: str) -> Dict[str, str]:
return {self.CACHE_KEY_FIELD_NAME: str(key)}
def _get_cache_key_filter_expression(self, key: str) -> Any:
from redisvl.query.filter import Tag # type: ignore[import-not-found, import-untyped]
return Tag(self.CACHE_KEY_FIELD_NAME) == str(key)
def _cache_hit_matches_key(self, cache_hit: Dict[str, Any], key: str) -> bool:
# Pre-isolation entries with no ``litellm_cache_key`` field cannot be
# safely reassigned to a caller's scope and are treated as misses.
cached_key = cache_hit.get(self.CACHE_KEY_FIELD_NAME)
if isinstance(cached_key, bytes):
cached_key = cached_key.decode("utf-8")
return cached_key is not None and str(cached_key) == str(key)
def _get_ttl(self, **kwargs) -> Optional[int]:
"""
Get the TTL (time-to-live) value for cache entries.
@ -188,7 +269,7 @@ class RedisSemanticCache(BaseCache):
Store a value in the semantic cache.
Args:
key: The cache key (not directly used in semantic caching)
key: The cache key used to isolate semantic cache entries
value: The response value to cache
**kwargs: Additional arguments including 'messages' for the prompt
and optional 'ttl' for time-to-live
@ -206,12 +287,15 @@ class RedisSemanticCache(BaseCache):
prompt = get_str_from_messages(messages)
value_str = str(value)
store_kwargs: Dict[str, Any] = {
"filters": self._get_cache_filters(key),
}
# Get TTL and store in Redis semantic cache
ttl = self._get_ttl(**kwargs)
if ttl is not None:
self.llmcache.store(prompt, value_str, ttl=int(ttl))
else:
self.llmcache.store(prompt, value_str)
store_kwargs["ttl"] = int(ttl)
self.llmcache.store(prompt, value_str, **store_kwargs)
except Exception as e:
print_verbose(
f"Error setting {value_str or value} in the Redis semantic cache: {str(e)}"
@ -222,7 +306,7 @@ class RedisSemanticCache(BaseCache):
Retrieve a semantically similar cached response.
Args:
key: The cache key (not directly used in semantic caching)
key: The cache key used to isolate semantic cache entries
**kwargs: Additional arguments including 'messages' for the prompt
Returns:
@ -235,18 +319,29 @@ class RedisSemanticCache(BaseCache):
messages = kwargs.get("messages", [])
if not messages:
print_verbose("No messages provided for semantic cache lookup")
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
return None
prompt = get_str_from_messages(messages)
# Check the cache for semantically similar prompts
results = self.llmcache.check(prompt=prompt)
# Check the cache for semantically similar prompts in this exact
# LiteLLM cache-key scope.
check_kwargs: Dict[str, Any] = {
"prompt": prompt,
"filter_expression": self._get_cache_key_filter_expression(key),
}
results = self.llmcache.check(**check_kwargs)
# Return None if no similar prompts found
if not results:
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
return None
# Process the best matching result
cache_hit = results[0]
if not self._cache_hit_matches_key(cache_hit=cache_hit, key=key):
print_verbose("Redis semantic-cache hit did not match cache key scope")
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
return None
vector_distance = float(cache_hit["vector_distance"])
# Convert vector distance back to similarity score
@ -257,6 +352,9 @@ class RedisSemanticCache(BaseCache):
cached_prompt = cache_hit["prompt"]
cached_response = cache_hit["response"]
# update kwargs["metadata"] with similarity, don't rewrite the original metadata
kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity
print_verbose(
f"Cache hit: similarity threshold: {self.similarity_threshold}, "
f"actual similarity: {similarity}, "
@ -267,6 +365,7 @@ class RedisSemanticCache(BaseCache):
return self._get_cache_logic(cached_response=cached_response)
except Exception as e:
print_verbose(f"Error retrieving from Redis semantic cache: {str(e)}")
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
async def _get_async_embedding(self, prompt: str, **kwargs) -> List[float]:
"""
@ -321,7 +420,7 @@ class RedisSemanticCache(BaseCache):
Asynchronously store a value in the semantic cache.
Args:
key: The cache key (not directly used in semantic caching)
key: The cache key used to isolate semantic cache entries
value: The response value to cache
**kwargs: Additional arguments including 'messages' for the prompt
and optional 'ttl' for time-to-live
@ -341,21 +440,20 @@ class RedisSemanticCache(BaseCache):
# Generate embedding for the value (response) to cache
prompt_embedding = await self._get_async_embedding(prompt, **kwargs)
store_kwargs: Dict[str, Any] = {
"vector": prompt_embedding,
"filters": self._get_cache_filters(key),
}
# Get TTL and store in Redis semantic cache
ttl = self._get_ttl(**kwargs)
if ttl is not None:
await self.llmcache.astore(
prompt,
value_str,
vector=prompt_embedding, # Pass through custom embedding
ttl=ttl,
)
else:
await self.llmcache.astore(
prompt,
value_str,
vector=prompt_embedding, # Pass through custom embedding
)
store_kwargs["ttl"] = ttl
await self.llmcache.astore(
prompt,
value_str,
**store_kwargs,
)
except Exception as e:
print_verbose(f"Error in async_set_cache: {str(e)}")
@ -364,7 +462,7 @@ class RedisSemanticCache(BaseCache):
Asynchronously retrieve a semantically similar cached response.
Args:
key: The cache key (not directly used in semantic caching)
key: The cache key used to isolate semantic cache entries
**kwargs: Additional arguments including 'messages' for the prompt
Returns:
@ -385,17 +483,25 @@ class RedisSemanticCache(BaseCache):
# Generate embedding for the prompt
prompt_embedding = await self._get_async_embedding(prompt, **kwargs)
# Check the cache for semantically similar prompts
results = await self.llmcache.acheck(prompt=prompt, vector=prompt_embedding)
# Check the cache for semantically similar prompts in this exact
# LiteLLM cache-key scope.
check_kwargs: Dict[str, Any] = {
"prompt": prompt,
"vector": prompt_embedding,
"filter_expression": self._get_cache_key_filter_expression(key),
}
results = await self.llmcache.acheck(**check_kwargs)
# handle results / cache hit
if not results:
kwargs.setdefault("metadata", {})[
"semantic-similarity"
] = 0.0 # TODO why here but not above??
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
return None
cache_hit = results[0]
if not self._cache_hit_matches_key(cache_hit=cache_hit, key=key):
print_verbose("Redis semantic-cache hit did not match cache key scope")
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
return None
vector_distance = float(cache_hit["vector_distance"])
# Convert vector distance back to similarity

View file

@ -202,6 +202,12 @@ DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET = int(
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET = int(
os.getenv("DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET", 4096)
)
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET = int(
os.getenv("DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET", 8192)
)
DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET = int(
os.getenv("DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET", 16384)
)
MAX_TOKEN_TRIMMING_ATTEMPTS = int(
os.getenv("MAX_TOKEN_TRIMMING_ATTEMPTS", 10)
) # Maximum number of attempts to trim the message
@ -399,6 +405,8 @@ BEDROCK_MAX_POLICY_SIZE = int(os.getenv("BEDROCK_MAX_POLICY_SIZE", 75))
BEDROCK_MIN_THINKING_BUDGET_TOKENS = int(
os.getenv("BEDROCK_MIN_THINKING_BUDGET_TOKENS", 1024)
)
# Anthropic's Messages API rejects thinking.budget_tokens < 1024.
ANTHROPIC_MIN_THINKING_BUDGET_TOKENS = 1024
REPLICATE_POLLING_DELAY_SECONDS = float(
os.getenv("REPLICATE_POLLING_DELAY_SECONDS", 0.5)
)

View file

@ -10,6 +10,7 @@ import contextvars
import time
import uuid as uuid_module
from functools import partial
from types import MappingProxyType
from typing import Any, Coroutine, Dict, Literal, Optional, Union, cast
import httpx
@ -85,6 +86,16 @@ bedrock_files_instance = BedrockFilesHandler()
#################################################
def _add_trusted_model_credentials_to_litellm_params(
litellm_params_dict: Dict[str, Any], kwargs: Dict[str, Any]
) -> None:
trusted_model_credentials = kwargs.get("_litellm_internal_model_credentials")
if isinstance(trusted_model_credentials, type(MappingProxyType({}))):
litellm_params_dict["_litellm_internal_model_credentials"] = (
trusted_model_credentials
)
@client
async def acreate_file(
file: FileTypes,
@ -373,6 +384,10 @@ def file_retrieve(
)
if provider_config is not None:
litellm_params_dict = get_litellm_params(**kwargs)
_add_trusted_model_credentials_to_litellm_params(
litellm_params_dict=litellm_params_dict,
kwargs=kwargs,
)
litellm_params_dict["api_key"] = optional_params.api_key
litellm_params_dict["api_base"] = optional_params.api_base
@ -497,6 +512,10 @@ def file_delete(
pass
optional_params = GenericLiteLLMParams(**kwargs)
litellm_params_dict = get_litellm_params(**kwargs)
_add_trusted_model_credentials_to_litellm_params(
litellm_params_dict=litellm_params_dict,
kwargs=kwargs,
)
### TIMEOUT LOGIC ###
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
# set timeout for 10 minutes by default
@ -846,6 +865,10 @@ def file_content(
try:
optional_params = GenericLiteLLMParams(**kwargs)
litellm_params_dict = get_litellm_params(**kwargs)
_add_trusted_model_credentials_to_litellm_params(
litellm_params_dict=litellm_params_dict,
kwargs=kwargs,
)
### TIMEOUT LOGIC ###
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
client = kwargs.get("client")
@ -993,6 +1016,7 @@ def file_content(
vertex_location=vertex_ai_location,
timeout=timeout,
max_retries=optional_params.max_retries,
litellm_params=litellm_params_dict,
)
elif custom_llm_provider == "bedrock":
response = bedrock_files_instance.file_content(

View file

@ -6,12 +6,14 @@ import time
from litellm._uuid import uuid
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
from urllib.parse import quote
from litellm._logging import verbose_logger
from litellm.constants import LITELLM_ASYNCIO_QUEUE_MAXSIZE
from litellm.integrations.additional_logging_utils import AdditionalLoggingUtils
from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase
from litellm.litellm_core_utils.cloud_storage_security import (
sanitize_cloud_object_component,
)
from litellm.proxy._types import CommonProxyErrors
from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus
from litellm.types.integrations.gcs_bucket import *
@ -335,7 +337,11 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils):
_litellm_params = kwargs.get("litellm_params", None) or {}
_metadata = _litellm_params.get("metadata", None) or {}
if "gcs_log_id" in _metadata:
object_name = _metadata["gcs_log_id"]
safe_log_id = sanitize_cloud_object_component(
_metadata.get("gcs_log_id"), fallback=""
)
if safe_log_id:
object_name = f"{current_date}/custom-{uuid.uuid4().hex}-{safe_log_id}"
return object_name
@ -367,8 +373,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils):
request_date_str=date_str,
response_id=request_id,
)
encoded_object_name = quote(object_name, safe="")
response = await self.download_gcs_object(encoded_object_name)
response = await self.download_gcs_object(object_name)
if response is not None:
loaded_response = json.loads(response)

View file

@ -11,6 +11,10 @@ from litellm.integrations.gcs_bucket.gcs_bucket_mock_client import (
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.litellm_core_utils.cloud_storage_security import (
encode_gcs_object_name_for_url,
split_configured_cloud_bucket_name,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -133,8 +137,8 @@ class GCSBucketBase(CustomBatchLogger):
- Returns: bucket_name="my-bucket", object_name="my-folder/dev/my-object"
"""
if "/" in bucket_name:
bucket_name, prefix = bucket_name.split("/", 1)
bucket_name, prefix = split_configured_cloud_bucket_name(bucket_name)
if prefix:
object_name = f"{prefix}/{object_name}"
return bucket_name, object_name
return bucket_name, object_name
@ -248,6 +252,7 @@ class GCSBucketBase(CustomBatchLogger):
bucket_name=bucket_name,
object_name=object_name,
)
object_name = encode_gcs_object_name_for_url(object_name)
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{object_name}?alt=media"
@ -288,6 +293,7 @@ class GCSBucketBase(CustomBatchLogger):
bucket_name=bucket_name,
object_name=object_name,
)
object_name = encode_gcs_object_name_for_url(object_name)
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{object_name}"
@ -334,10 +340,11 @@ class GCSBucketBase(CustomBatchLogger):
bucket_name=bucket_name,
object_name=object_name,
)
encoded_object_name = encode_gcs_object_name_for_url(object_name)
response = await self.async_httpx_client.post(
headers=headers,
url=f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}",
url=f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={encoded_object_name}",
data=json_logged_payload,
)

View file

@ -1929,7 +1929,7 @@ class PrometheusLogger(CustomLogger):
or _litellm_params_metadata.get("user_agent"),
}
def set_llm_deployment_failure_metrics(self, request_kwargs: dict):
def set_llm_deployment_failure_metrics(self, request_kwargs: dict): # noqa: PLR0915
"""
Sets Failure metrics when an LLM API call fails
@ -2007,17 +2007,32 @@ class PrometheusLogger(CustomLogger):
if code is not None:
exception_status = str(code)
# Create enum_values for the label factory (always create for use in different metrics)
# On LiteLLM-side rejects (no deployment picked), route request_kwargs["model"]
# into requested_model and leave deployment-scoped labels empty.
deployment_selected = bool(model_id)
if deployment_selected:
label_litellm_model_name = litellm_model_name
label_model_id = model_id
label_api_base = api_base
label_api_provider = llm_provider
label_requested_model = model_group or litellm_model_name
else:
label_litellm_model_name = ""
label_model_id = ""
label_api_base = ""
label_api_provider = ""
label_requested_model = litellm_model_name or model_group or ""
enum_values = UserAPIKeyLabelValues(
litellm_model_name=litellm_model_name,
model_id=model_id,
api_base=api_base,
api_provider=llm_provider,
litellm_model_name=label_litellm_model_name,
model_id=label_model_id,
api_base=label_api_base,
api_provider=label_api_provider,
exception_status=exception_status,
exception_class=(
self._get_exception_class_name(exception) if exception else None
),
requested_model=model_group or litellm_model_name,
requested_model=label_requested_model,
hashed_api_key=hashed_api_key,
api_key_alias=api_key_alias,
team=team,
@ -2031,12 +2046,14 @@ class PrometheusLogger(CustomLogger):
log these labels
["litellm_model_name", "model_id", "api_base", "api_provider"]
"""
self.set_deployment_partial_outage(
litellm_model_name=litellm_model_name or "",
model_id=model_id,
api_base=api_base,
api_provider=llm_provider or "",
)
# Only mark a deployment outage when one was actually picked.
if deployment_selected:
self.set_deployment_partial_outage(
litellm_model_name=litellm_model_name or "",
model_id=model_id,
api_base=api_base,
api_provider=llm_provider or "",
)
_deployment_label_ctx = PrometheusLabelFactoryContext(enum_values)
if exception is not None:
PrometheusLogger._inc_labeled_counter(

View file

@ -0,0 +1,175 @@
import posixpath
import re
from types import MappingProxyType
from typing import Any, Mapping, Optional, Sequence, Tuple, cast
from urllib.parse import quote, unquote
from litellm._uuid import uuid
VERTEX_AI_MANAGED_GCS_PREFIX = "litellm-vertex-files/"
BEDROCK_MANAGED_S3_BATCH_PREFIX = "litellm-bedrock-files-"
BEDROCK_MANAGED_S3_UPLOAD_PREFIX = "litellm-bedrock-files/"
BEDROCK_MANAGED_S3_OUTPUT_PREFIX = "litellm-batch-outputs/"
BEDROCK_MANAGED_S3_PREFIXES = (
BEDROCK_MANAGED_S3_BATCH_PREFIX,
BEDROCK_MANAGED_S3_UPLOAD_PREFIX,
BEDROCK_MANAGED_S3_OUTPUT_PREFIX,
)
_MAPPING_PROXY_TYPE: type = type(MappingProxyType({}))
_SAFE_OBJECT_COMPONENT_PATTERN = re.compile(r"[^A-Za-z0-9._-]+")
def sanitize_cloud_object_component(
value: Optional[str], fallback: str = "file"
) -> str:
if not isinstance(value, str):
return fallback
component = posixpath.basename(value.replace("\\", "/")).strip()
if component in {"", ".", ".."}:
return fallback
component = "".join(
"_" if ord(char) < 32 or ord(char) == 127 else char for char in component
)
component = _SAFE_OBJECT_COMPONENT_PATTERN.sub("_", component)
component = component.strip("._")
if not component:
return fallback
return component[:255]
def sanitize_cloud_object_path(value: Optional[str], fallback: str = "file") -> str:
if not isinstance(value, str):
return fallback
segments = []
for segment in value.replace("\\", "/").split("/"):
sanitized_segment = sanitize_cloud_object_component(segment, fallback="")
if sanitized_segment:
segments.append(sanitized_segment)
if not segments:
return fallback
return "/".join(segments)
def build_managed_cloud_object_name(
prefix: str, filename: Optional[str], fallback_filename: str = "file"
) -> str:
safe_filename = sanitize_cloud_object_component(
filename, fallback=fallback_filename
)
return f"{prefix}{uuid.uuid4().hex}-{safe_filename}"
def _validate_cloud_object_path(object_name: str) -> None:
if not object_name:
raise ValueError("Cloud storage object name is required")
if object_name.startswith("/"):
raise ValueError("Cloud storage object name must be relative")
if any(ord(char) < 32 or ord(char) == 127 for char in object_name):
raise ValueError("Cloud storage object name contains control characters")
segments = object_name.split("/")
if any(segment in {".", ".."} for segment in segments):
raise ValueError("Cloud storage object name contains an invalid path segment")
if "" in segments[:-1]:
raise ValueError("Cloud storage object name contains an invalid path segment")
def split_configured_cloud_bucket_name(bucket_name: str) -> Tuple[str, str]:
if not isinstance(bucket_name, str) or not bucket_name.strip():
raise ValueError("Cloud storage bucket name is required")
bucket_name = bucket_name.strip()
if "://" in bucket_name or "?" in bucket_name or "#" in bucket_name:
raise ValueError(
"Cloud storage bucket name must not include a URI scheme or query"
)
if any(ord(char) < 32 or ord(char) == 127 for char in bucket_name):
raise ValueError("Cloud storage bucket name contains control characters")
bucket, _, prefix = bucket_name.partition("/")
if not bucket:
raise ValueError("Cloud storage bucket name is required")
if "\\" in bucket:
raise ValueError("Cloud storage bucket name contains an invalid separator")
prefix = prefix.strip("/")
if prefix:
_validate_cloud_object_path(prefix)
return bucket, prefix
def encode_gcs_object_name_for_url(object_name: str) -> str:
return quote(unquote(object_name), safe="")
def encode_s3_object_key_for_url(object_key: str) -> str:
return quote(unquote(object_key), safe="/")
def should_allow_legacy_cloud_file_ids(
litellm_params: Optional[Mapping[str, Any]] = None,
) -> bool:
value = None
if isinstance(litellm_params, Mapping):
trusted_model_credentials = litellm_params.get(
"_litellm_internal_model_credentials"
)
if isinstance(trusted_model_credentials, _MAPPING_PROXY_TYPE):
value = cast(Mapping[str, Any], trusted_model_credentials).get(
"allow_legacy_cloud_file_ids"
)
if isinstance(value, bool):
return value
if isinstance(value, str):
return value.strip().lower() in {"1", "true", "yes", "on"}
return False
def validate_managed_cloud_file_id(
file_id: str,
scheme: str,
configured_bucket_name: str,
allowed_object_prefixes: Sequence[str],
allow_legacy_cloud_file_ids: bool = False,
) -> Tuple[str, str]:
decoded_file_id = unquote(file_id)
if not decoded_file_id.startswith(scheme):
raise ValueError(f"file_id must be a {scheme} URI")
full_path = decoded_file_id[len(scheme) :]
if "/" not in full_path:
raise ValueError("file_id must include a cloud storage object name")
bucket_name, object_name = full_path.split("/", 1)
configured_bucket, configured_prefix = split_configured_cloud_bucket_name(
configured_bucket_name
)
if bucket_name != configured_bucket:
raise ValueError("file_id bucket does not match the configured storage bucket")
_validate_cloud_object_path(object_name)
allowed_prefixes = tuple(allowed_object_prefixes)
if configured_prefix:
allowed_prefixes = tuple(
f"{configured_prefix.rstrip('/')}/{prefix}" for prefix in allowed_prefixes
)
if object_name.startswith(allowed_prefixes):
return bucket_name, object_name
if allow_legacy_cloud_file_ids:
if configured_prefix and not object_name.startswith(
f"{configured_prefix.rstrip('/')}/"
):
raise ValueError(
"file_id object does not match the configured storage prefix"
)
return bucket_name, object_name
raise ValueError("file_id must reference a LiteLLM-managed storage object")

View file

@ -37,8 +37,6 @@ _supported_callback_params = [
"langfuse_secret_key",
"langfuse_host",
"langfuse_prompt_version",
"gcs_bucket_name",
"gcs_path_service_account",
"langsmith_api_key",
"langsmith_project",
"langsmith_base_url",
@ -57,6 +55,11 @@ _supported_callback_params = [
"lunary_public_key",
]
_request_blocked_callback_params = {
"gcs_bucket_name",
"gcs_path_service_account",
}
def initialize_standard_callback_dynamic_params(
kwargs: Optional[Dict] = None,
@ -64,13 +67,15 @@ def initialize_standard_callback_dynamic_params(
"""
Initialize the standard callback dynamic params from the kwargs
checks if langfuse_secret_key, gcs_bucket_name in kwargs and sets the corresponding attributes in StandardCallbackDynamicParams
checks supported request callback params in kwargs and sets the corresponding attributes in StandardCallbackDynamicParams
"""
standard_callback_dynamic_params = StandardCallbackDynamicParams()
if kwargs:
# 1. Check top-level kwargs
for param in _supported_callback_params:
if param in _request_blocked_callback_params:
continue
if param in kwargs:
_param_value = kwargs.get(param)
validate_no_callback_env_reference(
@ -86,6 +91,8 @@ def initialize_standard_callback_dynamic_params(
if isinstance(metadata, dict):
for param in _supported_callback_params:
if param in _request_blocked_callback_params:
continue
if param not in standard_callback_dynamic_params and param in metadata:
_param_value = metadata.get(param)
validate_no_callback_env_reference(

View file

@ -21,7 +21,7 @@ Admins can opt out via two ``litellm`` globals (wired from proxy config):
import socket
from ipaddress import ip_address, ip_network
from typing import Any, List, Set, Tuple
from typing import Any, List, Optional, Set, Tuple
from urllib.parse import quote, urlparse, urlunparse
import httpx
@ -110,6 +110,85 @@ def _normalize_host(host: str) -> str:
return host.lower().rstrip(".")
def _default_port_for_scheme(scheme: str) -> int:
return 443 if scheme == "https" else 80
def _parse_url_destination_allowlist_entry(
entry: str,
) -> Optional[Tuple[str, Optional[str], Optional[int]]]:
"""Parse an admin allowlist entry into host, optional scheme, optional port.
Entries may be bare hosts (``api.example.com``), host+port
(``api.example.com:8443``), or origins (``https://api.example.com``).
URL paths are intentionally ignored so admins can paste an api_base value.
"""
entry = entry.strip()
if not entry:
return None
has_scheme = "://" in entry
parsed = urlparse(entry if has_scheme else f"//{entry}")
if has_scheme and parsed.scheme not in _ALLOWED_SCHEMES:
return None
if parsed.username is not None or parsed.password is not None:
return None
if not parsed.hostname:
return None
try:
port = parsed.port
except ValueError:
return None
scheme: Optional[str] = parsed.scheme if has_scheme else None
if scheme is not None and port is None:
port = _default_port_for_scheme(scheme)
return _normalize_host(parsed.hostname), scheme, port
def is_url_destination_allowed_by_host(url: str, allowed_hosts: List[str]) -> bool:
"""Return True when a credential-bearing provider URL is admin-allowlisted.
This does not fetch, resolve, or rewrite URLs. It only answers whether the
destination origin is explicitly trusted by configuration. Use ``safe_get``
for user-controlled content fetches that require SSRF protection.
"""
parsed = urlparse(url)
if parsed.scheme not in _ALLOWED_SCHEMES:
return False
if parsed.username is not None or parsed.password is not None:
return False
if not parsed.hostname:
return False
try:
effective_port = parsed.port or _default_port_for_scheme(parsed.scheme)
except ValueError:
return False
normalized_host = _normalize_host(parsed.hostname)
configured_entries = (
[allowed_hosts] if isinstance(allowed_hosts, str) else allowed_hosts
)
for entry in configured_entries or []:
if not isinstance(entry, str):
continue
parsed_entry = _parse_url_destination_allowlist_entry(entry)
if parsed_entry is None:
continue
allowed_host, allowed_scheme, allowed_port = parsed_entry
if allowed_host != normalized_host:
continue
if allowed_scheme is not None and allowed_scheme != parsed.scheme:
continue
if allowed_port is not None and allowed_port != effective_port:
continue
return True
return False
def _format_host_header(hostname: str, port: int, default_port: int) -> str:
"""Build an RFC 7230 Host header value, bracketing IPv6 literals."""
bracketed = f"[{hostname}]" if ":" in hostname else hostname
@ -185,7 +264,7 @@ def validate_url(url: str) -> Tuple[str, str]:
raise SSRFError("URL has no hostname")
port = parsed.port
default_port = 443 if parsed.scheme == "https" else 80
default_port = _default_port_for_scheme(parsed.scheme)
effective_port = port if port is not None else default_port
host_header = _format_host_header(hostname, effective_port, default_port)

View file

@ -1,18 +1,31 @@
import json
import re
import time
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
from typing import (
TYPE_CHECKING,
Any,
Dict,
List,
NoReturn,
Optional,
Tuple,
Union,
cast,
)
import httpx
import litellm
from litellm.constants import (
ANTHROPIC_MIN_THINKING_BUDGET_TOKENS,
ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES,
DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS,
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET,
RESPONSE_FORMAT_TOOL_NAME,
)
from litellm.litellm_core_utils.core_helpers import map_finish_reason
@ -92,6 +105,22 @@ else:
LoggingClass = Any
REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT: Dict[str, str] = {
"low": "low",
"minimal": "low",
"medium": "medium",
"high": "high",
"xhigh": "xhigh",
"max": "max",
}
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING = (
"Dropping unsupported `output_config` for model=%s "
"(drop_params=True). Effort is only supported on Opus 4.5+, "
"Sonnet 4.6+, and Mythos Preview."
)
class AnthropicConfig(AnthropicModelInfo, BaseConfig):
"""
Reference: https://docs.anthropic.com/claude/reference/messages_post
@ -202,17 +231,96 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _supports_effort_level(model: str, level: str) -> bool:
"""Check ``supports_{level}_reasoning_effort`` in the model map.
Mirrors the pattern used in ``openai/chat/gpt_5_transformation.py`` so
that adding support for a new effort level is a pure model-map change.
Strips bedrock/vertex prefixes so a provider-routed Claude still
resolves to the Anthropic model-map entry.
"""
key = f"supports_{level}_reasoning_effort"
try:
return _supports_factory(
if _supports_factory(
model=model,
custom_llm_provider="anthropic",
key=f"supports_{level}_reasoning_effort",
)
key=key,
):
return True
except Exception:
return False
pass
candidates = [model]
for prefix in (
"bedrock/converse/",
"bedrock/invoke/",
"bedrock/",
"vertex_ai/",
):
if model.startswith(prefix):
candidates.append(model[len(prefix) :])
try:
from litellm.llms.bedrock.common_utils import BedrockModelInfo
base = BedrockModelInfo.get_base_model(model)
if base:
candidates.append(base)
candidates.append(f"bedrock/{base}")
except Exception:
pass
try:
import litellm
for cand in candidates:
if cand in litellm.model_cost and (
litellm.model_cost[cand].get(key) is True
):
return True
except Exception:
pass
return False
@staticmethod
def _validate_effort_for_model(model: str, effort: Optional[str]) -> Optional[str]:
"""Return ``None`` if ``effort`` is allowed on ``model``, else an error message."""
if effort == "max" and not (
AnthropicConfig._is_claude_4_6_model(model)
or AnthropicConfig._is_claude_4_7_model(model)
or AnthropicConfig._supports_effort_level(model, "max")
):
return f"effort='max' is not supported by this model. Got model: {model}"
if effort == "xhigh" and not AnthropicConfig._supports_effort_level(
model, "xhigh"
):
return f"effort='xhigh' is not supported by this model. Got model: {model}"
return None
@staticmethod
def _model_supports_effort_param(model: str) -> bool:
"""Whether the model accepts ``output_config.effort`` at all."""
return any(
AnthropicConfig._supports_effort_level(model, level)
for level in ("low", "minimal", "medium", "high", "xhigh", "max")
)
@staticmethod
def _raise_invalid_reasoning_effort(
model: str, value: Any, llm_provider: str
) -> NoReturn:
"""Raise a ``BadRequestError`` for an unrecognised ``reasoning_effort``.
Args:
model: The model id the request was routed to (surfaced in the error).
value: The offending ``reasoning_effort`` value supplied by the caller.
llm_provider: Provider tag for the raised exception (``"anthropic"``,
``"bedrock_converse"``, ``"databricks"``, ...).
Raises:
litellm.exceptions.BadRequestError: Always.
"""
raise litellm.exceptions.BadRequestError(
message=(
f"Invalid reasoning_effort: {value!r}. "
f"Must be one of: 'minimal', 'low', 'medium', "
f"'high', 'xhigh', 'max', 'none'"
),
model=model,
llm_provider=llm_provider,
)
def get_supported_openai_params(self, model: str):
params = [
@ -794,12 +902,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _map_reasoning_effort(
reasoning_effort: Optional[Union[REASONING_EFFORT, str]],
model: str,
llm_provider: str = "anthropic",
) -> Optional[AnthropicThinkingParam]:
if reasoning_effort is None or reasoning_effort == "none":
return None
if AnthropicConfig._is_claude_4_6_model(
model
) or AnthropicConfig._is_claude_4_7_model(model):
if AnthropicConfig._is_adaptive_thinking_model(model):
return AnthropicThinkingParam(
type="adaptive",
)
@ -818,13 +925,34 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
)
elif reasoning_effort == "xhigh":
return AnthropicThinkingParam(
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET,
)
elif reasoning_effort == "max":
return AnthropicThinkingParam(
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET,
)
elif reasoning_effort == "minimal":
return AnthropicThinkingParam(
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET,
budget_tokens=max(
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET,
ANTHROPIC_MIN_THINKING_BUDGET_TOKENS,
),
)
else:
raise ValueError(f"Unmapped reasoning effort: {reasoning_effort}")
raise litellm.exceptions.BadRequestError(
message=(
f"Unmapped reasoning effort: {reasoning_effort!r}. "
f"Must be one of: 'minimal', 'low', 'medium', 'high', "
f"'xhigh', 'max', 'none'."
),
model=model,
llm_provider=llm_provider,
)
def _extract_json_schema_from_response_format(
self, value: Optional[dict]
@ -1089,27 +1217,25 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
optional_params["thinking"] = value
elif param == "reasoning_effort" and isinstance(value, str):
mapped_thinking = AnthropicConfig._map_reasoning_effort(
reasoning_effort=value, model=model
reasoning_effort=value,
model=model,
llm_provider=self.custom_llm_provider or "anthropic",
)
if mapped_thinking is None:
optional_params.pop("thinking", None)
optional_params.pop("output_config", None)
else:
optional_params["thinking"] = mapped_thinking
# For Claude 4.6+ models, effort is controlled via output_config,
# not thinking budget_tokens. Map reasoning_effort to output_config.
if AnthropicConfig._is_claude_4_6_model(
model
) or AnthropicConfig._is_claude_4_7_model(model):
effort_map = {
"low": "low",
"minimal": "low",
"medium": "medium",
"high": "high",
"xhigh": "xhigh",
"max": "max",
}
mapped_effort = effort_map.get(value, value)
if AnthropicConfig._is_adaptive_thinking_model(model):
mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(
value
)
if mapped_effort is None:
AnthropicConfig._raise_invalid_reasoning_effort(
model=model,
value=value,
llm_provider=self.custom_llm_provider or "anthropic",
)
optional_params["output_config"] = {"effort": mapped_effort}
elif param == "web_search_options" and isinstance(value, dict):
hosted_web_search_tool = self.map_web_search_tool(
@ -1532,29 +1658,31 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
output_config = optional_params.get("output_config")
if not output_config or not isinstance(output_config, dict):
return
if litellm.drop_params is True and not self._model_supports_effort_param(model):
litellm.verbose_logger.warning(
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,
model,
)
optional_params.pop("output_config", None)
data.pop("output_config", None)
return
effort = output_config.get("effort")
valid_efforts = ["high", "medium", "low", "xhigh", "max"]
if effort and effort not in valid_efforts:
raise ValueError(
f"Invalid effort value: {effort}. Must be one of: "
f"'high', 'medium', 'low', 'xhigh', 'max'"
if effort is not None and effort not in valid_efforts:
raise litellm.exceptions.BadRequestError(
message=(
f"Invalid effort value: {effort!r}. Must be one of: "
f"'high', 'medium', 'low', 'xhigh', 'max'"
),
model=model,
llm_provider=self.custom_llm_provider or "anthropic",
)
# ``max`` is for Opus 4.6+ output effort (not Sonnet 4.6, not Opus 4.5).
# Accept known Opus 4.6/4.7 id patterns and/or ``supports_max_reasoning_effort``
# in the model map (same pattern as ``xhigh`` below).
if effort == "max" and not (
self._is_opus_4_6_model(model)
or self._is_opus_4_7_model(model)
or self._supports_effort_level(model, "max")
):
raise ValueError(
f"effort='max' is not supported by this model. Got model: {model}"
)
# ``xhigh`` is data-driven via ``supports_xhigh_reasoning_effort`` so
# enabling it for a new model is a pure model-map change.
if effort == "xhigh" and not self._supports_effort_level(model, "xhigh"):
raise ValueError(
f"effort='xhigh' is not supported by this model. Got model: {model}"
gate_error = self._validate_effort_for_model(model, effort)
if gate_error is not None:
raise litellm.exceptions.BadRequestError(
message=gate_error,
model=model,
llm_provider=self.custom_llm_provider or "anthropic",
)
data["output_config"] = output_config

View file

@ -273,7 +273,18 @@ class AnthropicModelInfo(BaseLLMModelInfo):
@staticmethod
def _is_adaptive_thinking_model(model: str) -> bool:
"""Claude 4.6+ models use adaptive thinking with output_config effort."""
"""Claude 4.6+ models use adaptive thinking with ``output_config.effort``."""
from litellm.utils import _supports_factory
try:
if _supports_factory(
model=model,
custom_llm_provider=None,
key="supports_adaptive_thinking",
):
return True
except Exception:
pass
return AnthropicModelInfo._is_claude_4_6_model(
model
) or AnthropicModelInfo._is_claude_4_7_model(model)

View file

@ -47,6 +47,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
"inference_geo",
"speed",
"output_config",
"reasoning_effort",
# TODO: Add Anthropic `metadata` support
# "metadata",
]
@ -166,6 +167,62 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
return headers, api_base
@staticmethod
def _translate_reasoning_effort_to_anthropic(
model: str, optional_params: Dict
) -> None:
"""Map OpenAI-style ``reasoning_effort`` to native Anthropic params.
Caller-supplied ``thinking`` / ``output_config`` win over the alias.
``effort='none'`` clears both. Invalid efforts raise a 400.
"""
from litellm.exceptions import BadRequestError as _BadRequestError
from litellm.llms.anthropic.chat.transformation import (
REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT,
AnthropicConfig,
)
reasoning_effort = optional_params.pop("reasoning_effort", None)
if not isinstance(reasoning_effort, str):
return
try:
mapped_thinking = AnthropicConfig._map_reasoning_effort(
reasoning_effort=reasoning_effort, model=model
)
except _BadRequestError as e:
raise AnthropicError(message=str(e.message), status_code=400)
if mapped_thinking is None:
optional_params.pop("thinking", None)
optional_params.pop("output_config", None)
return
optional_params.setdefault("thinking", mapped_thinking)
if AnthropicModelInfo._is_adaptive_thinking_model(model):
mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(
reasoning_effort
)
if mapped_effort is None:
raise AnthropicError(
message=(
f"Invalid reasoning_effort: {reasoning_effort!r}. "
f"Must be one of: 'minimal', 'low', 'medium', 'high', "
f"'xhigh', 'max', 'none'"
),
status_code=400,
)
gate_error = AnthropicConfig._validate_effort_for_model(
model, mapped_effort
)
if gate_error is not None:
raise AnthropicError(message=gate_error, status_code=400)
existing_output_config = optional_params.get("output_config")
if not isinstance(existing_output_config, dict):
existing_output_config = {}
existing_output_config.setdefault("effort", mapped_effort)
optional_params["output_config"] = existing_output_config
@staticmethod
def _translate_legacy_thinking_for_adaptive_model(
model: str, optional_params: Dict
@ -217,6 +274,11 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
status_code=400,
)
self._translate_reasoning_effort_to_anthropic(
model=model,
optional_params=anthropic_messages_optional_request_params,
)
self._translate_legacy_thinking_for_adaptive_model(
model=model,
optional_params=anthropic_messages_optional_request_params,

View file

@ -12,6 +12,23 @@ if TYPE_CHECKING:
pass
def _promote_extra_body_to_optional_params(optional_params: dict) -> None:
"""Promote anthropic-native passthrough keys out of ``extra_body``.
``azure_ai`` is an OpenAI-compatible provider, so non-OpenAI kwargs like
``output_config`` get auto-routed into ``extra_body`` by
``add_provider_specific_params_to_optional_params``. For the Azure→Anthropic
route those keys must reach the request body and be validated, so promote
them. ``setdefault`` keeps explicit top-level values authoritative.
"""
extra_body = optional_params.get("extra_body")
if not isinstance(extra_body, dict) or not extra_body:
return
for k, v in extra_body.items():
optional_params.setdefault(k, v)
optional_params.pop("extra_body", None)
class AzureAnthropicConfig(AnthropicConfig):
"""
Azure Anthropic configuration that extends AnthropicConfig.
@ -39,6 +56,8 @@ class AzureAnthropicConfig(AnthropicConfig):
1. API key via 'api-key' header
2. Azure AD token via 'Authorization: Bearer <token>' header
"""
_promote_extra_body_to_optional_params(optional_params)
# Convert dict to GenericLiteLLMParams if needed
if isinstance(litellm_params, dict):
# Ensure api_key is included if provided
@ -101,7 +120,8 @@ class AzureAnthropicConfig(AnthropicConfig):
Transform request using parent AnthropicConfig, then remove unsupported params.
Azure Anthropic doesn't support extra_body, max_retries, or stream_options parameters.
"""
# Call parent transform_request
_promote_extra_body_to_optional_params(optional_params)
data = super().transform_request(
model=model,
messages=messages,

View file

@ -87,9 +87,7 @@ class BaseConfig(ABC):
return {
k: v
for k, v in cls.__dict__.items()
if not k.startswith("__")
and not k.startswith("_abc")
and not k.startswith("_is_base_class")
if not k.startswith("_")
and not isinstance(
v,
(

View file

@ -31,7 +31,11 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
_bedrock_converse_messages_pt,
_bedrock_tools_pt,
)
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.llms.anthropic.chat.transformation import (
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,
REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT,
AnthropicConfig,
)
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.types.llms.bedrock import *
from litellm.types.llms.openai import (
@ -189,7 +193,7 @@ class AmazonConverseConfig(BaseConfig):
return {
k: v
for k, v in cls.__dict__.items()
if not k.startswith("__")
if not k.startswith("_")
and not isinstance(
v,
(
@ -410,52 +414,65 @@ class AmazonConverseConfig(BaseConfig):
"""
Handle the reasoning_effort parameter based on the model type.
Different model families handle reasoning effort differently:
- GPT-OSS models: Keep reasoning_effort as-is (passed to additionalModelRequestFields)
- Nova 2 models: Transform to reasoningConfig structure
- Other models (Anthropic, etc.): Convert to thinking parameter
Args:
model: The model identifier
reasoning_effort: The reasoning effort value
optional_params: Dictionary of optional parameters to update in-place
Examples:
>>> config = AmazonConverseConfig()
>>> params = {}
>>> config._handle_reasoning_effort_parameter("gpt-oss-model", "high", params)
>>> params
{'reasoning_effort': 'high'}
>>> params = {}
>>> config._handle_reasoning_effort_parameter("amazon.nova-2-lite-v1:0", "high", params)
>>> params
{'reasoningConfig': {'type': 'enabled', 'maxReasoningEffort': 'high'}}
>>> params = {}
>>> config._handle_reasoning_effort_parameter("anthropic.claude-3", "high", params)
>>> params
{'thinking': {'type': 'enabled', 'budget_tokens': 10000}}
- GPT-OSS models: passed through unchanged via additionalModelRequestFields.
- Nova 2 models: transformed to reasoningConfig.
- Anthropic models: mapped to ``thinking`` (and ``output_config.effort`` on
adaptive Claude 4.6 / 4.7).
"""
if "gpt-oss" in model:
# GPT-OSS models: keep reasoning_effort as-is
# It will be passed through to additionalModelRequestFields
optional_params["reasoning_effort"] = reasoning_effort
elif self._is_nova_2_model(model):
# Nova 2 models: transform to reasoningConfig
reasoning_config = self._transform_reasoning_effort_to_reasoning_config(
reasoning_effort
)
optional_params.update(reasoning_config)
else:
# Anthropic and other models: convert to thinking parameter
mapped_thinking = AnthropicConfig._map_reasoning_effort(
reasoning_effort=reasoning_effort, model=model
reasoning_effort=reasoning_effort,
model=model,
llm_provider="bedrock_converse",
)
if mapped_thinking is None:
optional_params.pop("thinking", None)
optional_params.pop("output_config", None)
else:
optional_params["thinking"] = mapped_thinking
if AnthropicConfig._is_adaptive_thinking_model(model):
mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(
reasoning_effort
)
if mapped_effort is None:
AnthropicConfig._raise_invalid_reasoning_effort(
model=model,
value=reasoning_effort,
llm_provider="bedrock_converse",
)
self._validate_anthropic_adaptive_effort(
model=model, effort=mapped_effort
)
optional_params["output_config"] = {"effort": mapped_effort}
@staticmethod
def _validate_anthropic_adaptive_effort(model: str, effort: str) -> None:
"""Validate ``output_config.effort`` for adaptive-thinking Claude 4.6/4.7."""
valid_efforts = {"high", "medium", "low", "xhigh", "max"}
if effort not in valid_efforts:
raise litellm.exceptions.BadRequestError(
message=(
f"Invalid reasoning_effort/output_config.effort value: "
f"{effort!r}. Must be one of: 'low', 'medium', 'high', "
f"'xhigh', or 'max'."
),
model=model,
llm_provider="bedrock_converse",
)
error = AnthropicConfig._validate_effort_for_model(model=model, effort=effort)
if error is not None:
raise litellm.exceptions.BadRequestError(
message=error,
model=model,
llm_provider="bedrock_converse",
)
@staticmethod
def _clamp_thinking_budget_tokens(optional_params: dict) -> None:
@ -1196,9 +1213,11 @@ class AmazonConverseConfig(BaseConfig):
+ supported_config_params
)
inference_params.pop("json_mode", None) # used for handling json_schema
# Anthropic-only key. Bedrock expects `outputConfig` (camelCase) and
# will reject `output_config` if it leaks through pass-through routes.
inference_params.pop("output_config", None)
# Anthropic-only ``output_config`` (snake_case) — re-attached to
# ``additionalModelRequestFields`` for Anthropic models below. The
# Bedrock-native ``outputConfig`` (camelCase) is handled separately.
anthropic_output_config = inference_params.pop("output_config", None)
# Extract requestMetadata before processing other parameters
request_metadata = inference_params.pop("requestMetadata", None)
@ -1208,9 +1227,6 @@ class AmazonConverseConfig(BaseConfig):
output_config: Optional[OutputConfigBlock] = inference_params.pop(
"outputConfig", None
)
inference_params.pop(
"output_config", None
) # Bedrock Converse doesn't support it
# keep supported params in 'inference_params', and set all model-specific params in 'additional_request_params'
additional_request_params = {
@ -1253,6 +1269,27 @@ class AmazonConverseConfig(BaseConfig):
additional_request_params
)
if anthropic_output_config is not None and isinstance(
anthropic_output_config, dict
):
base_model = BedrockModelInfo.get_base_model(model)
if base_model.startswith("anthropic"):
if (
litellm.drop_params is True
and not AnthropicConfig._model_supports_effort_param(model)
):
litellm.verbose_logger.warning(
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,
model,
)
else:
effort = anthropic_output_config.get("effort")
if effort is not None:
self._validate_anthropic_adaptive_effort(
model=model, effort=effort
)
additional_request_params["output_config"] = anthropic_output_config
return (
inference_params,
additional_request_params,
@ -1376,9 +1413,25 @@ class AmazonConverseConfig(BaseConfig):
# Append pre-formatted tools (systemTool etc.) after transformation
bedrock_tools.extend(pre_formatted_tools)
# Opus 4.5 gates ``output_config.effort`` behind a beta header;
# Claude 4.6/4.7 accept it without one.
base_model = BedrockModelInfo.get_base_model(model)
if base_model.startswith("anthropic"):
output_config = additional_request_params.get("output_config")
if (
isinstance(output_config, dict)
and output_config.get("effort") is not None
and not AnthropicConfig._is_adaptive_thinking_model(model)
):
from litellm.types.llms.anthropic import (
ANTHROPIC_EFFORT_BETA_HEADER,
)
if ANTHROPIC_EFFORT_BETA_HEADER not in anthropic_beta_list:
anthropic_beta_list.append(ANTHROPIC_EFFORT_BETA_HEADER)
# Set anthropic_beta in additional_request_params if we have any beta features
# ONLY apply to Anthropic/Claude models - other models (e.g., Qwen, Llama) don't support this field
base_model = BedrockModelInfo.get_base_model(model)
if anthropic_beta_list and base_model.startswith("anthropic"):
additional_request_params["anthropic_beta"] = anthropic_beta_list

View file

@ -169,7 +169,6 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
anthropic_request.pop("model", None)
anthropic_request.pop("stream", None)
anthropic_request.pop("output_format", None)
anthropic_request.pop("output_config", None)
if "anthropic_version" not in anthropic_request:
anthropic_request["anthropic_version"] = self.anthropic_version

View file

@ -1,10 +1,17 @@
import asyncio
import base64
from typing import Any, Coroutine, Optional, Tuple, Union
import os
from types import MappingProxyType
from typing import Any, Coroutine, Mapping, Optional, Tuple, Union, cast
import httpx
from litellm import LlmProviders
from litellm.litellm_core_utils.cloud_storage_security import (
BEDROCK_MANAGED_S3_PREFIXES,
should_allow_legacy_cloud_file_ids,
validate_managed_cloud_file_id,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.openai import (
FileContentRequest,
@ -35,7 +42,7 @@ class BedrockFilesHandler(BaseAWSLLM):
The file ID can be in two formats:
1. Base64-encoded unified file ID containing: llm_output_file_id,s3://bucket/path
2. Direct S3 URI: s3://bucket/path
2. Direct S3 URI: s3://bucket/litellm-managed-prefix/path
Args:
file_id: Encoded file ID or direct S3 URI
@ -58,14 +65,19 @@ class BedrockFilesHandler(BaseAWSLLM):
except Exception:
pass
# If not base64 encoded or doesn't contain llm_output_file_id, assume it's already an S3 URI
# If not base64 encoded or doesn't contain llm_output_file_id, accept only
# explicit S3 URIs. Bucket and key validation happens before any S3 call.
if file_id.startswith("s3://"):
return file_id
# If it doesn't start with s3://, assume it's a direct S3 URI and add the prefix
return f"s3://{file_id}"
raise ValueError("file_id must be a managed LiteLLM S3 file id")
def _parse_s3_uri(self, s3_uri: str) -> Tuple[str, str]:
def _parse_s3_uri(
self,
s3_uri: str,
configured_bucket_name: str,
allow_legacy_cloud_file_ids: bool = False,
) -> Tuple[str, str]:
"""
Parse S3 URI to extract bucket name and object key.
@ -75,21 +87,34 @@ class BedrockFilesHandler(BaseAWSLLM):
Returns:
Tuple of (bucket_name, object_key)
"""
if not s3_uri.startswith("s3://"):
raise ValueError(
f"Invalid S3 URI format: {s3_uri}. Expected format: s3://bucket-name/path/to/file"
return validate_managed_cloud_file_id(
file_id=s3_uri,
scheme="s3://",
configured_bucket_name=configured_bucket_name,
allowed_object_prefixes=BEDROCK_MANAGED_S3_PREFIXES,
allow_legacy_cloud_file_ids=allow_legacy_cloud_file_ids,
)
def _get_configured_s3_bucket_name(self, litellm_params: dict) -> str:
trusted_model_credentials = litellm_params.get(
"_litellm_internal_model_credentials"
)
bucket_name = None
if isinstance(trusted_model_credentials, type(MappingProxyType({}))):
trusted_model_credentials_mapping = cast(
Mapping[str, Any], trusted_model_credentials
)
# Remove 's3://' prefix
path = s3_uri[5:]
if "/" in path:
bucket_name, object_key = path.split("/", 1)
else:
bucket_name = path
object_key = ""
return bucket_name, object_key
candidate_bucket_name = trusted_model_credentials_mapping.get(
"s3_bucket_name"
)
if isinstance(candidate_bucket_name, str):
bucket_name = candidate_bucket_name
bucket_name = bucket_name or os.getenv("AWS_S3_BUCKET_NAME")
if not bucket_name:
raise ValueError(
"S3 bucket_name is required. Set 's3_bucket_name' in proxy config or AWS_S3_BUCKET_NAME for Bedrock file content retrieval."
)
return bucket_name
async def afile_content(
self,
@ -119,7 +144,14 @@ class BedrockFilesHandler(BaseAWSLLM):
# Extract S3 URI from file ID
s3_uri = self._extract_s3_uri_from_file_id(file_id)
bucket_name, object_key = self._parse_s3_uri(s3_uri)
configured_bucket_name = self._get_configured_s3_bucket_name(optional_params)
bucket_name, object_key = self._parse_s3_uri(
s3_uri=s3_uri,
configured_bucket_name=configured_bucket_name,
allow_legacy_cloud_file_ids=should_allow_legacy_cloud_file_ids(
optional_params
),
)
# Get AWS credentials
aws_region_name = self._get_aws_region_name(

View file

@ -2,6 +2,7 @@ import json
import os
import time
from typing import Any, Dict, List, Optional, Tuple, Union
from urllib.parse import unquote
import httpx
from httpx import Headers, Response
@ -10,6 +11,14 @@ from openai.types.file_deleted import FileDeleted
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.files.utils import FilesAPIUtils
from litellm.litellm_core_utils.cloud_storage_security import (
BEDROCK_MANAGED_S3_BATCH_PREFIX,
BEDROCK_MANAGED_S3_UPLOAD_PREFIX,
build_managed_cloud_object_name,
encode_s3_object_key_for_url,
sanitize_cloud_object_component,
split_configured_cloud_bucket_name,
)
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.files.transformation import (
@ -116,10 +125,13 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
if _model.startswith("bedrock/"):
_model = _model[8:]
# Replace colons with hyphens for Bedrock S3 URI compliance
_model = _model.replace(":", "-")
safe_model = sanitize_cloud_object_component(
_model.replace(":", "-"), fallback="model"
)
object_name = f"litellm-bedrock-files-{_model}-{uuid.uuid4()}.jsonl"
object_name = (
f"{BEDROCK_MANAGED_S3_BATCH_PREFIX}{safe_model}-{uuid.uuid4()}.jsonl"
)
return object_name
def get_object_name(
@ -146,12 +158,13 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
if len(openai_jsonl_content) > 0:
return self._get_s3_object_name_from_batch_jsonl(openai_jsonl_content)
## 2. If not jsonl, return the filename
## 2. If not jsonl, store under a server-generated managed object name
filename = extracted_file_data.get("filename")
if filename:
return filename
## 3. If no file name, return timestamp
return str(int(time.time()))
return build_managed_cloud_object_name(
prefix=BEDROCK_MANAGED_S3_UPLOAD_PREFIX,
filename=filename,
fallback_filename="file",
)
def get_complete_file_url(
self,
@ -172,6 +185,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
raise ValueError(
"S3 bucket_name is required. Set 's3_bucket_name' in litellm_params or AWS_S3_BUCKET_NAME env var"
)
bucket_name, object_prefix = split_configured_cloud_bucket_name(bucket_name)
s3_region_name = litellm_params.get("s3_region_name") or optional_params.get(
"s3_region_name"
@ -188,14 +202,17 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
raise ValueError("purpose is required")
extracted_file_data = extract_file_data(file_data)
object_name = self.get_object_name(extracted_file_data, purpose)
if object_prefix:
object_name = f"{object_prefix}/{object_name}"
encoded_object_name = encode_s3_object_key_for_url(object_name)
# S3 endpoint URL format
s3_endpoint_url = (
optional_params.get("s3_endpoint_url")
or f"https://s3.{aws_region_name}.amazonaws.com"
)
).rstrip("/")
return f"{s3_endpoint_url}/{bucket_name}/{object_name}"
return f"{s3_endpoint_url}/{bucket_name}/{encoded_object_name}"
def get_supported_openai_params(
self, model: str
@ -532,10 +549,12 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
if match1:
# Pattern: https://s3.region.amazonaws.com/bucket/key
region, bucket, key = match1.groups()
key = unquote(key)
s3_uri = f"s3://{bucket}/{key}"
elif match2:
# Pattern: https://bucket.s3.region.amazonaws.com/key
bucket, region, key = match2.groups()
key = unquote(key)
s3_uri = f"s3://{bucket}/{key}"
else:
# Fallback: try to extract bucket and key from URL path
@ -545,6 +564,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
path_parts = parsed.path.lstrip("/").split("/", 1)
if len(path_parts) >= 2:
bucket, key = path_parts[0], path_parts[1]
key = unquote(key)
s3_uri = f"s3://{bucket}/{key}"
else:
raise ValueError(f"Unable to parse S3 URL: {https_url}")
@ -722,7 +742,12 @@ class BedrockJsonlFilesTransformation:
# Remove bedrock/ prefix if present
if _model.startswith("bedrock/"):
_model = _model[8:]
object_name = f"litellm-bedrock-files-{_model}-{uuid.uuid4()}.jsonl"
safe_model = sanitize_cloud_object_component(
_model.replace(":", "-"), fallback="model"
)
object_name = (
f"{BEDROCK_MANAGED_S3_BATCH_PREFIX}{safe_model}-{uuid.uuid4()}.jsonl"
)
return object_name
def _get_content_from_openai_file(self, openai_file_content: FileTypes) -> str:

View file

@ -12,9 +12,14 @@ from typing import (
import httpx
import litellm
from litellm.anthropic_beta_headers_manager import filter_and_transform_beta_headers
from litellm.constants import BEDROCK_MIN_THINKING_BUDGET_TOKENS
from litellm.litellm_core_utils.litellm_logging import verbose_logger
from litellm.llms.anthropic.chat.transformation import (
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,
AnthropicConfig,
)
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
@ -580,6 +585,17 @@ class AmazonAnthropicClaudeMessagesConfig(
if filtered_betas:
anthropic_messages_request["anthropic_beta"] = filtered_betas
if (
litellm.drop_params is True
and "output_config" in anthropic_messages_request
and not AnthropicConfig._model_supports_effort_param(model)
):
verbose_logger.warning(
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,
model,
)
anthropic_messages_request.pop("output_config", None)
# 7. Final safety net: filter top-level fields to the Bedrock Invoke allowlist.
# Catches Anthropic-only extensions (context_management, output_config, speed,
# mcp_servers, ...) and any future additions Claude Code may start sending.

View file

@ -56,7 +56,10 @@ from litellm.types.utils import (
Usage,
)
from ...anthropic.chat.transformation import AnthropicConfig
from ...anthropic.chat.transformation import (
REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT,
AnthropicConfig,
)
from ...openai_like.chat.transformation import OpenAILikeChatConfig
from ..common_utils import DatabricksBase, DatabricksException
@ -330,9 +333,30 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
) # unsupported for claude models - if json_schema -> convert to tool call
if "reasoning_effort" in non_default_params and "claude" in model:
optional_params["thinking"] = AnthropicConfig._map_reasoning_effort(
reasoning_effort=non_default_params.get("reasoning_effort"), model=model
reasoning_effort_value = non_default_params.get("reasoning_effort")
mapped_thinking = AnthropicConfig._map_reasoning_effort(
reasoning_effort=reasoning_effort_value,
model=model,
llm_provider="databricks",
)
if mapped_thinking is None:
optional_params.pop("thinking", None)
optional_params.pop("output_config", None)
else:
optional_params["thinking"] = mapped_thinking
if AnthropicConfig._is_adaptive_thinking_model(model):
mapped_effort: Optional[str] = None
if isinstance(reasoning_effort_value, str):
mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(
reasoning_effort_value
)
if mapped_effort is None:
AnthropicConfig._raise_invalid_reasoning_effort(
model=model,
value=reasoning_effort_value,
llm_provider="databricks",
)
optional_params["output_config"] = {"effort": mapped_effort}
optional_params.pop("reasoning_effort", None)
## handle thinking tokens
self.update_optional_params_with_thinking_tokens(

View file

@ -1,6 +1,6 @@
import asyncio
import time
import urllib.parse
from urllib.parse import unquote
from typing import Any, Coroutine, Optional, Tuple, Union
import httpx
@ -10,6 +10,11 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import (
GCSBucketBase,
GCSLoggingConfig,
)
from litellm.litellm_core_utils.cloud_storage_security import (
VERTEX_AI_MANAGED_GCS_PREFIX,
should_allow_legacy_cloud_file_ids,
validate_managed_cloud_file_id,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.openai import (
CreateFileRequest,
@ -114,34 +119,31 @@ class VertexAIFilesHandler(GCSBucketBase):
)
)
def _extract_bucket_and_object_from_file_id(self, file_id: str) -> Tuple[str, str]:
def _extract_bucket_and_object_from_file_id(
self,
file_id: str,
configured_bucket_name: str,
litellm_params: Optional[dict] = None,
) -> Tuple[str, str]:
"""
Extract bucket name and object path from URL-encoded file_id.
Validate and extract bucket name and object path from file_id.
Expected format: gs%3A%2F%2Fbucket-name%2Fpath%2Fto%2Ffile
Which decodes to: gs://bucket-name/path/to/file
Expected format: gs://bucket-name/litellm-vertex-files/path/to/file
Returns:
tuple: (bucket_name, url_encoded_object_path)
tuple: (bucket_name, object_path)
- bucket_name: "bucket-name"
- url_encoded_object_path: "path%2Fto%2Ffile"
- object_path: "litellm-vertex-files/path/to/file"
"""
decoded_path = urllib.parse.unquote(file_id)
if decoded_path.startswith("gs://"):
full_path = decoded_path[5:] # Remove 'gs://' prefix
else:
full_path = decoded_path
if "/" in full_path:
bucket_name, object_path = full_path.split("/", 1)
else:
bucket_name = full_path
object_path = ""
encoded_object_path = urllib.parse.quote(object_path, safe="")
return bucket_name, encoded_object_path
return validate_managed_cloud_file_id(
file_id=file_id,
scheme="gs://",
configured_bucket_name=configured_bucket_name,
allowed_object_prefixes=(VERTEX_AI_MANAGED_GCS_PREFIX,),
allow_legacy_cloud_file_ids=should_allow_legacy_cloud_file_ids(
litellm_params
),
)
async def afile_content(
self,
@ -151,6 +153,7 @@ class VertexAIFilesHandler(GCSBucketBase):
vertex_location: Optional[str],
timeout: Union[float, httpx.Timeout],
max_retries: Optional[int],
litellm_params: Optional[dict] = None,
) -> HttpxBinaryResponseContent:
"""
Download file content from GCS bucket for VertexAI files.
@ -170,23 +173,30 @@ class VertexAIFilesHandler(GCSBucketBase):
if not file_id:
raise ValueError("file_id is required in file_content_request")
bucket_name, encoded_object_path = self._extract_bucket_and_object_from_file_id(
file_id
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(
kwargs={}
)
bucket_name, object_path = self._extract_bucket_and_object_from_file_id(
file_id=file_id,
configured_bucket_name=gcs_logging_config["bucket_name"],
litellm_params=litellm_params,
)
download_kwargs = {
"standard_callback_dynamic_params": {"gcs_bucket_name": bucket_name}
"standard_callback_dynamic_params": {
"gcs_bucket_name": bucket_name,
"gcs_path_service_account": gcs_logging_config["path_service_account"],
}
}
file_content = await self.download_gcs_object(
object_name=encoded_object_path, **download_kwargs
object_name=object_path, **download_kwargs
)
decoded_file_id = unquote(file_id)
if file_content is None:
decoded_path = urllib.parse.unquote(file_id)
raise ValueError(f"Failed to download file from GCS: {decoded_path}")
raise ValueError(f"Failed to download file from GCS: {decoded_file_id}")
decoded_path = urllib.parse.unquote(file_id)
mock_response = httpx.Response(
status_code=200,
content=file_content,
@ -194,7 +204,7 @@ class VertexAIFilesHandler(GCSBucketBase):
"content-type": "application/octet-stream",
"content-length": str(len(file_content)),
},
request=httpx.Request(method="GET", url=decoded_path),
request=httpx.Request(method="GET", url=decoded_file_id),
)
# Apply transformation to convert Vertex AI batch outputs to OpenAI format
@ -225,6 +235,7 @@ class VertexAIFilesHandler(GCSBucketBase):
vertex_location: Optional[str],
timeout: Union[float, httpx.Timeout],
max_retries: Optional[int],
litellm_params: Optional[dict] = None,
) -> Union[
HttpxBinaryResponseContent, Coroutine[Any, Any, HttpxBinaryResponseContent]
]:
@ -253,6 +264,7 @@ class VertexAIFilesHandler(GCSBucketBase):
vertex_location=vertex_location,
timeout=timeout,
max_retries=max_retries,
litellm_params=litellm_params,
)
else:
return asyncio.run(
@ -263,5 +275,6 @@ class VertexAIFilesHandler(GCSBucketBase):
vertex_location=vertex_location,
timeout=timeout,
max_retries=max_retries,
litellm_params=litellm_params,
)
)

View file

@ -12,6 +12,15 @@ from openai.types.file_deleted import FileDeleted
import litellm
from litellm._uuid import uuid
from litellm.files.utils import FilesAPIUtils
from litellm.litellm_core_utils.cloud_storage_security import (
VERTEX_AI_MANAGED_GCS_PREFIX,
build_managed_cloud_object_name,
encode_gcs_object_name_for_url,
sanitize_cloud_object_path,
should_allow_legacy_cloud_file_ids,
split_configured_cloud_bucket_name,
validate_managed_cloud_file_id,
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
from litellm.llms.base_llm.chat.transformation import BaseLLMException
@ -248,7 +257,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
_model = openai_jsonl_content[0].get("body", {}).get("model", "")
if "publishers/google/models" not in _model:
_model = f"publishers/google/models/{_model}"
object_name = f"litellm-vertex-files/{_model}/{uuid.uuid4()}"
safe_model_path = sanitize_cloud_object_path(_model, fallback="model")
object_name = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
return object_name
def get_object_name(
@ -275,12 +285,19 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
if len(openai_jsonl_content) > 0:
return self._get_gcs_object_name_from_batch_jsonl(openai_jsonl_content)
## 2. If not jsonl, return the filename
## 2. If not jsonl, store under a server-generated managed object name
filename = extracted_file_data.get("filename")
if filename:
return filename
## 3. If no file name, return timestamp
return str(int(time.time()))
return build_managed_cloud_object_name(
prefix=f"{VERTEX_AI_MANAGED_GCS_PREFIX}uploads/",
filename=filename,
fallback_filename="file",
)
def _get_configured_bucket_name(self, litellm_params: Dict) -> str:
bucket_name = litellm_params.get("bucket_name") or os.getenv("GCS_BUCKET_NAME")
if not bucket_name:
raise ValueError("GCS bucket_name is required")
return bucket_name
def get_complete_file_url(
self,
@ -294,13 +311,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"""
Get the complete url for the request
"""
bucket_name = (
litellm_params.get("bucket_name")
or litellm_params.get("litellm_metadata", {}).pop("gcs_bucket_name", None)
or os.getenv("GCS_BUCKET_NAME")
)
if not bucket_name:
raise ValueError("GCS bucket_name is required")
bucket_name = self._get_configured_bucket_name(litellm_params)
bucket_name, object_prefix = split_configured_cloud_bucket_name(bucket_name)
file_data = data.get("file")
purpose = data.get("purpose")
if file_data is None:
@ -309,9 +321,10 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
raise ValueError("purpose is required")
extracted_file_data = extract_file_data(file_data)
object_name = self.get_object_name(extracted_file_data, purpose)
endpoint = (
f"upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}"
)
if object_prefix:
object_name = f"{object_prefix}/{object_name}"
encoded_object_name = encode_gcs_object_name_for_url(object_name)
endpoint = f"upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={encoded_object_name}"
api_base = api_base or "https://storage.googleapis.com"
if not api_base:
raise ValueError("api_base is required")
@ -450,27 +463,23 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
status_code=status_code, message=error_message, headers=headers
)
def _parse_gcs_uri(self, file_id: str) -> Tuple[str, str]:
def _parse_gcs_uri(
self, file_id: str, litellm_params: Optional[Dict] = None
) -> Tuple[str, str]:
"""
Parse a GCS URI (gs://bucket/path/to/object) into (bucket, url-encoded-object-path).
Handles both raw and URL-encoded input.
Validate a managed GCS file_id and return (bucket, url-encoded-object-path).
"""
import urllib.parse
decoded = urllib.parse.unquote(file_id)
if decoded.startswith("gs://"):
full_path = decoded[5:]
else:
full_path = decoded
if "/" in full_path:
bucket_name, object_path = full_path.split("/", 1)
else:
bucket_name = full_path
object_path = ""
encoded_object = urllib.parse.quote(object_path, safe="")
return bucket_name, encoded_object
configured_bucket_name = self._get_configured_bucket_name(litellm_params or {})
bucket_name, object_path = validate_managed_cloud_file_id(
file_id=file_id,
scheme="gs://",
configured_bucket_name=configured_bucket_name,
allowed_object_prefixes=(VERTEX_AI_MANAGED_GCS_PREFIX,),
allow_legacy_cloud_file_ids=should_allow_legacy_cloud_file_ids(
litellm_params
),
)
return bucket_name, encode_gcs_object_name_for_url(object_path)
def transform_retrieve_file_request(
self,
@ -478,7 +487,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
optional_params: dict,
litellm_params: dict,
) -> tuple[str, dict]:
bucket, encoded_object = self._parse_gcs_uri(file_id)
bucket, encoded_object = self._parse_gcs_uri(file_id, litellm_params)
url = f"https://storage.googleapis.com/storage/v1/b/{bucket}/o/{encoded_object}"
return url, {}
@ -510,7 +519,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
optional_params: dict,
litellm_params: dict,
) -> tuple[str, dict]:
bucket, encoded_object = self._parse_gcs_uri(file_id)
bucket, encoded_object = self._parse_gcs_uri(file_id, litellm_params)
url = f"https://storage.googleapis.com/storage/v1/b/{bucket}/o/{encoded_object}"
return url, {}
@ -554,7 +563,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
litellm_params: dict,
) -> tuple[str, dict]:
file_id = file_content_request.get("file_id", "")
bucket, encoded_object = self._parse_gcs_uri(file_id)
bucket, encoded_object = self._parse_gcs_uri(file_id, litellm_params)
url = f"https://storage.googleapis.com/storage/v1/b/{bucket}/o/{encoded_object}?alt=media"
return url, {}
@ -842,7 +851,8 @@ class VertexAIJsonlFilesTransformation(VertexGeminiConfig):
_model = openai_jsonl_content[0].get("body", {}).get("model", "")
if "publishers/google/models" not in _model:
_model = f"publishers/google/models/{_model}"
object_name = f"litellm-vertex-files/{_model}/{uuid.uuid4()}"
safe_model_path = sanitize_cloud_object_path(_model, fallback="model")
object_name = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
return object_name
def _map_openai_to_vertex_params(

View file

@ -159,10 +159,6 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert
"model", None
) # do not pass model in request body to vertex ai
# Vertex AI Claude accepts ``output_config.format`` (structured outputs)
# and ``output_format``, but rejects ``output_config.effort`` with 400
# "Extra inputs are not permitted". Sanitize in place so the supported
# bits flow through.
sanitize_vertex_anthropic_output_params(anthropic_messages_request)
return anthropic_messages_request

View file

@ -11,11 +11,9 @@ keeps the parent module's import surface narrow.
"""
# Keys inside ``output_config`` that Vertex AI Claude does not accept.
# Today only ``effort`` triggers "Extra inputs are not permitted"; add new
# entries here as Vertex parity drifts. Keep this list narrow — anything
# Vertex DOES accept (e.g. ``format`` for structured outputs) must be
# preserved so callers can rely on Anthropic-native features.
VERTEX_UNSUPPORTED_OUTPUT_CONFIG_KEYS: frozenset = frozenset({"effort"})
# Add an entry only when a 400 "Extra inputs are not permitted" is
# reproducible against the live Vertex endpoint.
VERTEX_UNSUPPORTED_OUTPUT_CONFIG_KEYS: frozenset = frozenset()
def sanitize_vertex_anthropic_output_params(data: dict) -> None:

View file

@ -106,11 +106,6 @@ class VertexAIAnthropicConfig(AnthropicConfig):
data.pop("model", None) # vertex anthropic doesn't accept 'model' parameter
# Vertex AI Claude accepts ``output_config.format`` (structured outputs /
# JSON Schema) but NOT ``output_config.effort`` — sending ``effort`` to
# Vertex returns 400 "Extra inputs are not permitted". Sanitize in place:
# forward the structured-output bits, drop the unsupported keys.
# Same treatment for the legacy top-level ``output_format`` field.
sanitize_vertex_anthropic_output_params(data)
tools = optional_params.get("tools")

View file

@ -223,8 +223,43 @@ class XAIChatConfig(OpenAIGPTConfig):
self._enhance_usage_with_xai_web_search_fields(response, raw_response_json)
except Exception as e:
verbose_logger.debug(f"Error extracting X.AI web search usage: {e}")
self._fold_reasoning_tokens_into_completion(response)
return response
@staticmethod
def _fold_reasoning_tokens_into_completion(model_response: ModelResponse) -> None:
"""Reconcile xAI Usage to the OpenAI invariant.
xAI accounts ``reasoning_tokens`` separately from
``completion_tokens`` while still summing them into ``total_tokens``.
OpenAI's contract (o1/o3) folds reasoning into ``completion_tokens``,
so fold here to keep ``total = prompt + completion``. Idempotent.
"""
usage = getattr(model_response, "usage", None)
if usage is None:
return
details = getattr(usage, "completion_tokens_details", None)
reasoning_tokens = (
int(getattr(details, "reasoning_tokens", 0) or 0) if details else 0
)
if reasoning_tokens <= 0:
return
prompt_tokens = int(getattr(usage, "prompt_tokens", 0) or 0)
completion_tokens = int(getattr(usage, "completion_tokens", 0) or 0)
total_tokens = int(getattr(usage, "total_tokens", 0) or 0)
if total_tokens == prompt_tokens + completion_tokens:
return
# Guard against double-counting if xAI changes accounting.
if total_tokens != prompt_tokens + completion_tokens + reasoning_tokens:
return
usage.completion_tokens = completion_tokens + reasoning_tokens
def _enhance_usage_with_xai_web_search_fields(
self, model_response: ModelResponse, raw_response_json: dict
) -> None:

View file

@ -25,16 +25,25 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
"""
# XAI-specific completion cost calculation
# For XAI models, completion is billed as (visible completion tokens + reasoning tokens)
# XAI-specific completion cost: completion is billed as visible + reasoning
# tokens. Detect when the transformation layer already folded them so we
# don't double-count; fall back to raw xAI shape for callers that bypass
# the transformation (e.g. proxy logs replayed into cost calc).
prompt_tokens = int(getattr(usage, "prompt_tokens", 0) or 0)
completion_tokens = int(getattr(usage, "completion_tokens", 0) or 0)
total_tokens = int(getattr(usage, "total_tokens", 0) or 0)
reasoning_tokens = 0
if hasattr(usage, "completion_tokens_details") and usage.completion_tokens_details:
reasoning_tokens = int(
getattr(usage.completion_tokens_details, "reasoning_tokens", 0) or 0
)
total_completion_tokens = completion_tokens + reasoning_tokens
already_normalised = total_tokens == prompt_tokens + completion_tokens
total_completion_tokens = (
completion_tokens
if already_normalised
else completion_tokens + reasoning_tokens
)
modified_usage = Usage(
prompt_tokens=usage.prompt_tokens,

View file

@ -977,6 +977,7 @@
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
@ -1162,6 +1163,21 @@
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"anthropic.claude-mythos-preview": {
"input_cost_per_token": 0,
"output_cost_per_token": 0,
"litellm_provider": "bedrock",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"supports_prompt_caching": false,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_tool_choice": true
},
"global.anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
@ -1307,6 +1323,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -1336,6 +1353,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -1365,6 +1383,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -1393,6 +1412,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -1421,6 +1441,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -1915,6 +1936,7 @@
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
@ -2038,6 +2060,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -9212,6 +9235,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -9347,6 +9371,7 @@
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
@ -9374,6 +9399,7 @@
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
@ -9477,7 +9503,6 @@
"us": 1.1,
"fast": 6.0
},
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"claude-opus-4-7-20260416": {
@ -9512,7 +9537,6 @@
"us": 1.1,
"fast": 6.0
},
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"claude-sonnet-4-20250514": {
@ -10790,6 +10814,7 @@
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_tool_choice": true
},
"databricks/databricks-claude-sonnet-4": {
@ -15548,7 +15573,7 @@
"mode": "embedding",
"output_cost_per_token": 0,
"output_vector_size": 3072,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal",
"supports_multimodal": true,
"uses_embed_content": true
},
@ -17150,7 +17175,8 @@
],
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": true
"supports_vision": true,
"supports_minimal_reasoning_effort": true
},
"github_copilot/claude-opus-4.6-fast": {
"litellm_provider": "github_copilot",
@ -17663,7 +17689,8 @@
"mode": "chat",
"output_cost_per_token": 2.5e-05,
"supports_function_calling": true,
"supports_vision": true
"supports_vision": true,
"supports_minimal_reasoning_effort": true
},
"gmi/anthropic/claude-sonnet-4.5": {
"input_cost_per_token": 3e-06,
@ -21142,7 +21169,7 @@
},
"gradient_ai/alibaba-qwen3-32b": {
"litellm_provider": "gradient_ai",
"max_tokens": 2048,
"max_tokens": 40960,
"mode": "chat",
"supported_endpoints": [
"/v1/chat/completions"
@ -21150,7 +21177,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 131072,
"max_output_tokens": 40960
},
"gradient_ai/anthropic-claude-3-opus": {
"input_cost_per_token": 1.5e-05,
@ -21164,7 +21193,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 200000,
"max_output_tokens": 1024
},
"gradient_ai/anthropic-claude-3.5-haiku": {
"input_cost_per_token": 8e-07,
@ -21178,7 +21209,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 200000,
"max_output_tokens": 1024
},
"gradient_ai/anthropic-claude-3.5-sonnet": {
"input_cost_per_token": 3e-06,
@ -21192,7 +21225,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 200000,
"max_output_tokens": 1024
},
"gradient_ai/anthropic-claude-3.7-sonnet": {
"input_cost_per_token": 3e-06,
@ -21206,7 +21241,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 200000,
"max_output_tokens": 1024
},
"gradient_ai/deepseek-r1-distill-llama-70b": {
"input_cost_per_token": 9.9e-07,
@ -21220,7 +21257,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 32768,
"max_output_tokens": 8000
},
"gradient_ai/llama3-8b-instruct": {
"input_cost_per_token": 2e-07,
@ -21234,7 +21273,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 8192,
"max_output_tokens": 512
},
"gradient_ai/llama3.3-70b-instruct": {
"input_cost_per_token": 6.5e-07,
@ -21248,7 +21289,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 128000,
"max_output_tokens": 2048
},
"gradient_ai/mistral-nemo-instruct-2407": {
"input_cost_per_token": 3e-07,
@ -21262,7 +21305,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 128000,
"max_output_tokens": 512
},
"gradient_ai/openai-gpt-4o": {
"litellm_provider": "gradient_ai",
@ -21274,7 +21319,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 128000,
"max_output_tokens": 16384
},
"gradient_ai/openai-gpt-4o-mini": {
"litellm_provider": "gradient_ai",
@ -21286,7 +21333,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 128000,
"max_output_tokens": 16384
},
"gradient_ai/openai-o3": {
"input_cost_per_token": 2e-06,
@ -21300,7 +21349,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 200000,
"max_output_tokens": 100000
},
"gradient_ai/openai-o3-mini": {
"input_cost_per_token": 1.1e-06,
@ -21314,7 +21365,9 @@
"supported_modalities": [
"text"
],
"supports_tool_choice": false
"supports_tool_choice": false,
"max_input_tokens": 200000,
"max_output_tokens": 100000
},
"lemonade/Qwen3-Coder-30B-A3B-Instruct-GGUF": {
"input_cost_per_token": 0,
@ -21614,11 +21667,13 @@
},
"heroku/claude-3-5-haiku": {
"litellm_provider": "heroku",
"max_tokens": 4096,
"max_tokens": 8192,
"mode": "chat",
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 200000,
"max_output_tokens": 8192
},
"heroku/claude-3-5-sonnet-latest": {
"litellm_provider": "heroku",
@ -21626,7 +21681,9 @@
"mode": "chat",
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 200000,
"max_output_tokens": 8192
},
"heroku/claude-3-7-sonnet": {
"litellm_provider": "heroku",
@ -21634,7 +21691,9 @@
"mode": "chat",
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 200000,
"max_output_tokens": 8192
},
"heroku/claude-4-sonnet": {
"litellm_provider": "heroku",
@ -21642,7 +21701,9 @@
"mode": "chat",
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 200000,
"max_output_tokens": 8192
},
"high/1024-x-1024/gpt-image-1": {
"input_cost_per_image": 0.167,
@ -22450,48 +22511,6 @@
"/v1/images/generations"
]
},
"luminous-base": {
"input_cost_per_token": 3e-05,
"litellm_provider": "aleph_alpha",
"max_tokens": 2048,
"mode": "completion",
"output_cost_per_token": 3.3e-05
},
"luminous-base-control": {
"input_cost_per_token": 3.75e-05,
"litellm_provider": "aleph_alpha",
"max_tokens": 2048,
"mode": "chat",
"output_cost_per_token": 4.125e-05
},
"luminous-extended": {
"input_cost_per_token": 4.5e-05,
"litellm_provider": "aleph_alpha",
"max_tokens": 2048,
"mode": "completion",
"output_cost_per_token": 4.95e-05
},
"luminous-extended-control": {
"input_cost_per_token": 5.625e-05,
"litellm_provider": "aleph_alpha",
"max_tokens": 2048,
"mode": "chat",
"output_cost_per_token": 6.1875e-05
},
"luminous-supreme": {
"input_cost_per_token": 0.000175,
"litellm_provider": "aleph_alpha",
"max_tokens": 2048,
"mode": "completion",
"output_cost_per_token": 0.0001925
},
"luminous-supreme-control": {
"input_cost_per_token": 0.00021875,
"litellm_provider": "aleph_alpha",
"max_tokens": 2048,
"mode": "chat",
"output_cost_per_token": 0.000240625
},
"max-x-max/50-steps/stability.stable-diffusion-xl-v0": {
"litellm_provider": "bedrock",
"max_input_tokens": 77,
@ -25963,12 +25982,14 @@
"input_cost_per_image": 0.0004,
"input_cost_per_token": 2.5e-07,
"litellm_provider": "openrouter",
"max_tokens": 200000,
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 1.25e-06,
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"max_input_tokens": 200000,
"max_output_tokens": 4096
},
"openrouter/anthropic/claude-3.5-sonnet": {
"input_cost_per_token": 3e-06,
@ -26087,6 +26108,7 @@
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 159,
@ -26105,6 +26127,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true,
@ -26126,6 +26149,7 @@
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -26174,6 +26198,29 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346
},
"openrouter/anthropic/claude-opus-4.7": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "openrouter",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"tool_use_system_prompt_tokens": 346
},
"openrouter/bytedance/ui-tars-1.5-7b": {
"input_cost_per_token": 1e-07,
"litellm_provider": "openrouter",
@ -26539,18 +26586,22 @@
"openrouter/mancer/weaver": {
"input_cost_per_token": 5.625e-06,
"litellm_provider": "openrouter",
"max_tokens": 8000,
"max_tokens": 2000,
"mode": "chat",
"output_cost_per_token": 5.625e-06,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 8000,
"max_output_tokens": 2000
},
"openrouter/meta-llama/llama-3-70b-instruct": {
"input_cost_per_token": 5.9e-07,
"litellm_provider": "openrouter",
"max_tokens": 8192,
"max_tokens": 8000,
"mode": "chat",
"output_cost_per_token": 7.9e-07,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 8192,
"max_output_tokens": 8000
},
"openrouter/minimax/minimax-m2": {
"input_cost_per_token": 2.55e-07,
@ -26638,34 +26689,42 @@
"openrouter/mistralai/mistral-7b-instruct": {
"input_cost_per_token": 1.3e-07,
"litellm_provider": "openrouter",
"max_tokens": 8192,
"max_tokens": 8191,
"mode": "chat",
"output_cost_per_token": 1.3e-07,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 32768,
"max_output_tokens": 8191
},
"openrouter/mistralai/mistral-large": {
"input_cost_per_token": 8e-06,
"litellm_provider": "openrouter",
"max_tokens": 32000,
"max_tokens": 8191,
"mode": "chat",
"output_cost_per_token": 2.4e-05,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 128000,
"max_output_tokens": 8191
},
"openrouter/mistralai/mistral-small-3.1-24b-instruct": {
"input_cost_per_token": 1e-07,
"litellm_provider": "openrouter",
"max_tokens": 32000,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 3e-07,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 131072,
"max_output_tokens": 131072
},
"openrouter/mistralai/mistral-small-3.2-24b-instruct": {
"input_cost_per_token": 1e-07,
"litellm_provider": "openrouter",
"max_tokens": 32000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 3e-07,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 128000,
"max_output_tokens": 128000
},
"openrouter/mistralai/mixtral-8x22b-instruct": {
"input_cost_per_token": 6.5e-07,
@ -26673,7 +26732,9 @@
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 6.5e-07,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 65536,
"max_output_tokens": 65536
},
"openrouter/moonshotai/kimi-k2.5": {
"cache_read_input_token_cost": 1e-07,
@ -26693,26 +26754,32 @@
"openrouter/openai/gpt-3.5-turbo": {
"input_cost_per_token": 1.5e-06,
"litellm_provider": "openrouter",
"max_tokens": 4095,
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 2e-06,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 16385,
"max_output_tokens": 4096
},
"openrouter/openai/gpt-3.5-turbo-16k": {
"input_cost_per_token": 3e-06,
"litellm_provider": "openrouter",
"max_tokens": 16383,
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 4e-06,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 16385,
"max_output_tokens": 4096
},
"openrouter/openai/gpt-4": {
"input_cost_per_token": 3e-05,
"litellm_provider": "openrouter",
"max_tokens": 8192,
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 6e-05,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 8191,
"max_output_tokens": 4096
},
"openrouter/openai/gpt-4.1": {
"cache_read_input_token_cost": 5e-07,
@ -27220,10 +27287,12 @@
"openrouter/undi95/remm-slerp-l2-13b": {
"input_cost_per_token": 1.875e-06,
"litellm_provider": "openrouter",
"max_tokens": 6144,
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 1.875e-06,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 6144,
"max_output_tokens": 4096
},
"openrouter/x-ai/grok-4": {
"input_cost_per_token": 3e-06,
@ -28026,7 +28095,8 @@
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false,
"supports_function_calling": true
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true
},
"perplexity/anthropic/claude-sonnet-4-5": {
"litellm_provider": "perplexity",
@ -29844,14 +29914,16 @@
"together_ai/deepseek-ai/DeepSeek-V3.1": {
"input_cost_per_token": 6e-07,
"litellm_provider": "together_ai",
"max_tokens": 128000,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 1.7e-06,
"source": "https://www.together.ai/models/deepseek-v3-1",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"max_input_tokens": 128000,
"max_output_tokens": 16384
},
"together_ai/meta-llama/Llama-3.2-3B-Instruct-Turbo": {
"litellm_provider": "together_ai",
@ -30503,6 +30575,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -30531,6 +30604,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -30558,6 +30632,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -31138,6 +31213,7 @@
"output_cost_per_token": 2.5e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_minimal_reasoning_effort": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -32330,6 +32406,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -32356,6 +32433,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -32522,6 +32600,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -34426,6 +34505,7 @@
"output_cost_per_token": 1.5e-05,
"source": "https://x.ai/api#pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": false,
"supports_tool_choice": true,
"supports_web_search": true
@ -34441,6 +34521,7 @@
"output_cost_per_token": 1.5e-05,
"source": "https://x.ai/api#pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": false,
"supports_tool_choice": true,
"supports_web_search": true
@ -34456,6 +34537,7 @@
"output_cost_per_token": 2.5e-05,
"source": "https://x.ai/api#pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": false,
"supports_tool_choice": true,
"supports_web_search": true
@ -34471,6 +34553,7 @@
"output_cost_per_token": 2.5e-05,
"source": "https://x.ai/api#pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": false,
"supports_tool_choice": true,
"supports_web_search": true
@ -34486,6 +34569,7 @@
"output_cost_per_token": 1.5e-05,
"source": "https://x.ai/api#pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": false,
"supports_tool_choice": true,
"supports_web_search": true
@ -34502,6 +34586,7 @@
"output_cost_per_token": 5e-07,
"source": "https://x.ai/api#pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": false,
"supports_tool_choice": true,
@ -34519,6 +34604,7 @@
"output_cost_per_token": 5e-07,
"source": "https://x.ai/api#pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": false,
"supports_tool_choice": true,
@ -34535,6 +34621,7 @@
"output_cost_per_token": 4e-06,
"source": "https://x.ai/api#pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": false,
"supports_tool_choice": true,
@ -34551,6 +34638,7 @@
"output_cost_per_token": 4e-06,
"source": "https://x.ai/api#pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": false,
"supports_tool_choice": true,
@ -34567,6 +34655,7 @@
"output_cost_per_token": 4e-06,
"source": "https://x.ai/api#pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": false,
"supports_tool_choice": true,
@ -34583,6 +34672,7 @@
"output_cost_per_token": 5e-07,
"source": "https://x.ai/api#pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": false,
"supports_tool_choice": true,
@ -34598,38 +34688,41 @@
"output_cost_per_token": 1.5e-05,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true,
"supports_web_search": true
},
"xai/grok-4-fast-reasoning": {
"cache_read_input_token_cost": 5e-08,
"input_cost_per_token": 2e-07,
"input_cost_per_token_above_128k_tokens": 4e-07,
"litellm_provider": "xai",
"max_input_tokens": 2000000.0,
"max_output_tokens": 2000000.0,
"max_tokens": 2000000.0,
"mode": "chat",
"input_cost_per_token": 2e-07,
"input_cost_per_token_above_128k_tokens": 4e-07,
"output_cost_per_token": 5e-07,
"output_cost_per_token_above_128k_tokens": 1e-06,
"cache_read_input_token_cost": 5e-08,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true,
"supports_web_search": true
},
"xai/grok-4-fast-non-reasoning": {
"cache_read_input_token_cost": 5e-08,
"input_cost_per_token": 2e-07,
"input_cost_per_token_above_128k_tokens": 4e-07,
"litellm_provider": "xai",
"max_input_tokens": 2000000.0,
"max_output_tokens": 2000000.0,
"cache_read_input_token_cost": 5e-08,
"max_tokens": 2000000.0,
"mode": "chat",
"input_cost_per_token": 2e-07,
"input_cost_per_token_above_128k_tokens": 4e-07,
"output_cost_per_token": 5e-07,
"output_cost_per_token_above_128k_tokens": 1e-06,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true,
"supports_web_search": true
},
@ -34645,6 +34738,7 @@
"output_cost_per_token_above_128k_tokens": 3e-05,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true,
"supports_web_search": true
},
@ -34660,6 +34754,7 @@
"output_cost_per_token_above_128k_tokens": 3e-05,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true,
"supports_web_search": true
},
@ -34677,6 +34772,7 @@
"source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning",
"supports_audio_input": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
@ -34697,6 +34793,7 @@
"source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning",
"supports_audio_input": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
@ -34717,6 +34814,7 @@
"source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning",
"supports_audio_input": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
@ -34737,6 +34835,7 @@
"source": "https://docs.x.ai/docs/models/grok-4-1-fast-non-reasoning",
"supports_audio_input": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
@ -34756,6 +34855,7 @@
"source": "https://docs.x.ai/docs/models/grok-4-1-fast-non-reasoning",
"supports_audio_input": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
@ -34772,6 +34872,7 @@
"output_cost_per_token": 6e-06,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": true,
@ -34788,6 +34889,7 @@
"output_cost_per_token": 6e-06,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": true,
@ -34820,6 +34922,7 @@
"output_cost_per_token": 6e-06,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
@ -34848,6 +34951,7 @@
"output_cost_per_token": 1.5e-06,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
@ -34862,6 +34966,7 @@
"output_cost_per_token": 1.5e-06,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
@ -34876,6 +34981,7 @@
"output_cost_per_token": 1.5e-06,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
@ -34907,6 +35013,20 @@
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"zai.glm-5": {
"input_cost_per_token": 1e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 3.2e-06,
"supports_function_calling": true,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"zai.glm-4.7-flash": {
"input_cost_per_token": 7e-08,
"litellm_provider": "bedrock_converse",
@ -39492,6 +39612,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -39716,6 +39837,87 @@
}
]
},
"zai.glm-5": {
"input_cost_per_token": 1e-06,
"output_cost_per_token": 3.2e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"bedrock/us-east-1/zai.glm-5": {
"input_cost_per_token": 1e-06,
"output_cost_per_token": 3.2e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 200000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"bedrock/us-west-2/zai.glm-5": {
"input_cost_per_token": 1e-06,
"output_cost_per_token": 3.2e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 200000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"minimax.minimax-m2.5": {
"input_cost_per_token": 3e-07,
"output_cost_per_token": 1.2e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"bedrock/us-east-1/minimax.minimax-m2.5": {
"input_cost_per_token": 3e-07,
"output_cost_per_token": 1.2e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"bedrock/us-west-2/minimax.minimax-m2.5": {
"input_cost_per_token": 3e-07,
"output_cost_per_token": 1.2e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"bedrock/us-gov-east-1/anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.5e-06,
"cache_read_input_token_cost": 1.2e-07,

View file

@ -190,6 +190,8 @@ class LitellmTableNames(str, enum.Enum):
PROXY_MODEL_TABLE_NAME = "LiteLLM_ProxyModelTable"
MANAGED_FILE_TABLE_NAME = "LiteLLM_ManagedFileTable"
TOOL_TABLE_NAME = "LiteLLM_ToolTable"
CACHE_CONFIG_TABLE_NAME = "LiteLLM_CacheConfig"
CONFIG_OVERRIDES_TABLE_NAME = "LiteLLM_ConfigOverrides"
class Litellm_EntityType(enum.Enum):
@ -565,6 +567,7 @@ class LiteLLMRoutes(enum.Enum):
"/team/available",
"/team/permissions_list",
"/team/permissions_update",
"/team/permissions_bulk_update",
"/team/daily/activity",
# model
"/model/new",
@ -614,13 +617,12 @@ class LiteLLMRoutes(enum.Enum):
"/",
"/health/liveliness",
"/health/liveness",
"/health/readiness",
"/test",
"/config/yaml",
"/metrics",
"/litellm/.well-known/litellm-ui-config",
"/.well-known/litellm-ui-config",
"/public/model_hub",
"/public/model_hub/info",
"/public/agent_hub",
"/public/mcp_hub",
"/public/skill_hub",

View file

@ -31,6 +31,7 @@ from litellm.types.agents import (
from litellm.litellm_core_utils.litellm_logging import _get_masked_values
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.proxy.management_endpoints.common_daily_activity import (
DailySpendMetadata,
SpendAnalyticsPaginatedResponse,
)
@ -973,7 +974,64 @@ async def get_agent_daily_activity(
exclude_agent_ids.split(",") if exclude_agent_ids else None
)
where_condition = {}
# Without scoping, an empty `agent_ids` query returned every agent's
# spend/token rows on the proxy. Restrict non-admin callers to the
# agents they're permitted to invoke (or that they created), and
# intersect their explicit `agent_ids` filter with the same allowlist.
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
AgentRequestHandler,
)
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
where_condition: Dict[str, Any] = {}
if not _user_has_admin_view(user_api_key_dict):
permitted_agent_ids = await AgentRequestHandler.get_allowed_agents(
user_api_key_auth=user_api_key_dict
)
# `get_allowed_agents` returns an empty list when the caller's key
# and team carry no agent restrictions. For activity scoping that's
# not "see everything" — fall back to the agents the caller
# created so they cannot enumerate other tenants' agents.
# Guard against `user_id is None`: a literal None in Prisma
# `where={"created_by": None}` resolves to ``created_by IS NULL``
# and would expose every ownerless agent's rows.
if not permitted_agent_ids:
if user_api_key_dict.user_id is None:
permitted_agent_ids = []
else:
owned_records = await prisma_client.db.litellm_agentstable.find_many(
where={"created_by": user_api_key_dict.user_id}
)
permitted_agent_ids = [a.agent_id for a in owned_records]
if agent_ids_list:
permitted_agent_id_set = set(permitted_agent_ids)
agent_ids_list = [
aid for aid in agent_ids_list if aid in permitted_agent_id_set
]
else:
agent_ids_list = list(permitted_agent_ids)
# No accessible agents → return an empty page without querying.
if not agent_ids_list:
return SpendAnalyticsPaginatedResponse(
results=[],
metadata=DailySpendMetadata(
total_spend=0.0,
total_prompt_tokens=0,
total_completion_tokens=0,
total_tokens=0,
total_api_requests=0,
total_successful_requests=0,
total_failed_requests=0,
total_cache_read_input_tokens=0,
total_cache_creation_input_tokens=0,
page=page,
total_pages=0,
has_more=False,
),
)
if agent_ids_list:
where_condition["agent_id"] = {"in": list(agent_ids_list)}

View file

@ -3547,7 +3547,13 @@ async def _check_team_member_budget(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
if default_budget is not None:
# Treat 0 on the team default as "no cap".
# Per-member rows still respect 0 as an explicit admin disable.
if (
default_budget is not None
and default_budget.max_budget is not None
and default_budget.max_budget > 0
):
team_member_budget = default_budget.max_budget
if team_member_budget is not None:

View file

@ -2,7 +2,7 @@ import os
import re
import sys
from functools import lru_cache
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union
from typing import Any, Dict, FrozenSet, List, Mapping, Optional, Tuple, Union
from fastapi import HTTPException, Request, status
@ -173,9 +173,71 @@ def _allow_model_level_clientside_configurable_parameters(
# threat shape should be added here.
_NESTED_CONFIG_KEYS: Tuple[str, ...] = ("litellm_embedding_config",)
# Banned root-level params. Same list applies to every entry in
# ``_NESTED_CONFIG_KEYS`` because those dicts get spread as ``**kwargs``
# into the same outbound calls.
# Metadata containers that carry per-request configuration consumed by the
# observability callbacks. The same banned-param list applies — a value
# under ``metadata.langfuse_host`` redirects the same Langfuse client and
# leaks the same credentials as the root-level ``langfuse_host``, but the
# original check only walked the request-body root, so the metadata path
# was an unintentional bypass.
_NESTED_METADATA_KEYS: Tuple[str, ...] = ("metadata", "litellm_metadata")
# Banned request-body params. The same list applies to every entry in
# ``_NESTED_CONFIG_KEYS`` (dicts spread as ``**kwargs`` into outbound
# calls) and ``_NESTED_METADATA_KEYS`` (dicts read directly by integration
# callbacks), so a single banned name is enforced wherever the field can
# reach the call path from.
# Per-request observability params that are SAFE to accept from clients.
# These describe the request being logged (prompt version, sampling rate)
# without choosing the destination or the credentials, so they don't
# contribute to the data-exfil primitive that the rest of
# ``_supported_callback_params`` does.
_SAFE_CLIENT_CALLBACK_PARAMS: FrozenSet[str] = frozenset(
{
"langfuse_prompt_version",
"langsmith_sampling_rate",
}
)
# Observability fields that integrations read from the request body or
# metadata but that are not (yet) listed in ``_supported_callback_params``.
# Listed here so the proxy bans them today; the long-term cleanup is to
# fold these into the canonical allowlist so they share one source of
# truth with the rest.
_EXTRA_BANNED_OBSERVABILITY_PARAMS: FrozenSet[str] = frozenset(
{
"posthog_api_url",
"phoenix_project_name",
"wandb_api_key",
"weave_project_id",
}
)
def _build_banned_observability_params() -> FrozenSet[str]:
"""Derive the observability ban list from the canonical allowlist.
``_supported_callback_params`` and ``_request_blocked_callback_params`` in
``litellm/litellm_core_utils/initialize_dynamic_callback_params.py`` is
the single place that enumerates every observability field integrations
resolve from kwargs/metadata, plus fields that integration code explicitly
blocks from request-supplied callback params. Subtract the small set of
informational fields (``_SAFE_CLIENT_CALLBACK_PARAMS``) and union with the
extras the canonical allowlist hasn't caught up to yet. New integrations
added to the canonical allowlist are banned by default, which is the safe
failure mode.
"""
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
_request_blocked_callback_params,
_supported_callback_params,
)
return (
(frozenset(_supported_callback_params) - _SAFE_CLIENT_CALLBACK_PARAMS)
| frozenset(_request_blocked_callback_params)
| _EXTRA_BANNED_OBSERVABILITY_PARAMS
)
_BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = (
"api_base",
"base_url",
@ -190,11 +252,6 @@ _BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = (
# tokens) to the attacker's host, or coerces the proxy into
# authenticating against the attacker's host with admin secrets.
"aws_bedrock_runtime_endpoint",
"langsmith_base_url",
"langfuse_host",
"posthog_host",
"braintrust_host",
"slack_webhook_url",
# Provider-specific endpoint overrides that flow into the outbound
# request via ``optional_params``. Same threat as ``api_base``:
# ``s3_endpoint_url`` redirects Bedrock file uploads to attacker
@ -203,6 +260,11 @@ _BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = (
"s3_endpoint_url",
"sagemaker_base_url",
"deployment_url",
# Observability credentials, hosts, and project identifiers: derived
# from the canonical ``_supported_callback_params`` allowlist so new
# integrations are covered automatically. Sorted for stable iteration
# order and reviewable diffs.
*sorted(_build_banned_observability_params()),
)
@ -221,6 +283,8 @@ def _check_banned_params(
if param not in body:
continue
if general_settings.get("allow_client_side_credentials") is True:
# Proxy-wide opt-in: every banned param is permitted, exit
# entirely so the rest of the loop doesn't waste work.
return
if (
_allow_model_level_clientside_configurable_parameters(
@ -231,7 +295,12 @@ def _check_banned_params(
)
is True
):
return
# Per-param opt-in: only THIS param is permitted by the
# deployment's ``configurable_clientside_auth_params``. Skip
# to the next banned param so a body that pairs an allowed
# ``api_base`` with an unallowed ``langfuse_host`` is still
# rejected for the second field.
continue
raise ValueError(
f"Rejected Request: {param} is not allowed in request body. "
"Clientside passthrough requires explicit admin opt-in via "
@ -275,9 +344,33 @@ def is_request_body_safe(
nested = request_body.get(nested_key)
if isinstance(nested, dict):
_check_banned_params(nested, general_settings, llm_router, model)
for metadata_key in _NESTED_METADATA_KEYS:
metadata = _coerce_metadata_to_dict(request_body.get(metadata_key))
if metadata is not None:
_check_banned_params(metadata, general_settings, llm_router, model)
return True
def _coerce_metadata_to_dict(value: Any) -> Optional[Dict[str, Any]]:
"""Return ``value`` as a dict, parsing it from JSON if delivered as a string.
Multipart/form-data and ``extra_body`` callers send ``litellm_metadata``
as a JSON-encoded string; the proxy parses it into a dict later in
``add_litellm_data_to_request``, but the auth-time bouncer runs first
and would otherwise miss the banned-param check on a still-stringified
metadata blob.
"""
if isinstance(value, dict):
return value
if isinstance(value, str):
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
parsed = safe_json_loads(value)
if isinstance(parsed, dict):
return parsed
return None
async def pre_db_read_auth_checks(
request: Request,
request_data: dict,

View file

@ -6,6 +6,7 @@ from fastapi import HTTPException, Request, status
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import (
CommonProxyErrors,
KeyManagementRoutes,
LiteLLM_UserTable,
LiteLLMRoutes,
LitellmUserRoles,
@ -14,6 +15,49 @@ from litellm.proxy._types import (
from .auth_checks_organization import _user_is_org_admin
# Management write routes denied to PROXY_ADMIN_VIEW_ONLY. Adding a new write
# endpoint to a management router REQUIRES adding it here too — the surrounding
# check falls through to "allow" if the route is not matched, which previously
# let view-only admins call /team/block, /team/unblock, /key/bulk_update, etc.
_PROXY_ADMIN_VIEW_ONLY_BLOCKED_ROUTES = frozenset(
[
# user
"/user/new",
"/user/delete",
"/user/bulk_update",
# team
"/team/new",
"/team/update",
"/team/delete",
"/team/block",
"/team/unblock",
"/team/permissions_update",
"/team/permissions_bulk_update",
# model
"/model/new",
"/model/update",
"/model/delete",
# JWT key mapping
"/jwt/key/mapping/new",
"/jwt/key/mapping/update",
"/jwt/key/mapping/delete",
# key management — keep in sync with KeyManagementRoutes write entries
KeyManagementRoutes.KEY_GENERATE.value,
KeyManagementRoutes.KEY_UPDATE.value,
KeyManagementRoutes.KEY_DELETE.value,
KeyManagementRoutes.KEY_REGENERATE.value,
KeyManagementRoutes.KEY_GENERATE_SERVICE_ACCOUNT.value,
KeyManagementRoutes.KEY_BLOCK.value,
KeyManagementRoutes.KEY_UNBLOCK.value,
KeyManagementRoutes.KEY_BULK_UPDATE.value,
]
)
# Suffixes for `/key/{key_id}/...` path-parameterized write routes that the
# enum templates with `{key_id}`. The blocklist above can't match templated
# paths directly because the request route carries the resolved key id.
_PROXY_ADMIN_VIEW_ONLY_BLOCKED_KEY_SUFFIXES = ("/regenerate", "/reset_spend")
class RouteChecks:
@staticmethod
@ -664,6 +708,31 @@ class RouteChecks:
detail=f"user not allowed to access this OpenAI routes, role= {_user_role}",
)
# Check if this is a write operation on management routes
if RouteChecks.check_route_access(
route=route, allowed_routes=LiteLLMRoutes.management_routes.value
):
# For management routes, only allow read operations or specific allowed updates
if route == "/user/update":
# Check the Request params are valid for PROXY_ADMIN_VIEW_ONLY
if request_data is not None and isinstance(request_data, dict):
_params_updated = request_data.keys()
for param in _params_updated:
if param not in ["user_email", "password"]:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"user not allowed to access this route, role= {_user_role}. Trying to access: {route} and updating invalid param: {param}. only user_email and password can be updated",
)
elif route in _PROXY_ADMIN_VIEW_ONLY_BLOCKED_ROUTES or (
route.startswith("/key/")
and route.endswith(_PROXY_ADMIN_VIEW_ONLY_BLOCKED_KEY_SUFFIXES)
):
# Block write operations for PROXY_ADMIN_VIEW_ONLY
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"user not allowed to access this route, role= {_user_role}. Trying to access: {route}",
)
# Allow read operations on management routes (like /user/info, /team/info, /model/info)
method = request.method.upper() if request is not None else "GET"
is_safe_method = method in RouteChecks._SAFE_HTTP_METHODS

View file

@ -87,6 +87,23 @@ except ImportError as e:
user_api_key_service_logger_obj = ServiceLogging() # used for tracking latency on OTEL
def _normalize_public_auth_route(route: str) -> str:
if route != "/" and route.endswith("/"):
return route.rstrip("/")
return route
def _route_requires_auth_despite_public(
route: str, general_settings: Optional[dict]
) -> bool:
normalized_route = _normalize_public_auth_route(route)
if normalized_route == "/metrics":
return litellm.require_auth_for_metrics_endpoint is not False
return False
custom_litellm_key_header = APIKeyHeader(
name=SpecialHeaders.custom_litellm_api_key.value,
auto_error=False,
@ -714,7 +731,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
"""
######## Route Checks Before Reading DB / Cache for "token" ################
if (
if not _route_requires_auth_despite_public(
route=route, general_settings=general_settings
) and (
route in LiteLLMRoutes.public_routes.value # type: ignore
or route_in_additonal_public_routes(current_route=route)
):
@ -1698,7 +1717,7 @@ async def _run_centralized_common_checks(
user_custom_auth,
)
# Public routes (e.g. /health/readiness, /metrics) are exempt from
# Public routes (e.g. /health/liveness) are exempt from
# auth in the builder — the wrapper must not retroactively apply
# authz on top, or k8s readiness probes and other unauthenticated
# callers get 401.

View file

@ -50,7 +50,10 @@ def configure_gc_thresholds():
configure_gc_thresholds()
@router.get("/debug/asyncio-tasks")
@router.get(
"/debug/asyncio-tasks",
dependencies=[Depends(user_api_key_auth)],
)
async def get_active_tasks_stats():
"""
Returns:
@ -103,7 +106,11 @@ if os.environ.get("LITELLM_PROFILE", "false").lower() == "true":
tracemalloc.start(10)
@router.get("/memory-usage", include_in_schema=False)
@router.get(
"/memory-usage",
dependencies=[Depends(user_api_key_auth)],
include_in_schema=False,
)
async def memory_usage():
# Take a snapshot of the current memory usage
snapshot = tracemalloc.take_snapshot()
@ -711,7 +718,11 @@ async def configure_gc_thresholds_endpoint(
}
@router.get("/otel-spans", include_in_schema=False)
@router.get(
"/otel-spans",
dependencies=[Depends(user_api_key_auth)],
include_in_schema=False,
)
async def get_otel_spans():
from litellm.proxy.proxy_server import open_telemetry_logger

View file

@ -0,0 +1,236 @@
"""
Shared helpers for guardrail hooks: extract text from a request body
regardless of whether it uses Chat Completions ``messages``, Responses-API
``input``, or multimodal list-format ``content`` parts.
Hooks that only check ``data["messages"]`` for string content silently
skip the other shapes — these helpers normalise that so every hook sees
every text fragment.
"""
from typing import Any, Callable, Dict, FrozenSet, Iterator, List
# Call types whose body carries free-form chat / prompt text that
# text-content guardrails (banned keywords, content moderation, secret
# detection, …) should inspect. The proxy ingress passes ``route_type``
# straight through as ``call_type``, so the literal values here are
# what the guardrail dispatcher actually receives:
#
# /v1/chat/completions -> "acompletion"
# /v1/responses -> "aresponses"
#
# ``"completion"`` is included for SDK / internal callers that invoke
# ``pre_call_hook`` directly with the sync name. Embedding, moderation,
# audio, and transcription endpoints are deliberately excluded — text
# guardrails on those paths are a separate scope.
TEXT_CONTENT_CALL_TYPES: FrozenSet[str] = frozenset(
{"completion", "acompletion", "aresponses"}
)
def is_text_content_call_type(call_type: str) -> bool:
"""Return True if ``call_type`` carries free-form text that text
guardrails should inspect (Chat Completions or Responses API)."""
return call_type in TEXT_CONTENT_CALL_TYPES
def _iter_text_parts_in_content(content: Any) -> Iterator[str]:
"""Yield text fragments from a ``message.content`` value (string or
multimodal list). Non-text parts (images, audio, …) are skipped."""
if isinstance(content, str):
if content:
yield content
elif isinstance(content, list):
for part in content:
if isinstance(part, str):
# A bare string in a content/input list is itself a text
# fragment (Responses-API mixed-list shape).
if part:
yield part
continue
if not isinstance(part, dict):
continue
if part.get("type") == "text":
text = part.get("text")
if isinstance(text, str) and text:
yield text
def _coerce_input_to_messages(input_value: Any) -> List[Dict[str, Any]]:
"""Coerce a Responses-API ``data["input"]`` value into chat-style messages."""
if isinstance(input_value, str):
return [{"role": "user", "content": input_value}]
if isinstance(input_value, list):
if input_value and all(
isinstance(item, dict) and "role" in item for item in input_value
):
return list(input_value)
# Mixed lists (content-part dicts + bare strings) and pure
# string/dict lists all become a single user message; the content
# iterator below handles each element type uniformly.
return [{"role": "user", "content": input_value}]
return []
def _iter_inspection_messages(data: Dict[str, Any]) -> Iterator[Dict[str, Any]]:
"""Yield every message-like dict, walking ``messages`` AND ``input``."""
messages = data.get("messages")
if isinstance(messages, list):
yield from messages
yield from _coerce_input_to_messages(data.get("input"))
def iter_message_text(data: Dict[str, Any]) -> Iterator[str]:
"""Yield every text fragment from ``messages`` AND ``input``.
Walks every role (user, assistant, system, …) — guardrails inspect
the entire conversation, not just user turns.
"""
for message in _iter_inspection_messages(data):
if not isinstance(message, dict):
continue
yield from _iter_text_parts_in_content(message.get("content"))
def walk_user_text(data: Dict[str, Any], visit: Callable[[str], str]) -> int:
"""Rewrite every text fragment in place via ``visit``.
Mutates ``data["messages"]`` and ``data["input"]``. Returns the number
of fragments visited so callers can short-circuit when nothing was
inspected.
"""
visited = 0
def _rewrite_content(content: Any) -> Any:
nonlocal visited
if isinstance(content, str):
if content:
visited += 1
return visit(content)
return content
if isinstance(content, list):
new_parts: List[Any] = []
for part in content:
if isinstance(part, str) and part:
visited += 1
new_parts.append(visit(part))
elif (
isinstance(part, dict)
and part.get("type") == "text"
and isinstance(part.get("text"), str)
and part["text"]
):
visited += 1
new_parts.append({**part, "text": visit(part["text"])})
else:
new_parts.append(part)
return new_parts
return content
messages = data.get("messages")
if isinstance(messages, list):
for message in messages:
if isinstance(message, dict) and "content" in message:
message["content"] = _rewrite_content(message["content"])
input_value = data.get("input")
if isinstance(input_value, str):
if input_value:
visited += 1
data["input"] = visit(input_value)
return visited
if isinstance(input_value, list):
# List of full messages: rewrite each message's content.
if input_value and all(
isinstance(item, dict) and "role" in item for item in input_value
):
for item in input_value:
if "content" in item:
item["content"] = _rewrite_content(item["content"])
return visited
# List of content parts and/or bare strings: rewrite in place.
for idx, item in enumerate(input_value):
if isinstance(item, str) and item:
visited += 1
input_value[idx] = visit(item)
elif (
isinstance(item, dict)
and item.get("type") == "text"
and isinstance(item.get("text"), str)
and item["text"]
):
visited += 1
input_value[idx] = {**item, "text": visit(item["text"])}
return visited
return visited
def apply_redacted_messages_back(
data: Dict[str, Any], redacted_messages: List[Dict[str, Any]]
) -> None:
"""Write redacted messages back to whichever field(s) the caller used.
Mask/anonymize paths take a synthesised messages list (from
:func:`build_inspection_messages`), get a redacted version back from a
third-party guardrail, and need to rewrite the request body. Writing
only to ``data["messages"]`` leaves the Responses-API ``data["input"]``
field untouched, so the unredacted text still reaches the LLM.
This helper updates both fields when both are present.
"""
if "messages" in data:
data["messages"] = redacted_messages
if isinstance(data.get("input"), str):
text_parts: List[str] = []
for msg in redacted_messages:
if not isinstance(msg, dict):
continue
text_parts.extend(_iter_text_parts_in_content(msg.get("content")))
data["input"] = "\n".join(text_parts)
def has_non_string_content(data: Dict[str, Any]) -> bool:
"""Return True if any inspected content is not a plain string.
Used by hooks whose mask/redact path operates on string offsets and
therefore cannot preserve multimodal non-text parts. Such hooks should
degrade to block-on-detect when this returns True so image/audio parts
are not silently stripped during in-place masking.
"""
messages = data.get("messages")
if isinstance(messages, list):
for message in messages:
if isinstance(message, dict) and not isinstance(
message.get("content"), str
):
if message.get("content") is not None:
return True
input_value = data.get("input")
if input_value is not None and not isinstance(input_value, str):
return True
return False
def build_inspection_messages(data: Dict[str, Any]) -> List[Dict[str, str]]:
"""Synthesize a chat-style messages list for posting to a guardrail API.
Each returned message has a plain-string ``content`` — multimodal text
parts are joined with newlines and Responses-API ``input`` is lifted
into synthetic messages. Messages with no inspectable text are dropped.
Hooks that POST ``{"messages": [...]}`` to an external service should
call this instead of ``data.get("messages", [])`` so the Responses API
and multimodal content are covered.
"""
flattened: List[Dict[str, str]] = []
for message in _iter_inspection_messages(data):
if not isinstance(message, dict):
continue
text = "\n".join(_iter_text_parts_in_content(message.get("content")))
if not text:
continue
role = message.get("role", "user") or "user"
flattened.append({"role": role, "content": text})
return flattened

View file

@ -22,6 +22,11 @@ from litellm.llms.custom_httpx.http_handler import (
httpxSpecialProvider,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import (
apply_redacted_messages_back,
build_inspection_messages,
has_non_string_content,
)
from litellm.types.utils import (
CallTypesLiteral,
Choices,
@ -101,10 +106,11 @@ class AimGuardrail(CustomGuardrail):
user_email=user_email,
litellm_call_id=call_id,
)
# Covers multimodal list content + Responses-API input.
response = await self.async_handler.post(
f"{self.api_base}/fw/v1/analyze",
headers=headers,
json={"messages": data.get("messages", [])},
json={"messages": build_inspection_messages(data)},
)
response.raise_for_status()
res = response.json()
@ -137,13 +143,31 @@ class AimGuardrail(CustomGuardrail):
redacted_chat = res.get("redacted_chat")
if not redacted_chat:
return data
data["messages"] = [
# Aim returns text-only redacted messages. Overwriting
# ``data["messages"]`` with that would silently strip image/audio
# parts from a multimodal request — degrade to block so the
# multimodal payload is never silently rewritten.
if has_non_string_content(data):
raise HTTPException(
status_code=400,
detail=(
"Aim: anonymize action requested for multimodal input "
"but mask-in-place would drop non-text parts. Send the "
"request with plain string content to use anonymize, "
"or rely on block-mode policies."
),
)
redacted_messages = [
{
"role": message["role"],
"content": message["content"],
}
for message in redacted_chat["all_redacted_messages"]
]
# Write back to ``messages`` AND ``input``. The Responses-API
# backend reads ``input``; writing only to ``messages`` would let
# unredacted text reach the LLM for ``/v1/responses`` calls.
apply_redacted_messages_back(data, redacted_messages)
return data
async def call_aim_guardrail_on_output(
@ -162,7 +186,7 @@ class AimGuardrail(CustomGuardrail):
litellm_call_id=call_id,
),
json={
"messages": request_data.get("messages", [])
"messages": build_inspection_messages(request_data)
+ [{"role": "assistant", "content": output}]
},
)
@ -233,15 +257,33 @@ class AimGuardrail(CustomGuardrail):
user_api_key_dict: UserAPIKeyAuth,
response: Union[Any, ModelResponse, EmbeddingResponse, ImageResponse],
) -> Any:
if (
isinstance(response, ModelResponse)
and response.choices
and isinstance(response.choices[0], Choices)
):
content = response.choices[0].message.content or ""
aim_output_guardrail_result = await self.call_aim_guardrail_on_output(
data, content, hook="output", key_alias=user_api_key_dict.key_alias
)
if not (isinstance(response, ModelResponse) and response.choices):
return response
# Inspect every choice — when ``n>1`` the additional completions
# used to bypass Aim entirely because the hook only inspected
# ``choices[0]``. Run inspections concurrently so multi-completion
# responses don't pay an n× latency penalty.
choices_to_inspect = [c for c in response.choices if isinstance(c, Choices)]
if not choices_to_inspect:
return response
# ``return_exceptions=True`` lets every inspection finish even if
# one fails — without it, the first exception would propagate and
# leave the remaining tasks running in the background.
results = await asyncio.gather(
*(
self.call_aim_guardrail_on_output(
data,
choice.message.content or "",
hook="output",
key_alias=user_api_key_dict.key_alias,
)
for choice in choices_to_inspect
),
return_exceptions=True,
)
for choice, aim_output_guardrail_result in zip(choices_to_inspect, results):
if isinstance(aim_output_guardrail_result, BaseException):
raise aim_output_guardrail_result
if aim_output_guardrail_result and aim_output_guardrail_result.get(
"detection_message"
):
@ -252,7 +294,7 @@ class AimGuardrail(CustomGuardrail):
if aim_output_guardrail_result and aim_output_guardrail_result.get(
"redacted_output"
):
response.choices[0].message.content = aim_output_guardrail_result.get(
choice.message.content = aim_output_guardrail_result.get(
"redacted_output"
)
return response

View file

@ -254,15 +254,16 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
) -> Any:
from litellm.types.utils import Choices, ModelResponse
if (
isinstance(response, ModelResponse)
and response.choices
and isinstance(response.choices[0], Choices)
):
content = response.choices[0].message.content or ""
await self.async_make_request(
text=content,
)
if isinstance(response, ModelResponse) and response.choices:
for choice in response.choices:
if not isinstance(choice, Choices):
continue
content = _message_content_to_text(choice.message.content)
if not content:
continue
await self.async_make_request(
text=content,
)
return response
async def async_post_call_streaming_hook(
@ -279,3 +280,16 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
error_returned = json.dumps({"error": e.detail})
return f"data: {error_returned}\n\n"
def _message_content_to_text(content: Any) -> str:
if isinstance(content, str):
return content
if isinstance(content, list):
text_parts = [
item.get("text")
for item in content
if isinstance(item, dict) and isinstance(item.get("text"), str)
]
return "\n".join(part for part in text_parts if part)
return ""

View file

@ -20,6 +20,7 @@ from litellm.llms.custom_httpx.http_handler import (
httpxSpecialProvider,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import iter_message_text
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.guardrails.guardrail_hooks.ibm import (
IBMDetectorDetection,
@ -463,65 +464,53 @@ class IBMGuardrailDetector(CustomGuardrail):
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return data
_messages = data.get("messages")
if _messages:
contents_to_check: List[str] = []
for message in _messages:
_content = message.get("content")
if isinstance(_content, str):
contents_to_check.append(_content)
# Covers multimodal list content + Responses-API input.
contents_to_check: List[str] = list(iter_message_text(data))
if contents_to_check:
if self.is_detector_server:
# Call detector server with all contents at once
result = await self._call_detector_server(
contents=contents_to_check,
request_data=data,
event_type=GuardrailEventHooks.pre_call,
)
if contents_to_check:
if self.is_detector_server:
# Call detector server with all contents at once
result = await self._call_detector_server(
contents=contents_to_check,
verbose_proxy_logger.debug(
"IBM Detector Server async_pre_call_hook result: %s", result
)
# Check if any detections were found
has_violations = False
for message_detections in result:
filtered = self._filter_detections_by_threshold(message_detections)
if filtered:
has_violations = True
break
if has_violations and self.block_on_detection:
error_message = self._create_error_message_detector_server(result)
raise ValueError(error_message)
else:
# Call orchestrator for each content separately
for content in contents_to_check:
orchestrator_result = await self._call_orchestrator(
content=content,
request_data=data,
event_type=GuardrailEventHooks.pre_call,
)
verbose_proxy_logger.debug(
"IBM Detector Server async_pre_call_hook result: %s", result
"IBM Orchestrator async_pre_call_hook result: %s",
orchestrator_result,
)
# Check if any detections were found
has_violations = False
for message_detections in result:
filtered = self._filter_detections_by_threshold(
message_detections
)
if filtered:
has_violations = True
break
if has_violations and self.block_on_detection:
error_message = self._create_error_message_detector_server(
result
)
raise ValueError(error_message)
else:
# Call orchestrator for each content separately
for content in contents_to_check:
orchestrator_result = await self._call_orchestrator(
content=content,
request_data=data,
event_type=GuardrailEventHooks.pre_call,
)
verbose_proxy_logger.debug(
"IBM Orchestrator async_pre_call_hook result: %s",
orchestrator_result,
)
filtered = self._filter_detections_by_threshold(
filtered = self._filter_detections_by_threshold(orchestrator_result)
if filtered and self.block_on_detection:
error_message = self._create_error_message_orchestrator(
orchestrator_result
)
if filtered and self.block_on_detection:
error_message = self._create_error_message_orchestrator(
orchestrator_result
)
raise ValueError(error_message)
raise ValueError(error_message)
# Add guardrail to applied guardrails header
add_guardrail_to_applied_guardrails_header(
@ -550,65 +539,53 @@ class IBMGuardrailDetector(CustomGuardrail):
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return
_messages = data.get("messages")
if _messages:
contents_to_check: List[str] = []
for message in _messages:
_content = message.get("content")
if isinstance(_content, str):
contents_to_check.append(_content)
# Covers multimodal list content + Responses-API input.
contents_to_check: List[str] = list(iter_message_text(data))
if contents_to_check:
if self.is_detector_server:
# Call detector server with all contents at once
result = await self._call_detector_server(
contents=contents_to_check,
request_data=data,
event_type=GuardrailEventHooks.during_call,
)
if contents_to_check:
if self.is_detector_server:
# Call detector server with all contents at once
result = await self._call_detector_server(
contents=contents_to_check,
verbose_proxy_logger.debug(
"IBM Detector Server async_moderation_hook result: %s", result
)
# Check if any detections were found
has_violations = False
for message_detections in result:
filtered = self._filter_detections_by_threshold(message_detections)
if filtered:
has_violations = True
break
if has_violations and self.block_on_detection:
error_message = self._create_error_message_detector_server(result)
raise ValueError(error_message)
else:
# Call orchestrator for each content separately
for content in contents_to_check:
orchestrator_result = await self._call_orchestrator(
content=content,
request_data=data,
event_type=GuardrailEventHooks.during_call,
)
verbose_proxy_logger.debug(
"IBM Detector Server async_moderation_hook result: %s", result
"IBM Orchestrator async_moderation_hook result: %s",
orchestrator_result,
)
# Check if any detections were found
has_violations = False
for message_detections in result:
filtered = self._filter_detections_by_threshold(
message_detections
)
if filtered:
has_violations = True
break
if has_violations and self.block_on_detection:
error_message = self._create_error_message_detector_server(
result
)
raise ValueError(error_message)
else:
# Call orchestrator for each content separately
for content in contents_to_check:
orchestrator_result = await self._call_orchestrator(
content=content,
request_data=data,
event_type=GuardrailEventHooks.during_call,
)
verbose_proxy_logger.debug(
"IBM Orchestrator async_moderation_hook result: %s",
orchestrator_result,
)
filtered = self._filter_detections_by_threshold(
filtered = self._filter_detections_by_threshold(orchestrator_result)
if filtered and self.block_on_detection:
error_message = self._create_error_message_orchestrator(
orchestrator_result
)
if filtered and self.block_on_detection:
error_message = self._create_error_message_orchestrator(
orchestrator_result
)
raise ValueError(error_message)
raise ValueError(error_message)
# Add guardrail to applied guardrails header
add_guardrail_to_applied_guardrails_header(

View file

@ -13,6 +13,11 @@ from litellm.llms.custom_httpx.http_handler import (
httpxSpecialProvider,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import (
apply_redacted_messages_back,
build_inspection_messages,
has_non_string_content,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import AllMessageValues
@ -214,18 +219,26 @@ class LakeraAIGuardrail(CustomGuardrail):
)
return data
new_messages: Optional[List[AllMessageValues]] = data.get("messages")
if new_messages is None:
# Covers multimodal list content + Responses-API input.
new_messages = build_inspection_messages(data)
if not new_messages:
verbose_proxy_logger.warning(
"Lakera AI: not running guardrail. No messages in data"
"Lakera AI: not running guardrail. No inspectable text in data"
)
return data
# Mask-in-place uses offsets returned by Lakera and can only
# preserve non-text parts (images, audio, …) when the original
# content is a plain string. For multimodal/Responses-API input
# we degrade to block-on-detect so we never silently strip image
# parts while attempting to redact text.
is_multimodal_input = has_non_string_content(data)
#########################################################
########## 1. Make the Lakera AI v2 guard API request ##########
#########################################################
lakera_guardrail_response, masked_entity_count = await self.call_v2_guard(
messages=new_messages,
messages=new_messages, # type: ignore[arg-type]
request_data=data,
event_type=GuardrailEventHooks.pre_call,
)
@ -234,13 +247,20 @@ class LakeraAIGuardrail(CustomGuardrail):
########## 2. Handle flagged content ##########
#########################################################
if lakera_guardrail_response.get("flagged") is True:
# If only PII violations exist, mask the PII
if self._is_only_pii_violation(lakera_guardrail_response):
data["messages"] = self._mask_pii_in_messages(
messages=new_messages,
# If only PII violations exist, mask the PII (string input only).
if (
self._is_only_pii_violation(lakera_guardrail_response)
and not is_multimodal_input
):
redacted_messages = self._mask_pii_in_messages(
messages=new_messages, # type: ignore[arg-type]
lakera_response=lakera_guardrail_response,
masked_entity_count=masked_entity_count,
)
# Write back to ``messages`` AND ``input``. The Responses-API
# backend reads ``input``; writing only to ``messages``
# would let unredacted PII reach the LLM for /v1/responses.
apply_redacted_messages_back(data, list(redacted_messages)) # type: ignore[arg-type]
verbose_proxy_logger.debug(
"Lakera AI: Masked PII in messages instead of blocking request"
)
@ -252,7 +272,9 @@ class LakeraAIGuardrail(CustomGuardrail):
)
# Log violation but continue
elif self.on_flagged == "block":
# If there are other violations or not set to mask PII, raise exception
# Either non-PII violations, or PII on multimodal input
# (which cannot be masked in place without dropping
# image/audio parts) — raise the standard block error.
raise self._get_http_exception_for_blocked_guardrail(
lakera_guardrail_response
)
@ -280,18 +302,22 @@ class LakeraAIGuardrail(CustomGuardrail):
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return
new_messages: Optional[List[AllMessageValues]] = data.get("messages")
if new_messages is None:
new_messages = build_inspection_messages(data)
if not new_messages:
verbose_proxy_logger.warning(
"Lakera AI: not running guardrail. No messages in data"
"Lakera AI: not running guardrail. No inspectable text in data"
)
return
# See ``async_pre_call_hook`` — multimodal input degrades to
# block-on-detect because mask-in-place would drop image parts.
is_multimodal_input = has_non_string_content(data)
#########################################################
########## 1. Make the Lakera AI v2 guard API request ##########
#########################################################
lakera_guardrail_response, masked_entity_count = await self.call_v2_guard(
messages=new_messages,
messages=new_messages, # type: ignore[arg-type]
request_data=data,
event_type=GuardrailEventHooks.during_call,
)
@ -300,25 +326,28 @@ class LakeraAIGuardrail(CustomGuardrail):
########## 2. Handle flagged content ##########
#########################################################
if lakera_guardrail_response.get("flagged") is True:
# If only PII violations exist, mask the PII
if self._is_only_pii_violation(lakera_guardrail_response):
data["messages"] = self._mask_pii_in_messages(
messages=new_messages,
if (
self._is_only_pii_violation(lakera_guardrail_response)
and not is_multimodal_input
):
redacted_messages = self._mask_pii_in_messages(
messages=new_messages, # type: ignore[arg-type]
lakera_response=lakera_guardrail_response,
masked_entity_count=masked_entity_count,
)
# Write back to ``messages`` AND ``input``. The Responses-API
# backend reads ``input``; writing only to ``messages``
# would let unredacted PII reach the LLM for /v1/responses.
apply_redacted_messages_back(data, list(redacted_messages)) # type: ignore[arg-type]
verbose_proxy_logger.debug(
"Lakera AI: Masked PII in messages instead of blocking request"
)
else:
# Check on_flagged setting
if self.on_flagged == "monitor":
verbose_proxy_logger.warning(
"Lakera Guardrail: Monitoring mode - violation detected but allowing request"
)
# Log violation but continue
elif self.on_flagged == "block":
# If there are other violations or not set to mask PII, raise exception
raise self._get_http_exception_for_blocked_guardrail(
lakera_guardrail_response
)

View file

@ -50,6 +50,11 @@ from litellm.llms.custom_httpx.http_handler import (
httpxSpecialProvider,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import (
apply_redacted_messages_back,
build_inspection_messages,
has_non_string_content,
)
from litellm.types.guardrails import GuardrailEventHooks
import litellm
@ -366,16 +371,19 @@ class LassoGuardrail(CustomGuardrail):
LassoGuardrailAPIError: If the Lasso API call fails
HTTPException: If blocking violations are detected
"""
messages: List[Dict[str, str]] = data.get("messages", [])
# Covers multimodal list content + Responses-API input.
messages: List[Dict[str, str]] = build_inspection_messages(data)
if not messages:
return data
if self.mask:
# Lasso's classifix endpoint returns masked text that we copy back
# into ``data["messages"]``. For multimodal/Responses-API input we
# would silently strip image/audio parts, so fall back to the
# classify endpoint (which still raises on BLOCK actions) and
# leave the original payload intact.
if self.mask and not has_non_string_content(data):
return await self._handle_masking(data, cache, message_type, messages)
else:
return await self._handle_classification(
data, cache, message_type, messages
)
return await self._handle_classification(data, cache, message_type, messages)
async def _handle_classification(
self,
@ -413,8 +421,9 @@ class LassoGuardrail(CustomGuardrail):
self._process_lasso_response(response)
# Apply masking to messages if violations detected and masked messages are available
if response.get("violations_detected") and response.get("messages"):
data["messages"] = response["messages"]
redacted_messages = response.get("messages")
if response.get("violations_detected") and redacted_messages:
apply_redacted_messages_back(data, list(redacted_messages))
self._log_masking_applied(message_type, dict(response))
return data

View file

@ -1873,8 +1873,9 @@ class ContentFilterGuardrail(CustomGuardrail):
and the UI Request Lifecycle panel. Mirrors apply_guardrail's finally-block
contract.
"""
accumulated_full_text = ""
yielded_masked_text_len = 0
accumulated_text_by_choice: Dict[int, str] = {}
yielded_masked_text_len_by_choice: Dict[int, int] = {}
latest_detections_by_choice: Dict[int, List[ContentFilterDetection]] = {}
buffer_size = 50 # Increased buffer to catch patterns split across many chunks
start_time = datetime.now()
@ -1890,79 +1891,90 @@ class ContentFilterGuardrail(CustomGuardrail):
try:
async for item in response:
if isinstance(item, ModelResponseStream) and item.choices:
delta_content = ""
is_final = False
for choice in item.choices:
if hasattr(choice, "delta") and choice.delta:
content = getattr(choice.delta, "content", None)
if content and isinstance(content, str):
delta_content += content
if getattr(choice, "finish_reason", None):
is_final = True
if not (hasattr(choice, "delta") and choice.delta):
continue
accumulated_full_text += delta_content
choice_index = getattr(choice, "index", 0)
if not isinstance(choice_index, int):
choice_index = 0
# Check for blocking or apply masking
# Add a space at the end if it's the final chunk to trigger word boundaries (\b)
text_to_check = accumulated_full_text
if is_final:
text_to_check += " "
content = getattr(choice.delta, "content", None)
is_final = bool(getattr(choice, "finish_reason", None))
if isinstance(content, str) and content:
accumulated_text_by_choice[choice_index] = (
accumulated_text_by_choice.get(choice_index, "")
+ content
)
elif not is_final:
continue
try:
# Reset before each scan: _filter_single_text scans the
# whole accumulated buffer every chunk, so previous-chunk
# matches are guaranteed to be re-found. Keeping only the
# latest scan's detections avoids N× duplication in the
# final log row. BLOCK still records correctly because
# handlers append to detections before raising.
detections.clear()
masked_text = self._filter_single_text(
text_to_check, detections=detections
text_to_check = accumulated_text_by_choice.get(choice_index, "")
if not text_to_check:
continue
# Add a space at the end if it's the final chunk to trigger word boundaries (\b)
text_to_scan = text_to_check + (" " if is_final else "")
choice_detections: List[ContentFilterDetection] = []
try:
# _filter_single_text scans the whole accumulated
# choice buffer every chunk, so previous-chunk
# matches are guaranteed to be re-found. Keeping
# only each choice's latest scan avoids duplicate
# detections in the final log row.
masked_text = self._filter_single_text(
text_to_scan, detections=choice_detections
)
if is_final and masked_text.endswith(" "):
masked_text = masked_text[:-1]
latest_detections_by_choice[choice_index] = (
choice_detections
)
except HTTPException:
latest_detections_by_choice[choice_index] = (
choice_detections
)
raise
except Exception as e:
verbose_proxy_logger.error(
f"ContentFilterGuardrail: Error in masking: {e}"
)
masked_text = text_to_scan # Fallback to current text
# Determine how much can be safely yielded
if is_final:
safe_to_yield_len = len(masked_text)
else:
safe_to_yield_len = max(0, len(masked_text) - buffer_size)
yielded_masked_text_len = yielded_masked_text_len_by_choice.get(
choice_index, 0
)
if is_final and masked_text.endswith(" "):
masked_text = masked_text[:-1]
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.error(
f"ContentFilterGuardrail: Error in masking: {e}"
)
masked_text = text_to_check # Fallback to current text
if safe_to_yield_len > yielded_masked_text_len:
new_masked_content = masked_text[
yielded_masked_text_len:safe_to_yield_len
]
choice.delta.content = new_masked_content
yielded_masked_text_len_by_choice[choice_index] = (
safe_to_yield_len
)
else:
# Hold content by yielding empty content on this choice
# while preserving chunk metadata and other choices.
choice.delta.content = ""
# Determine how much can be safely yielded
if is_final:
safe_to_yield_len = len(masked_text)
else:
safe_to_yield_len = max(0, len(masked_text) - buffer_size)
if safe_to_yield_len > yielded_masked_text_len:
new_masked_content = masked_text[
yielded_masked_text_len:safe_to_yield_len
]
# Modify the chunk to contain only the new masked content
if (
item.choices
and hasattr(item.choices[0], "delta")
and item.choices[0].delta
):
item.choices[0].delta.content = new_masked_content
yielded_masked_text_len = safe_to_yield_len
yield item
else:
# Hold content by yielding empty content chunk (keeps metadata/structure)
if (
item.choices
and hasattr(item.choices[0], "delta")
and item.choices[0].delta
):
item.choices[0].delta.content = ""
yield item
yield item
else:
# Not a ModelResponseStream or no choices - yield as is
yield item
# Any remaining content (should have been handled by is_final, but just in case)
if yielded_masked_text_len < len(accumulated_full_text):
if any(
yielded_masked_text_len_by_choice.get(choice_index, 0)
< len(accumulated_text)
for choice_index, accumulated_text in accumulated_text_by_choice.items()
):
# We already reached the end of the generator
pass
except HTTPException:
@ -1973,6 +1985,11 @@ class ContentFilterGuardrail(CustomGuardrail):
exception_str = str(e)
raise e
finally:
detections = [
detection
for choice_detections in latest_detections_by_choice.values()
for detection in choice_detections
]
self._count_masked_entities(detections, masked_entity_count)
self._log_guardrail_information(
request_data=request_data,

View file

@ -187,11 +187,28 @@ def _extract_user_text(messages: List) -> str:
def _extract_response_text(response: Any) -> str:
"""Extract text from LLM response object."""
"""Extract text from every LLM response choice."""
if hasattr(response, "choices") and response.choices:
choice = response.choices[0]
if hasattr(choice, "message") and choice.message:
return choice.message.content or ""
text_parts: List[str] = []
for choice in response.choices:
if hasattr(choice, "message") and choice.message:
text = _content_to_text(choice.message.content)
if text:
text_parts.append(text)
return "\n".join(text_parts)
return ""
def _content_to_text(content: Any) -> str:
if isinstance(content, str):
return content
if isinstance(content, list):
text_parts = [
block.get("text")
for block in content
if isinstance(block, dict) and isinstance(block.get("text"), str)
]
return " ".join(part for part in text_parts if part)
return ""

View file

@ -480,21 +480,32 @@ class XecGuardGuardrail(CustomGuardrail):
choices = response.get("choices")
if not choices:
return None
first = choices[0]
if hasattr(first, "message"):
message = first.message
elif isinstance(first, dict):
message = first.get("message")
text_parts: List[str] = []
for choice in choices:
content = XecGuardGuardrail._extract_choice_content(choice)
text = XecGuardGuardrail._content_to_text(content)
if text:
text_parts.append(text)
return "\n".join(text_parts) or None
@staticmethod
def _extract_choice_content(choice: Any) -> Any:
if hasattr(choice, "message"):
message = choice.message
elif isinstance(choice, dict):
message = choice.get("message")
else:
return None
if message is None:
return None
if hasattr(message, "content"):
content = message.content
elif isinstance(message, dict):
content = message.get("content")
else:
return None
return message.content
if isinstance(message, dict):
return message.get("content")
return None
@staticmethod
def _content_to_text(content: Any) -> Optional[str]:
if isinstance(content, str) and content:
return content
if isinstance(content, list):

View file

@ -1447,14 +1447,11 @@ def callback_name(callback):
return str(callback)
@router.get(
"/health/readiness",
tags=["health"],
dependencies=[Depends(user_api_key_auth)],
)
async def health_readiness(response: Response):
async def _get_health_readiness_details(
response: Optional[Response] = None,
) -> Dict[str, Any]:
"""
Unprotected endpoint for checking if worker can receive requests
Detailed health payload for authenticated diagnostics.
"""
from litellm.proxy.proxy_server import prisma_client, version
@ -1473,7 +1470,7 @@ async def health_readiness(response: Response):
success_callback_names = litellm.success_callback
# check Cache
cache_type = None
cache_type: Any = None
if litellm.cache is not None:
from litellm.caching.caching import RedisSemanticCache
@ -1482,6 +1479,7 @@ async def health_readiness(response: Response):
if isinstance(litellm.cache.cache, RedisSemanticCache):
# ping the cache
# TODO: @ishaan-jaff - we should probably not ping the cache on every /health/readiness check
index_info: Any
try:
index_info = await litellm.cache.cache._index_info()
except Exception as e:
@ -1499,7 +1497,7 @@ async def health_readiness(response: Response):
# serve requests that depend on persisted state (keys, budgets,
# spend logs). Return 503 so orchestrators take this pod out of
# rotation; "Not connected" (no DB configured at all) stays 200.
if db_health_status["status"] != "connected":
if response is not None and db_health_status["status"] != "connected":
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
return {
"status": "healthy",
@ -1526,6 +1524,52 @@ async def health_readiness(response: Response):
raise HTTPException(status_code=503, detail=f"Service Unhealthy ({str(e)})")
def _allow_public_health_readiness_details() -> bool:
from litellm.proxy.proxy_server import general_settings
return general_settings.get("allow_public_health_readiness_details") is True
async def _set_public_readiness_status(response: Response) -> None:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return
db_health_status = await _db_health_readiness_check()
if db_health_status["status"] != "connected":
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
@router.get(
"/health/readiness",
tags=["health"],
)
async def health_readiness(response: Response):
"""
Public readiness probe. Keep this low-detail for unauthenticated load
balancers by default. Admins can opt into the legacy detailed public
payload with general_settings.allow_public_health_readiness_details.
"""
if _allow_public_health_readiness_details():
return await _get_health_readiness_details(response=response)
await _set_public_readiness_status(response=response)
return {"status": "healthy"}
@router.get(
"/health/readiness/details",
tags=["health"],
dependencies=[Depends(user_api_key_auth)],
)
async def health_readiness_details(response: Response):
"""
Authenticated readiness diagnostics with DB/cache/callback metadata.
"""
return await _get_health_readiness_details(response=response)
@router.get(
"/health/backlog",
tags=["health"],
@ -1561,7 +1605,6 @@ async def health_liveliness():
@router.options(
"/health/readiness",
tags=["health"],
dependencies=[Depends(user_api_key_auth)],
)
async def health_readiness_options():
"""

View file

@ -8,6 +8,10 @@ from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import (
is_text_content_call_type,
iter_message_text,
)
class _PROXY_AzureContentSafety(
@ -118,10 +122,9 @@ class _PROXY_AzureContentSafety(
):
verbose_proxy_logger.debug("Inside Azure Content-Safety Pre-Call Hook")
try:
if call_type == "completion" and "messages" in data:
for m in data["messages"]:
if "content" in m and isinstance(m["content"], str):
await self.test_violation(content=m["content"], source="input")
if is_text_content_call_type(call_type):
for text in iter_message_text(data):
await self.test_violation(content=text, source="input")
except HTTPException as e:
raise e
@ -140,12 +143,16 @@ class _PROXY_AzureContentSafety(
response,
):
verbose_proxy_logger.debug("Inside Azure Content-Safety Post-Call Hook")
if isinstance(response, litellm.ModelResponse) and isinstance(
response.choices[0], litellm.utils.Choices
):
await self.test_violation(
content=response.choices[0].message.content or "", source="output"
)
if not isinstance(response, litellm.ModelResponse):
return
for choice in response.choices:
if not isinstance(choice, litellm.utils.Choices):
continue
message = getattr(choice, "message", None)
content = getattr(message, "content", None)
if isinstance(content, str):
await self.test_violation(content=content, source="output")
# async def async_post_call_streaming_hook(
# self,

View file

@ -5,7 +5,7 @@ import time
from collections import OrderedDict
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from fastapi import Request
from fastapi import HTTPException, Request
from pydantic import ValidationError as PydanticValidationError
from starlette.datastructures import Headers
@ -14,6 +14,7 @@ from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm._service_logger import ServiceLogging
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host
from litellm.proxy._types import (
AddTeamCallback,
CommonProxyErrors,
@ -60,6 +61,7 @@ from litellm.secret_managers.main import get_secret_bool
from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS
from litellm.types.services import ServiceTypes
from litellm.types.utils import (
CustomPricingLiteLLMParams,
LlmProviders,
ProviderSpecificHeader,
StandardLoggingUserAPIKeyMetadata,
@ -120,6 +122,19 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS = (
"pillar_response_headers",
"_guardrail_pipelines",
"_pipeline_managed_guardrails",
# Callback-registration fields. ``callbacks``, ``service_callback``,
# and ``logger_fn`` are read by ``litellm.utils.function_setup`` and
# appended to process-wide ``litellm.{input,success,failure,_async_*,
# service}_callback`` lists / ``litellm.user_logger_fn`` — one request
# poisons the worker for every subsequent caller.
# ``litellm_disabled_callbacks`` is the inverse primitive: the
# legitimate path reads it from key/team metadata, the request-body
# version silently turns off admin-configured audit/observability
# for the caller's request.
"callbacks",
"service_callback",
"logger_fn",
"litellm_disabled_callbacks",
)
_UNTRUSTED_METADATA_CONTROL_FIELDS = (
@ -154,6 +169,59 @@ _ALLOW_CLIENT_MESSAGE_REDACTION_OPT_OUT_METADATA_KEY = (
"allow_client_message_redaction_opt_out"
)
# Per-request pricing parameters mutate cost-tracking output and (via
# ``litellm.completion`` → ``register_model``) the process-wide
# ``litellm.model_cost`` map. Both effects belong to deployment configuration,
# not to user-supplied request bodies, so the proxy strips them before they
# reach the call path. Built from the Pydantic model so newly-added pricing
# fields are covered automatically.
_CLIENT_PRICING_CONTROL_FIELDS = 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 = frozenset({"model_info"})
_ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY = "allow_client_pricing_override"
# Request fields whose value, when URL-valued, becomes the outbound destination
# for a provider call. Letting a proxy caller pin the destination is an SSRF
# primitive (HuggingFace/Oobabooga `model`, Gemini files `file_id`); guard
# them centrally so SDK users keep working but proxy users default-deny.
_URL_DESTINATION_REQUEST_FIELDS = ("model", "file_id")
def _reject_url_valued_destinations(data: Dict[str, Any]) -> None:
"""Reject URL-valued ``model``/``file_id`` unless admin-allowlisted.
Some providers (HuggingFace, Oobabooga, Gemini files) accept a URL in the
identifier field and use it as the outbound destination. On the proxy that
is an SSRF primitive — a low-privilege caller can point traffic at any
host the proxy can reach, including internal services. Reject here at the
proxy boundary so SDK users (who legitimately pass URL-valued identifiers)
are unaffected, while admins can opt specific hosts back in via
``litellm.provider_url_destination_allowed_hosts``.
"""
allowed_hosts = getattr(litellm, "provider_url_destination_allowed_hosts", []) or []
for field in _URL_DESTINATION_REQUEST_FIELDS:
value = data.get(field)
if not isinstance(value, str) or not value.startswith(("http://", "https://")):
continue
if is_url_destination_allowed_by_host(value, allowed_hosts):
continue
raise HTTPException(
status_code=400,
detail={
"error": "invalid_request",
"param": field,
"message": (
f"URL-valued '{field}' is not allowed. Configure custom "
"endpoints with api_base instead, or add the destination "
"host to `provider_url_destination_allowed_hosts` in "
"litellm_settings."
),
},
)
def _strip_untrusted_request_header_controls(
headers: Any,
@ -212,6 +280,46 @@ def _key_or_team_allows_client_message_redaction_opt_out(
)
def _key_or_team_allows_client_pricing_override(
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
return _key_or_team_metadata_flag_is_true(
user_api_key_dict=user_api_key_dict,
metadata_key=_ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY,
)
def _strip_client_pricing_overrides(data: Dict[str, Any]) -> None:
"""Drop pricing overrides from the request body and any metadata variant.
Skipped only when the calling key/team carries
``allow_client_pricing_override: True`` in its metadata. Emits a
``debug``-level log line naming the dropped fields so operators can
trace why a client-supplied pricing override stopped being applied
(otherwise the strip is invisible from the caller's perspective).
"""
stripped: List[str] = []
for field in _CLIENT_PRICING_CONTROL_FIELDS:
if field in data:
stripped.append(field)
data.pop(field, None)
for metadata_key in ("metadata", "litellm_metadata"):
metadata = data.get(metadata_key)
if not isinstance(metadata, dict):
continue
for field in _CLIENT_PRICING_METADATA_FIELDS:
if field in metadata:
stripped.append(f"{metadata_key}.{field}")
metadata.pop(field, None)
if stripped:
verbose_proxy_logger.debug(
"Stripped client-supplied pricing fields from request body: %s. "
"Set `allow_client_pricing_override: true` on the key or team "
"metadata to keep these values.",
", ".join(stripped),
)
def _get_metadata_variable_name(request: Request) -> str:
"""
Helper to return what the "metadata" field should be called in the request data
@ -1109,6 +1217,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915
if _allow_client_mock_response and _internal_key in _CLIENT_MOCK_CONTROL_FIELDS:
continue
data.pop(_internal_key, None)
_reject_url_valued_destinations(data)
# Strip spoofable auth metadata from user-supplied metadata dict
_user_metadata = data.get("metadata")
if isinstance(_user_metadata, dict):
@ -1310,6 +1419,14 @@ async def add_litellm_data_to_request( # noqa: PLR0915
]:
_user_meta.pop(_k, None)
# Strip pricing overrides AFTER the litellm_metadata string-to-dict parse
# above, for the same reason as the user_api_key_* strip — JSON-string
# metadata (sent via multipart/form-data or extra_body) wouldn't be a
# dict yet at the earlier strip point and the isinstance(dict) guard
# would silently skip the field.
if not _key_or_team_allows_client_pricing_override(user_api_key_dict):
_strip_client_pricing_overrides(data)
# Strip caller-supplied routing/budget tags unless the admin has opted
# this key or team in via metadata.allow_client_tags=True. Tags drive
# tag-based routing and tag budget attribution — accepting them from

View file

@ -8,14 +8,23 @@ POST /cache/settings/test - Test cache connection with provided credentials
POST /cache/settings - Save cache settings to database
"""
import asyncio
import json
from typing import Any, Dict, List, Optional
from datetime import datetime, timezone
from typing import Any, Dict, List, Mapping, Optional
from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, Depends, Header, HTTPException
from pydantic import BaseModel, Field
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm._uuid import uuid
from litellm.proxy._types import (
AUDIT_ACTIONS,
LiteLLM_AuditLogs,
LitellmTableNames,
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.management_endpoints import (
CACHE_SETTINGS_FIELDS,
@ -26,6 +35,85 @@ from litellm.types.management_endpoints import (
router = APIRouter()
_REDACTED_VALUE = "***REDACTED***"
def _redact_settings(settings: Optional[Mapping[str, Any]]) -> Dict[str, Any]:
"""Replace every value in a settings map with a fixed marker.
Cache config carries Redis credentials (passwords, connection strings).
The audit-log row preserves the field names so a reader can see *which*
fields changed, but values are stripped so the audit table can't itself
become a credential-harvest sink.
"""
if not settings:
return {}
return {k: _REDACTED_VALUE for k in settings.keys()}
def _log_audit_task_exception(task: "asyncio.Task[None]") -> None:
"""Surface a fire-and-forget audit-log task failure as a warning.
``asyncio.create_task`` swallows exceptions silently — if the audit
write fails we'd otherwise lose the row without any signal.
"""
if task.cancelled():
return
exc = task.exception()
if exc is not None:
verbose_proxy_logger.warning(
"Failed to write cache-settings audit log: %s", exc
)
async def _emit_cache_settings_audit_log(
*,
action: AUDIT_ACTIONS,
before_settings: Optional[Mapping[str, Any]],
after_settings: Optional[Mapping[str, Any]],
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: Optional[str],
) -> None:
"""Emit an audit-log row for a /cache/settings mutation.
Mirrors the ``store_audit_logs``-gated pattern used in
``team_callback_endpoints.py``: fire-and-forget, no-op when audit
logging is disabled, with a done-callback that surfaces any task
exception. Captured under ``LiteLLM_CacheConfig`` so the row
co-locates with the table it mutates.
"""
if litellm.store_audit_logs is not True:
return
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
task = asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.CACHE_CONFIG_TABLE_NAME,
object_id="cache_config",
action=action,
updated_values=json.dumps(
{"settings": _redact_settings(after_settings)}, default=str
),
before_value=json.dumps(
{"settings": _redact_settings(before_settings)}, default=str
),
)
)
)
task.add_done_callback(_log_audit_task_exception)
class CacheSettingsManager:
"""
Manages cache settings initialization and updates.
@ -282,6 +370,10 @@ async def test_cache_connection(
async def update_cache_settings(
request: CacheSettingsUpdateRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
litellm_changed_by: Optional[str] = Header(
None,
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",
),
):
"""
Save cache settings to database and initialize cache.
@ -314,6 +406,19 @@ async def update_cache_settings(
try:
cache_settings = request.cache_settings.copy()
# Snapshot the prior settings (key set only — values get redacted in
# the audit row) so the audit-log entry shows which fields changed.
existing_row = await prisma_client.db.litellm_cacheconfig.find_unique(
where={"id": "cache_config"}
)
before_settings: Optional[Dict[str, Any]] = None
if existing_row is not None and existing_row.cache_settings:
try:
before_settings = json.loads(existing_row.cache_settings)
except (TypeError, ValueError):
before_settings = None
action: AUDIT_ACTIONS = "updated" if existing_row is not None else "created"
# Encrypt sensitive fields (keep redis_type for storage)
encrypted_settings = proxy_config._encrypt_env_variables(
environment_variables=cache_settings
@ -353,6 +458,18 @@ async def update_cache_settings(
# Switch on LLM response caching
proxy_config.switch_on_llm_response_caching()
# Cache settings carry Redis credentials and connection strings that
# control where LLM responses are cached. An admin (or compromised
# admin) flipping the cache backend silently is a data-routing
# pivot; emit an audit-log row so the action is traceable.
await _emit_cache_settings_audit_log(
action=action,
before_settings=before_settings,
after_settings=cache_settings,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
return {
"message": "Cache settings updated successfully",
"status": "success",

View file

@ -1,10 +1,13 @@
import asyncio
import json
import os
from typing import Any, Dict, Set
from datetime import datetime, timezone
from typing import Any, Dict, Mapping, Optional, Set
from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, Depends, Header, HTTPException
from pydantic import TypeAdapter
from litellm._uuid import uuid
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
@ -18,8 +21,11 @@ from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._types import (
AUDIT_ACTIONS,
CommonProxyErrors,
KeyManagementSystem,
LiteLLM_AuditLogs,
LitellmTableNames,
LitellmUserRoles,
UserAPIKeyAuth,
)
@ -32,6 +38,82 @@ from litellm.types.proxy.management_endpoints.config_overrides import (
router = APIRouter()
_AUDIT_REDACTED = "***REDACTED***"
def _redact_config(config: Optional[Mapping[str, Any]]) -> Dict[str, Any]:
"""Strip values from a config snapshot before audit-log emission.
Hashicorp Vault config carries ``vault_token``, ``approle_secret_id``,
``client_key`` etc. Persisting them verbatim into ``LiteLLM_AuditLogs``
would let anyone with read access to the audit table harvest the
proxy's KMS credentials. Keep keys, redact values.
"""
if not config:
return {}
return {k: _AUDIT_REDACTED for k in config.keys()}
def _log_audit_task_exception(task: "asyncio.Task[None]") -> None:
if task.cancelled():
return
exc = task.exception()
if exc is not None:
verbose_proxy_logger.warning(
"Failed to write hashicorp-vault config audit log: %s", exc
)
async def _emit_hashicorp_vault_audit_log(
*,
action: AUDIT_ACTIONS,
before_config: Optional[Mapping[str, Any]],
after_config: Optional[Mapping[str, Any]],
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: Optional[str],
) -> None:
"""Emit an audit-log row for a /config_overrides/hashicorp_vault mutation.
Mirrors the ``store_audit_logs``-gated pattern from
``team_callback_endpoints.py``. Captured under
``LiteLLM_ConfigOverrides`` so the row co-locates with the table it
mutates.
"""
import litellm
if litellm.store_audit_logs is not True:
return
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
task = asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.CONFIG_OVERRIDES_TABLE_NAME,
object_id="hashicorp_vault",
action=action,
updated_values=json.dumps(
{"config": _redact_config(after_config)}, default=str
),
before_value=json.dumps(
{"config": _redact_config(before_config)}, default=str
),
)
)
)
task.add_done_callback(_log_audit_task_exception)
# --- Hashicorp Vault constants ---
HASHICORP_ENV_VAR_MAPPING: Dict[str, str] = {
@ -144,6 +226,10 @@ def _clear_hashicorp_vault_state(proxy_config: Any) -> None:
async def update_hashicorp_vault_config(
config: HashicorpVaultConfig,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
litellm_changed_by: Optional[str] = Header(
None,
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",
),
):
"""
Update Hashicorp Vault secret manager configuration.
@ -171,6 +257,8 @@ async def update_hashicorp_vault_config(
existing_record = await prisma_client.db.litellm_configoverrides.find_unique(
where={"config_type": "hashicorp_vault"}
)
existing_decrypted: Optional[Dict[str, Any]] = None
env_values: Dict[str, Any] = {}
if existing_record is not None and existing_record.config_value is not None:
existing_data = _parse_config_value(existing_record.config_value)
existing_decrypted = proxy_config._decrypt_db_variables(existing_data)
@ -178,7 +266,8 @@ async def update_hashicorp_vault_config(
if field not in config_data and existing_decrypted.get(field):
config_data[field] = existing_decrypted[field]
else:
# No DB record yet — merge from current env vars
# No DB record (or DB record with null config_value) — merge from
# current env vars instead.
env_values = _get_current_env_values(HASHICORP_ENV_VAR_MAPPING)
for field in HASHICORP_ENV_VAR_MAPPING:
if field not in config_data and env_values.get(field):
@ -248,6 +337,22 @@ async def update_hashicorp_vault_config(
# Update change-detection cache so the background reload doesn't redundantly re-init
proxy_config._last_hashicorp_vault_config = safe_json_loads(config_value)
# Mutating the proxy's KMS config affects every secret retrieval going
# forward — emit an audit-log row so the action is traceable even
# though the secret_manager_client itself was just swapped under us.
# Action keys off row existence (a row with NULL ``config_value`` is
# still an update). ``before_config`` falls back to env vars when the
# row was absent or its ``config_value`` was NULL.
before_config = existing_decrypted if existing_decrypted is not None else env_values
action: AUDIT_ACTIONS = "updated" if existing_record is not None else "created"
await _emit_hashicorp_vault_audit_log(
action=action,
before_config=before_config,
after_config=config_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
return {
"message": "Hashicorp Vault configuration updated successfully",
"status": "success",
@ -321,6 +426,10 @@ async def get_hashicorp_vault_config(
)
async def delete_hashicorp_vault_config(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
litellm_changed_by: Optional[str] = Header(
None,
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",
),
):
"""Delete Hashicorp Vault configuration. Idempotent."""
from litellm.proxy.proxy_server import prisma_client, proxy_config
@ -337,11 +446,27 @@ async def delete_hashicorp_vault_config(
detail=CommonProxyErrors.db_not_connected_error.value,
)
# Capture the prior config before delete so the audit-log row can
# show *what* was removed (keys only — values get redacted).
existing_record = await prisma_client.db.litellm_configoverrides.find_unique(
where={"config_type": "hashicorp_vault"}
)
before_config: Optional[Dict[str, Any]] = None
if existing_record is not None and existing_record.config_value is not None:
try:
before_config = proxy_config._decrypt_db_variables(
_parse_config_value(existing_record.config_value)
)
except Exception:
before_config = None
# Delete DB record if it exists — ignore if not found
deleted = False
try:
await prisma_client.db.litellm_configoverrides.delete(
where={"config_type": "hashicorp_vault"}
)
deleted = True
except RecordNotFoundError:
verbose_proxy_logger.debug(
"No existing Hashicorp Vault config record to delete"
@ -349,6 +474,17 @@ async def delete_hashicorp_vault_config(
_clear_hashicorp_vault_state(proxy_config)
# Only emit audit log if a row was actually removed; an idempotent
# delete on a non-existent row produces no security-relevant change.
if deleted:
await _emit_hashicorp_vault_audit_log(
action="deleted",
before_config=before_config,
after_config=None,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
return {
"message": "Hashicorp Vault configuration deleted successfully",
"status": "success",

View file

@ -477,6 +477,42 @@ if MCP_AVAILABLE:
allowed_routes = getattr(user_api_key_dict, "allowed_routes", None)
return isinstance(allowed_routes, list) and len(allowed_routes) > 0
def _sanitize_mcp_server_for_non_admin(
mcp_server: LiteLLM_MCPServerTable,
) -> LiteLLM_MCPServerTable:
"""Strip credential-bearing fields for non-admin viewers.
Non-admin users may legitimately need to discover MCP servers
their team has access to (so they can pick one in the UI), but
they must never see fields that can carry bearer tokens or
upstream API keys. ``_redact_mcp_credentials`` already clears
the explicit ``credentials`` field; this layers on top to catch
the URL+headers+env vectors that the virtual-key sanitizer also
strips. Reset values match each field's declared default on
``LiteLLM_MCPServerTable`` (``None`` for Optional fields,
``[]``/``{}`` for required list/dict fields).
"""
sanitized = _redact_mcp_credentials(mcp_server)
# URL is the highest-impact vector: many MCP integrations embed
# the upstream API key directly in the path. spec_path can carry
# similar tokens in the OpenAPI spec URL.
sanitized.url = None
sanitized.spec_path = None
sanitized.static_headers = None
sanitized.extra_headers = []
sanitized.env = {}
sanitized.command = None
sanitized.args = []
sanitized.authorization_url = None
sanitized.token_url = None
sanitized.registration_url = None
return sanitized
def _sanitize_mcp_server_list_for_non_admin(
mcp_servers: Iterable[LiteLLM_MCPServerTable],
) -> List[LiteLLM_MCPServerTable]:
return [_sanitize_mcp_server_for_non_admin(s) for s in mcp_servers]
def _sanitize_mcp_server_for_virtual_key(
mcp_server: LiteLLM_MCPServerTable,
) -> LiteLLM_MCPServerTable:
@ -926,6 +962,12 @@ if MCP_AVAILABLE:
if is_restricted_virtual_key:
return _sanitize_mcp_server_list_for_virtual_key(redacted_mcp_servers)
# Non-admin authenticated users may see the server inventory but
# not credential-bearing fields like `url` (often contains bearer
# tokens) or headers/env (often contain Authorization).
if not _user_has_admin_view(user_api_key_dict):
return _sanitize_mcp_server_list_for_non_admin(redacted_mcp_servers)
return redacted_mcp_servers
@router.get(
@ -1293,6 +1335,8 @@ if MCP_AVAILABLE:
redacted = _redact_mcp_credentials(mcp_server)
if is_restricted_virtual_key:
return _sanitize_mcp_server_for_virtual_key(redacted)
if not _user_has_admin_view(user_api_key_dict):
return _sanitize_mcp_server_for_non_admin(redacted)
return redacted
@router.post(

View file

@ -500,6 +500,21 @@ async def update_organization(
if data.updated_by is None:
data.updated_by = user_api_key_dict.user_id
if data.organization_id is None:
raise HTTPException(
status_code=400,
detail={"error": "organization_id is required"},
)
# IDOR guard: only proxy admins / org admins of THIS org may update
# it. Without this, any authenticated key holder could rewrite
# another organization's metadata, budgets, and object permissions.
await _verify_org_access(
organization_id=data.organization_id,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
existing_organization_row = (
await prisma_client.db.litellm_organizationtable.find_unique(
where={"organization_id": data.organization_id},
@ -909,6 +924,16 @@ async def organization_member_add(
if prisma_client is None:
raise HTTPException(status_code=500, detail={"error": "No db connected"})
# IDOR guard: docstring says "Only proxy_admin or org_admin of
# organization, allowed to access this endpoint" — but the code
# never enforced that. Any authenticated key holder could add
# members to any org. Now gated explicitly.
await _verify_org_access(
organization_id=data.organization_id,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
# Check if organization exists
existing_organization_row = (
await prisma_client.db.litellm_organizationtable.find_unique(
@ -1018,6 +1043,16 @@ async def organization_member_update(
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
# IDOR guard: only proxy admins / org admins of THIS org may
# update member roles. The PROXY_ADMIN-target check below was
# the only access control; without this, any authenticated user
# could change any non-admin member's role in any org.
await _verify_org_access(
organization_id=data.organization_id,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
# Check if organization exists
existing_organization_row = (
await prisma_client.db.litellm_organizationtable.find_unique(
@ -1179,6 +1214,15 @@ async def organization_member_delete(
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
# IDOR guard: only proxy admins / org admins of THIS org may
# delete members. Without this, any authenticated key holder
# could remove any user from any org.
await _verify_org_access(
organization_id=data.organization_id,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
if data.user_email is not None and data.user_id is None:
existing_user_email_row = await find_member_if_email(
data.user_email, prisma_client

View file

@ -104,11 +104,16 @@ async def get_router_settings(
config = await proxy_config.get_config()
router_settings_from_config = config.get("router_settings", {})
# Get current values from llm_router if initialized
current_values = {}
current_values: Dict[str, Any] = {}
if llm_router is not None:
# Check all field names from the fields list
# Router exposes routing groups as private `_routing_groups`; the
# generic `hasattr` loop below would miss them.
current_values["routing_groups"] = [
group.model_dump() for group in llm_router._routing_groups.values()
]
for field in router_fields:
if field.field_name == "routing_groups":
continue
if hasattr(llm_router, field.field_name):
value = getattr(llm_router, field.field_name)
current_values[field.field_name] = value

View file

@ -19,6 +19,7 @@ from litellm._uuid import uuid
from litellm.proxy._types import (
AddTeamCallback,
LiteLLM_AuditLogs,
LiteLLM_TeamTable,
LitellmTableNames,
ProxyErrorTypes,
ProxyException,
@ -26,6 +27,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.team_endpoints import _verify_team_access
from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
router = APIRouter()
@ -207,6 +209,15 @@ async def add_team_callbacks(
},
)
# IDOR guard: only proxy admins / org admins / team admins of THIS
# team may write callback credentials. Without this, any
# authenticated key holder could overwrite another team's logging
# config (and read back the credentials they wrote).
await _verify_team_access(
team_obj=LiteLLM_TeamTable(**_existing_team.model_dump()),
user_api_key_dict=user_api_key_dict,
)
# store team callback settings in metadata
team_metadata = _existing_team.metadata
team_callback_settings: List[dict] = team_metadata.get(
@ -316,6 +327,14 @@ async def disable_team_logging(
detail={"error": f"Team id = {team_id} does not exist."},
)
# IDOR guard: only proxy admins / org admins / team admins of THIS
# team may disable its logging — otherwise any authenticated key
# holder can silence audit logging for any team.
await _verify_team_access(
team_obj=LiteLLM_TeamTable(**_existing_team.model_dump()),
user_api_key_dict=user_api_key_dict,
)
# Update team metadata to disable logging
team_metadata = _existing_team.metadata
before_metadata = copy.deepcopy(team_metadata)
@ -364,20 +383,18 @@ async def disable_team_logging(
},
}
except HTTPException:
# Legitimate 4xx (e.g. 403 from the access guard, 404 for an
# unknown team). Re-raise without the error-level log noise that
# the catch-all branch below would produce.
raise
except ProxyException:
raise
except Exception as e:
verbose_proxy_logger.error(
f"litellm.proxy.proxy_server.disable_team_logging(): Exception occurred - {str(e)}"
)
verbose_proxy_logger.debug(traceback.format_exc())
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "detail", f"Internal Server Error({str(e)})"),
type=ProxyErrorTypes.internal_server_error.value,
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR),
)
elif isinstance(e, ProxyException):
raise e
raise ProxyException(
message="Internal Server Error, " + str(e),
type=ProxyErrorTypes.internal_server_error.value,
@ -437,6 +454,14 @@ async def get_team_callbacks(
detail={"error": f"Team id = {team_id} does not exist."},
)
# IDOR guard: callback metadata holds third-party API credentials
# (Langfuse / Langsmith / GCS). Only proxy admins / org admins /
# team admins of THIS team may read them.
await _verify_team_access(
team_obj=LiteLLM_TeamTable(**_existing_team.model_dump()),
user_api_key_dict=user_api_key_dict,
)
# Retrieve team callback settings from metadata
team_metadata = _existing_team.metadata
team_callback_settings = team_metadata.get("callback_settings", {})
@ -454,6 +479,13 @@ async def get_team_callbacks(
},
}
except HTTPException:
# Legitimate 4xx (e.g. 403 from the access guard) — re-raise
# without the error-level log noise that the catch-all below
# would produce.
raise
except ProxyException:
raise
except Exception as e:
verbose_proxy_logger.error(
"litellm.proxy.proxy_server.get_team_callbacks(): Exception occurred - {}".format(

View file

@ -5028,24 +5028,30 @@ async def get_team_daily_activity(
}
# Check if user is team admin or has /team/daily/activity permission
# If not, filter by user's API keys
# If not, filter by user's API keys.
#
# Earlier this loop used `any-team admin -> set has_full_team_view=True
# for the entire request`, so an admin of one team that requested
# data for several teams would see API-key-level breakdowns for all
# of them. Require full view on EVERY requested team — if the caller
# only has admin/permission for a strict subset, fall back to
# filtering the entire response by their own API keys (they can re-
# request the admin-only teams separately to get the wider view).
user_api_keys: Optional[List[str]] = None
if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases:
# Check if user is team admin or has usage view permission for any team
has_full_team_view = False
has_full_team_view = True
for team_alias in team_aliases:
team_obj = LiteLLM_TeamTable(**team_alias.model_dump())
if _is_user_team_admin(
is_admin = _is_user_team_admin(
user_api_key_dict=user_api_key_dict, team_obj=team_obj
):
has_full_team_view = True
break
if _team_member_has_permission(
)
has_perm = _team_member_has_permission(
user_api_key_dict=user_api_key_dict,
team_obj=team_obj,
permission="/team/daily/activity",
):
has_full_team_view = True
)
if not (is_admin or has_perm):
has_full_team_view = False
break
# If user does not have full team view, filter by their API keys

View file

@ -20,13 +20,13 @@ class PrometheusAuthMiddleware:
"""
Middleware to authenticate requests to the metrics endpoint.
By default, auth is not run on the metrics endpoint.
By default, auth is run on the metrics endpoint.
Enabled by setting the following in proxy_config.yaml:
To allow unauthenticated metrics in proxy_config.yaml:
```yaml
litellm_settings:
require_auth_for_metrics_endpoint: true
require_auth_for_metrics_endpoint: false
```
"""
@ -39,8 +39,8 @@ class PrometheusAuthMiddleware:
await self.app(scope, receive, send)
return
# Only run auth if configured to do so
if litellm.require_auth_for_metrics_endpoint is True:
# Run auth by default; allow legacy public metrics only when explicitly disabled.
if litellm.require_auth_for_metrics_endpoint is not False:
# user_api_key_auth reads the request body, which consumes ASGI `receive`.
# Buffer those messages and replay them for the inner app; otherwise a
# successful auth would forward an exhausted receive and /metrics hangs.
@ -52,10 +52,29 @@ class PrometheusAuthMiddleware:
return message
request = Request(scope, receive_for_auth)
api_key = request.headers.get(_AUTHORIZATION_HEADER) or ""
try:
await user_api_key_auth(request=request, api_key=api_key)
await user_api_key_auth(
request=request,
api_key=request.headers.get(_AUTHORIZATION_HEADER) or "",
azure_api_key_header=request.headers.get(
SpecialHeaders.azure_authorization.value
)
or "",
anthropic_api_key_header=request.headers.get(
SpecialHeaders.anthropic_authorization.value
),
google_ai_studio_api_key_header=request.headers.get(
SpecialHeaders.google_ai_studio_authorization.value
),
azure_apim_header=request.headers.get(
SpecialHeaders.azure_apim_authorization.value
)
or "",
custom_litellm_key_header=request.headers.get(
SpecialHeaders.custom_litellm_api_key.value
),
)
except Exception as e:
# Send 401 response directly via ASGI protocol
error_message = getattr(e, "message", str(e))

View file

@ -2,6 +2,7 @@ import base64
import mimetypes
import re
from dataclasses import dataclass, field
from types import MappingProxyType
from typing import TYPE_CHECKING, List, Literal, Optional, Union
from litellm.types.utils import SpecialEnums
@ -298,6 +299,7 @@ def prepare_data_with_credentials(
data: dict,
credentials: dict,
file_id: Optional[str] = None,
include_internal_credentials: bool = False,
) -> None:
"""
Update data dictionary with model credentials (in-place).
@ -306,8 +308,14 @@ def prepare_data_with_credentials(
data: Data dictionary to update
credentials: Credentials from router
file_id: Optional original file_id to set (for decoded file IDs)
include_internal_credentials: Preserve an immutable server-side snapshot
for code paths that must distinguish proxy config from request params.
"""
data.update(credentials)
if include_internal_credentials:
data["_litellm_internal_model_credentials"] = MappingProxyType(
dict(credentials)
)
data.pop("custom_llm_provider", None)
if file_id is not None:

View file

@ -774,6 +774,7 @@ async def get_file_content( # noqa: PLR0915
data=data,
credentials=credentials, # type: ignore
file_id=original_file_id, # Use decoded file ID if from encoded ID
include_internal_credentials=True,
)
response = await litellm.afile_content(
custom_llm_provider=credentials["custom_llm_provider"], # type: ignore
@ -949,6 +950,7 @@ async def get_file(
data=data,
credentials=credentials, # type: ignore
file_id=original_file_id,
include_internal_credentials=True,
)
response = await litellm.afile_retrieve(**data) # type: ignore
@ -1149,6 +1151,7 @@ async def delete_file(
data=data,
credentials=credentials, # type: ignore
file_id=original_file_id,
include_internal_credentials=True,
)
response = await litellm.afile_delete(

View file

@ -5,7 +5,7 @@ from importlib.resources import files
from typing import Any, Dict, List, Optional
import litellm
from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, HTTPException
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_blog_posts import (
@ -14,8 +14,9 @@ from litellm.litellm_core_utils.get_blog_posts import (
GetBlogPosts,
get_blog_posts,
)
from litellm.proxy._types import CommonProxyErrors
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy._types import (
CommonProxyErrors,
)
from litellm.types.agents import AgentCard
from litellm.types.mcp import MCPPublicServer
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
@ -31,6 +32,7 @@ from litellm.types.utils import LlmProviders
router = APIRouter()
# ---------------------------------------------------------------------------
# /public/endpoints — helpers
# ---------------------------------------------------------------------------
@ -153,7 +155,6 @@ def _load_endpoints() -> List[Dict[str, Any]]:
@router.get(
"/public/model_hub",
tags=["public", "model management"],
dependencies=[Depends(user_api_key_auth)],
response_model=List[ModelGroupInfoProxy],
)
async def public_model_hub():
@ -208,7 +209,6 @@ async def public_model_hub():
@router.get(
"/public/agent_hub",
tags=["[beta] Agents", "public"],
dependencies=[Depends(user_api_key_auth)],
response_model=List[AgentCard],
)
async def get_agents():
@ -230,7 +230,6 @@ async def get_agents():
@router.get(
"/public/mcp_hub",
tags=["[beta] MCP", "public"],
dependencies=[Depends(user_api_key_auth)],
response_model=List[MCPPublicServer],
)
async def get_mcp_servers():

View file

@ -3079,7 +3079,11 @@ async def global_spend_models(
return response
@router.get("/provider/budgets", response_model=ProviderBudgetResponse)
@router.get(
"/provider/budgets",
dependencies=[Depends(user_api_key_auth)],
response_model=ProviderBudgetResponse,
)
async def provider_budgets() -> ProviderBudgetResponse:
"""
Provider Budget Routing - Get Budget, Spend Details https://docs.litellm.ai/docs/proxy/provider_budget_routing

View file

@ -99,6 +99,11 @@ class UISettings(BaseModel):
description="If true, requires authentication for accessing the public AI Hub.",
)
allow_public_health_readiness_details: bool = Field(
default=False,
description="If true, returns the legacy detailed payload from the unauthenticated /health/readiness endpoint.",
)
forward_client_headers_to_llm_api: bool = Field(
default=False,
description=(
@ -169,6 +174,7 @@ ALLOWED_UI_SETTINGS_FIELDS = {
"disable_team_admin_delete_team_user",
"enabled_ui_pages_internal_users",
"require_auth_for_public_ai_hub",
"allow_public_health_readiness_details",
"forward_client_headers_to_llm_api",
"forward_llm_provider_auth_headers",
"disable_agents_for_internal_users",
@ -183,6 +189,7 @@ ALLOWED_UI_SETTINGS_FIELDS = {
# Flags that must be synced from the persisted UISettings into
# general_settings at runtime (on both read and write).
_RUNTIME_GENERAL_SETTINGS_FLAGS = [
"allow_public_health_readiness_details",
"forward_client_headers_to_llm_api",
"forward_llm_provider_auth_headers",
"disable_agents_for_internal_users",

View file

@ -8,6 +8,9 @@ from litellm.proxy._types import CommonProxyErrors, 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.utils import jsonify_object
from litellm.proxy.vector_store_endpoints.management_endpoints import (
_resolve_embedding_config,
)
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store,
get_litellm_managed_vector_store,
@ -56,6 +59,30 @@ async def _update_request_data_with_litellm_managed_vector_store_registry(
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
# Resolve ``litellm_embedding_config`` here, at request-handling
# time, instead of at row-creation time. The resolved
# ``api_key`` / ``api_base`` / ``api_version`` lives only in
# this per-request ``data`` dict and is never persisted.
# Legacy rows that already carry a resolved (cleartext)
# ``litellm_embedding_config`` skip the lookup and pass through
# unchanged so the embed call keeps working.
embedding_model = litellm_params.get("litellm_embedding_model")
if embedding_model and not litellm_params.get("litellm_embedding_config"):
from litellm.proxy.proxy_server import prisma_client
resolved_config = await _resolve_embedding_config(
embedding_model=embedding_model, prisma_client=prisma_client
)
if resolved_config:
# Build a fresh dict via spread instead of mutating
# ``litellm_params`` in place — the registry hands back
# a reference to its cached object, so an in-place
# update would persist the resolved cleartext into the
# in-memory cache for the lifetime of the process.
litellm_params = {
**litellm_params,
"litellm_embedding_config": resolved_config,
}
data.update(litellm_params)
return data

View file

@ -16,6 +16,7 @@ from fastapi import APIRouter, Depends, HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import REDACTED_BY_LITELM_STRING
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
@ -45,6 +46,28 @@ _LITELLM_PARAMS_MASKER = SensitiveDataMasker()
_REDACT_LITELLM_PARAMS_MAX_DEPTH = 10
# Use-time embedding-config resolution runs on every vector-store request
# whose persisted row carries only a model reference (the post-fix shape).
# Without a cache, that's one ``litellm_proxymodeltable.find_first`` per
# request — the no-DB-in-critical-path rule. Hold the resolved config in
# memory for a short TTL so a hot model name pays the DB lookup at most
# once per ``_EMBEDDING_CONFIG_CACHE_TTL`` seconds. Cleartext credentials
# only ever live in process memory (never persisted, never echoed in
# management responses), so the cache doesn't widen the disclosure surface.
_EMBEDDING_CONFIG_CACHE_TTL = 60
_EMBEDDING_CONFIG_CACHE_MAX_SIZE = 256
_embedding_config_cache: Optional[InMemoryCache] = None
def _get_embedding_config_cache() -> InMemoryCache:
global _embedding_config_cache
if _embedding_config_cache is None:
_embedding_config_cache = InMemoryCache(
max_size_in_memory=_EMBEDDING_CONFIG_CACHE_MAX_SIZE,
default_ttl=_EMBEDDING_CONFIG_CACHE_TTL,
)
return _embedding_config_cache
def _redact_sensitive_litellm_params(litellm_params: Any, _depth: int = 0) -> Any:
"""
@ -303,6 +326,11 @@ async def _resolve_embedding_config(
This function first checks the router for config-defined models, then falls back
to the database. This allows users to use models defined in either location.
Results are cached in process memory for ``_EMBEDDING_CONFIG_CACHE_TTL``
seconds so the request-handling path doesn't hit the database on every
vector-store call. Negative results (model not found) are intentionally
not cached to avoid blocking a freshly-added model behind the TTL.
Args:
embedding_model: The embedding model string (e.g., "text-embedding-ada-002" or "azure/text-embedding-3-large")
prisma_client: The Prisma client instance
@ -314,6 +342,11 @@ async def _resolve_embedding_config(
if not embedding_model:
return None
cache = _get_embedding_config_cache()
cached = cache.get_cache(embedding_model)
if cached is not None:
return cached
# Import llm_router if not provided
if llm_router is None:
try:
@ -330,6 +363,7 @@ async def _resolve_embedding_config(
verbose_proxy_logger.debug(
f"Resolved embedding config from router for model {embedding_model}"
)
cache.set_cache(embedding_model, router_config)
return router_config
# Fall back to database
@ -341,6 +375,7 @@ async def _resolve_embedding_config(
verbose_proxy_logger.debug(
f"Resolved embedding config from database for model {embedding_model}"
)
cache.set_cache(embedding_model, db_config)
return db_config
verbose_proxy_logger.debug(
@ -432,20 +467,17 @@ async def create_vector_store_in_db(
if user_id is not None:
data_to_create["user_id"] = user_id
# Handle litellm_params - always provide at least an empty dict
# Handle litellm_params - always provide at least an empty dict.
# The earlier behaviour resolved ``litellm_embedding_config`` from the
# admin-configured router/DB model and persisted the cleartext result
# (``api_key``, ``api_base``, ``api_version``) into this row. That
# exposed every env-stored embedding-model credential on the
# ``/vector_store/{new,info,update,list}`` responses. Keep the user's
# raw ``litellm_embedding_model`` reference; resolution now happens in
# ``_update_request_data_with_litellm_managed_vector_store_registry``
# at request-handling time so the cleartext config exists only in
# per-request memory and never reaches the database.
if litellm_params:
# Auto-resolve embedding config if embedding model is provided but config is not
embedding_model = litellm_params.get("litellm_embedding_model")
if embedding_model and not litellm_params.get("litellm_embedding_config"):
resolved_config = await _resolve_embedding_config(
embedding_model=embedding_model, prisma_client=prisma_client
)
if resolved_config:
litellm_params["litellm_embedding_config"] = resolved_config
verbose_proxy_logger.info(
f"Auto-resolved embedding config for model {embedding_model}"
)
litellm_params_dict = GenericLiteLLMParams(**litellm_params).model_dump(
exclude_none=True
)
@ -531,10 +563,19 @@ async def new_vector_store(
user_id=user_api_key_dict.user_id,
)
# Apply the same litellm_params redaction the list / info / update
# endpoints already use, so a caller-supplied credential or a
# cleartext value persisted by an earlier proxy version doesn't
# come back in the response.
response_vs = LiteLLM_ManagedVectorStore(**new_vector_store)
response_vs["litellm_params"] = _redact_sensitive_litellm_params(
new_vector_store.get("litellm_params")
)
return {
"status": "success",
"message": f"Vector store {vector_store.get('vector_store_id')} created successfully",
"vector_store": new_vector_store,
"vector_store": response_vs,
}
except Exception as e:
verbose_proxy_logger.exception(f"Error creating vector store: {str(e)}")
@ -865,24 +906,15 @@ async def update_vector_store(
update_data["vector_store_metadata"]
)
# Handle litellm_params if provided
# Handle litellm_params if provided. As with the create path, the
# embedding-config auto-resolve previously persisted cleartext
# credentials into the row; resolution now happens at request-
# handling time in
# ``_update_request_data_with_litellm_managed_vector_store_registry``
# so this row only ever stores the user-supplied
# ``litellm_embedding_model`` reference.
if "litellm_params" in update_data:
_input_litellm_params: dict = update_data.get("litellm_params", {}) or {}
# Auto-resolve embedding config if embedding model is provided but config is not
embedding_model = _input_litellm_params.get("litellm_embedding_model")
if embedding_model and not _input_litellm_params.get(
"litellm_embedding_config"
):
resolved_config = await _resolve_embedding_config(
embedding_model=embedding_model, prisma_client=prisma_client
)
if resolved_config:
_input_litellm_params["litellm_embedding_config"] = resolved_config
verbose_proxy_logger.info(
f"Auto-resolved embedding config for model {embedding_model}"
)
litellm_params_dict = GenericLiteLLMParams(
**_input_litellm_params
).model_dump(exclude_none=True)

View file

@ -159,6 +159,7 @@ from litellm.types.router import (
RouterModelGroupAliasItem,
RouterRateLimitError,
RouterRateLimitErrorBasic,
RoutingGroup,
RoutingStrategy,
SearchToolTypedDict,
)
@ -308,6 +309,7 @@ class Router:
] = "simple-shuffle",
optional_pre_call_checks: Optional[OptionalPreCallChecks] = None,
routing_strategy_args: dict = {}, # just for latency-based
routing_groups: Optional[List[Union[RoutingGroup, dict]]] = None,
provider_budget_config: Optional[GenericBudgetConfigType] = None,
alerting_config: Optional[AlertingConfig] = None,
router_general_settings: Optional[
@ -347,8 +349,9 @@ class Router:
retry_after (int): Minimum time to wait before retrying a failed request. Defaults to 0.
allowed_fails (Optional[int]): Number of allowed fails before adding to cooldown. Defaults to None.
cooldown_time (float): Time to cooldown a deployment after failure in seconds. Defaults to 1.
routing_strategy (Literal["simple-shuffle", "least-busy", "usage-based-routing", "latency-based-routing", "cost-based-routing"]): Routing strategy. Defaults to "simple-shuffle".
routing_strategy_args (dict): Additional args for latency-based routing. Defaults to {}.
routing_strategy (Literal["simple-shuffle", "least-busy", "usage-based-routing", "latency-based-routing", "cost-based-routing"]): Routing strategy used for the implicit "default" group (any model not claimed by an entry in `routing_groups`). Defaults to "simple-shuffle".
routing_strategy_args (dict): Additional args for the default group's routing strategy (e.g. latency window). Defaults to {}.
routing_groups (Optional[List[RoutingGroup]]): Named subsets of `model_name`s that use a per-group routing strategy and args. Each model belongs to at most one explicit group; everything else lands in the implicit "default" group driven by `routing_strategy` / `routing_strategy_args`. Defaults to None.
alerting_config (AlertingConfig): Slack alerting configuration. Defaults to None.
provider_budget_config (ProviderBudgetConfig): Provider budget configuration. Use this to set llm_provider budget limits. example $100/day to OpenAI, $100/day to Azure, etc. Defaults to None.
deployment_affinity_ttl_seconds (int): TTL for user-key -> deployment affinity mapping. Defaults to 3600.
@ -541,7 +544,10 @@ class Router:
self.stream_timeout = stream_timeout
self.retry_after = retry_after
self.routing_strategy = routing_strategy
self.routing_strategy = self._normalize_strategy(routing_strategy)
self._routing_groups_input: Optional[List[Union[RoutingGroup, dict]]] = (
routing_groups
)
## SETTING FALLBACKS ##
### validate if it's set + in correct format
@ -612,6 +618,7 @@ class Router:
routing_strategy=routing_strategy,
routing_strategy_args=routing_strategy_args,
)
self._init_routing_groups(self._routing_groups_input)
self.access_groups = None
## USAGE TRACKING ##
if isinstance(litellm._async_success_callback, list):
@ -806,84 +813,355 @@ class Router:
if self.cache.redis_cache is None:
self.cache.redis_cache = cache
# Maps a routing strategy string to the attribute on `self` that holds
# the default group's strategy selector for that strategy. (The selectors
# double as `CustomLogger` callbacks, hence the legacy `*_logger` attrs.)
_DEFAULT_SELECTOR_ATTR_BY_STRATEGY: Dict[str, str] = {
"least-busy": "leastbusy_logger",
"usage-based-routing": "lowesttpm_logger",
"usage-based-routing-v2": "lowesttpm_logger_v2",
"latency-based-routing": "lowestlatency_logger",
"cost-based-routing": "lowestcost_logger",
}
@staticmethod
def _normalize_strategy(
strategy: Union[RoutingStrategy, str, None]
) -> Optional[str]:
if strategy is None:
return None
if isinstance(strategy, RoutingStrategy):
return strategy.value
return strategy
def _validate_routing_strategy(
self, routing_strategy: Union[RoutingStrategy, str, None]
) -> None:
# See: https://github.com/BerriAI/litellm/issues/11330
valid_strategy_strings = ["simple-shuffle"] + [s.value for s in RoutingStrategy]
if routing_strategy is None:
return
is_valid_string = (
isinstance(routing_strategy, str)
and routing_strategy in valid_strategy_strings
)
is_valid_enum = isinstance(routing_strategy, RoutingStrategy)
if not is_valid_string and not is_valid_enum:
raise ValueError(
f"Invalid routing_strategy: '{routing_strategy}'. "
f"Valid options: {valid_strategy_strings}. "
f"Check 'router_settings.routing_strategy' in your config.yaml "
f"or the 'routing_strategy' parameter if using the Router SDK directly."
)
def _build_strategy_selector(
self,
strategy: Union[RoutingStrategy, str],
routing_strategy_args: dict,
register_callbacks: bool = True,
) -> Optional[Any]:
"""
Constructs a strategy selector for a given strategy.
Returns None for `simple-shuffle` (no selector needed) and unknown
strategies.
"""
selector: Optional[Any] = None
match self._normalize_strategy(strategy):
case RoutingStrategy.LEAST_BUSY.value:
selector = LeastBusyLoggingHandler(router_cache=self.cache)
if register_callbacks:
if isinstance(litellm.input_callback, list):
litellm.input_callback.append(selector) # type: ignore
else:
litellm.input_callback = [selector] # type: ignore
case RoutingStrategy.USAGE_BASED_ROUTING.value:
selector = LowestTPMLoggingHandler(
router_cache=self.cache,
routing_args=routing_strategy_args,
)
case RoutingStrategy.USAGE_BASED_ROUTING_V2.value:
selector = LowestTPMLoggingHandler_v2(
router_cache=self.cache,
routing_args=routing_strategy_args,
)
case RoutingStrategy.LATENCY_BASED.value:
selector = LowestLatencyLoggingHandler(
router_cache=self.cache,
routing_args=routing_strategy_args,
)
case RoutingStrategy.COST_BASED.value:
selector = LowestCostLoggingHandler(
router_cache=self.cache,
routing_args={},
)
if (
selector is not None
and register_callbacks
and isinstance(litellm.callbacks, list)
):
litellm.logging_callback_manager.add_litellm_callback(selector) # type: ignore
return selector
def _unregister_router_selectors(self, selectors: List[Any]) -> None:
"""
Drop router-owned strategy selectors from litellm's global callback
lists by identity. Used before re-init (`routing_strategy_init` /
`_init_routing_groups`) so repeated `update_settings` calls don't
accumulate dead selectors that keep receiving callback events.
"""
selector_ids = {id(s) for s in selectors if s is not None}
if not selector_ids:
return
if isinstance(litellm.callbacks, list):
litellm.callbacks = [
c for c in litellm.callbacks if id(c) not in selector_ids
]
if isinstance(litellm.input_callback, list):
litellm.input_callback = [
c for c in litellm.input_callback if id(c) not in selector_ids
]
def routing_strategy_init(
self, routing_strategy: Union[RoutingStrategy, str], routing_strategy_args: dict
):
verbose_router_logger.info(f"Routing strategy: {routing_strategy}")
self._validate_routing_strategy(routing_strategy)
# Validate routing_strategy value to fail fast with helpful error
# See: https://github.com/BerriAI/litellm/issues/11330
# Derive valid strategies from RoutingStrategy enum + "simple-shuffle" (default, not in enum)
valid_strategy_strings = ["simple-shuffle"] + [s.value for s in RoutingStrategy]
self._unregister_router_selectors(
[
getattr(self, attr, None)
for attr in self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.values()
]
)
if routing_strategy is not None:
is_valid_string = (
isinstance(routing_strategy, str)
and routing_strategy in valid_strategy_strings
)
is_valid_enum = isinstance(routing_strategy, RoutingStrategy)
if not is_valid_string and not is_valid_enum:
self.leastbusy_logger: Optional[LeastBusyLoggingHandler] = None
self.lowesttpm_logger: Optional[LowestTPMLoggingHandler] = None
self.lowesttpm_logger_v2: Optional[LowestTPMLoggingHandler_v2] = None
self.lowestlatency_logger: Optional[LowestLatencyLoggingHandler] = None
self.lowestcost_logger: Optional[LowestCostLoggingHandler] = None
selector = self._build_strategy_selector(
strategy=routing_strategy,
routing_strategy_args=routing_strategy_args,
)
attr = self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.get(
self._normalize_strategy(routing_strategy) or ""
)
# TODO: legacy `self.<strategy>_logger` attributes are read directly by
# `get_settings()` and external callers. Fold the default group into
# `self._group_selectors["default"]` and drop these attribute writes —
# the dual storage is an antipattern preserved only for back-compat.
if attr is not None:
setattr(self, attr, selector)
def _init_routing_groups(
self,
groups_input: Optional[List[Union[RoutingGroup, dict]]],
) -> None:
"""
Validates and indexes `routing_groups`. Each `model_name` may belong to
at most one explicit group. Constructs per-group strategy selectors so
groups with different `routing_strategy_args` track independent state.
Models not claimed by any explicit group are served by the implicit
`"default"` group, whose selectors are the `self.<strategy>_logger`
attributes set up in `routing_strategy_init`.
"""
self._unregister_router_selectors(
[
sel
for selectors in getattr(self, "_group_selectors", {}).values()
for sel in selectors.values()
]
)
self._routing_groups: Dict[str, RoutingGroup] = {}
self._model_to_group: Dict[str, str] = {}
self._group_selectors: Dict[str, Dict[str, Any]] = {}
if not groups_input:
return
known_model_names = {
m.get("model_name") for m in (self.model_list or []) if m.get("model_name")
}
seen_group_names: set = set()
for raw in groups_input:
group = raw if isinstance(raw, RoutingGroup) else RoutingGroup(**raw)
if not group.group_name:
raise ValueError("routing_groups: group_name must be non-empty.")
if group.group_name == "default":
raise ValueError(
f"Invalid routing_strategy: '{routing_strategy}'. "
f"Valid options: {valid_strategy_strings}. "
f"Check 'router_settings.routing_strategy' in your config.yaml "
f"or the 'routing_strategy' parameter if using the Router SDK directly."
"routing_groups: 'default' is reserved for the implicit fallback group."
)
if group.group_name in seen_group_names:
raise ValueError(
f"routing_groups: group names must be unique, duplicate group_name '{group.group_name}'."
)
seen_group_names.add(group.group_name)
if (
routing_strategy == RoutingStrategy.LEAST_BUSY.value
or routing_strategy == RoutingStrategy.LEAST_BUSY
):
self.leastbusy_logger = LeastBusyLoggingHandler(router_cache=self.cache)
## add callback
if isinstance(litellm.input_callback, list):
litellm.input_callback.append(self.leastbusy_logger) # type: ignore
else:
litellm.input_callback = [self.leastbusy_logger] # type: ignore
if isinstance(litellm.callbacks, list):
litellm.logging_callback_manager.add_litellm_callback(self.leastbusy_logger) # type: ignore
elif (
routing_strategy == RoutingStrategy.USAGE_BASED_ROUTING.value
or routing_strategy == RoutingStrategy.USAGE_BASED_ROUTING
):
self.lowesttpm_logger = LowestTPMLoggingHandler(
router_cache=self.cache,
routing_args=routing_strategy_args,
self._validate_routing_strategy(group.routing_strategy)
for model_name in group.models:
if model_name in self._model_to_group:
raise ValueError(
f"routing_groups: model_name '{model_name}' appears in "
f"both '{self._model_to_group[model_name]}' and "
f"'{group.group_name}'. Each model may belong to at most one group."
)
if known_model_names and model_name not in known_model_names:
verbose_router_logger.warning(
"routing_groups: model_name '%s' (group '%s') is not in model_list; "
"the group entry will only take effect once a deployment with that "
"model_name is added.",
model_name,
group.group_name,
)
self._model_to_group[model_name] = group.group_name
self._routing_groups[group.group_name] = group
strategy_value = self._normalize_strategy(group.routing_strategy) or ""
group_selector = self._build_strategy_selector(
strategy=group.routing_strategy,
routing_strategy_args=group.routing_strategy_args or {},
)
if isinstance(litellm.callbacks, list):
litellm.logging_callback_manager.add_litellm_callback(self.lowesttpm_logger) # type: ignore
elif (
routing_strategy == RoutingStrategy.USAGE_BASED_ROUTING_V2.value
or routing_strategy == RoutingStrategy.USAGE_BASED_ROUTING_V2
):
self.lowesttpm_logger_v2 = LowestTPMLoggingHandler_v2(
router_cache=self.cache,
routing_args=routing_strategy_args,
self._group_selectors[group.group_name] = (
{strategy_value: group_selector} if group_selector is not None else {}
)
if isinstance(litellm.callbacks, list):
litellm.logging_callback_manager.add_litellm_callback(self.lowesttpm_logger_v2) # type: ignore
elif (
routing_strategy == RoutingStrategy.LATENCY_BASED.value
or routing_strategy == RoutingStrategy.LATENCY_BASED
):
self.lowestlatency_logger = LowestLatencyLoggingHandler(
router_cache=self.cache,
routing_args=routing_strategy_args,
def _get_routing_context(self, model: str) -> Tuple[Optional[str], Optional[Any]]:
"""
Resolves the routing strategy and selector to use for the given model.
Every model belongs to exactly one group: an explicit entry from
`routing_groups`, or the implicit `"default"` group driven by the
router's top-level `routing_strategy` / `routing_strategy_args`.
`self.routing_strategy` may be either a string or a `RoutingStrategy`
enum member (the constructor accepts both), so it is normalized to a
string here. Downstream call sites and `_select_deployment_*` arms
compare against string literals.
"""
group_name = self._model_to_group.get(model)
if group_name is None:
strategy = self._normalize_strategy(self.routing_strategy)
attr = self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.get(strategy or "")
selector = getattr(self, attr, None) if attr is not None else None
verbose_router_logger.debug(
"routing_group=default model=%s strategy=%s", model, strategy
)
if isinstance(litellm.callbacks, list):
litellm.logging_callback_manager.add_litellm_callback(self.lowestlatency_logger) # type: ignore
elif (
routing_strategy == RoutingStrategy.COST_BASED.value
or routing_strategy == RoutingStrategy.COST_BASED
):
self.lowestcost_logger = LowestCostLoggingHandler(
router_cache=self.cache,
routing_args={},
)
if isinstance(litellm.callbacks, list):
litellm.logging_callback_manager.add_litellm_callback(self.lowestcost_logger) # type: ignore
else:
pass
return strategy, selector
group = self._routing_groups[group_name]
strategy = self._normalize_strategy(group.routing_strategy)
selector = self._group_selectors.get(group_name, {}).get(strategy or "")
verbose_router_logger.debug(
"routing_group=%s model=%s strategy=%s", group_name, model, strategy
)
return strategy, selector
async def _select_deployment_async(
self,
*,
strategy: Optional[str],
selector: Optional[Any],
model: str,
healthy_deployments: list,
messages: Optional[List[Dict[str, str]]],
input: Optional[Union[str, List]],
request_kwargs: Optional[Dict],
) -> Optional[Any]:
"""
Asks the strategy selector for a deployment. Caller handles
`simple-shuffle` separately (it does not flow through a selector).
Returns None for unknown strategies or when the selector is missing.
"""
if selector is None:
return None
match strategy:
case "least-busy":
return await selector.async_get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments,
)
case "usage-based-routing":
# `LowestTPMLoggingHandler` (v1) only exposes the sync
# `get_available_deployments`. Mirror the pre-routing-groups
# top-level fallback by calling it inline so groups using v1
# still work from async callers.
return selector.get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
)
case "usage-based-routing-v2" | "cost-based-routing":
return await selector.async_get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
)
case "latency-based-routing":
return await selector.async_get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
request_kwargs=request_kwargs,
)
case _:
return None
def _select_deployment_sync(
self,
*,
strategy: Optional[str],
selector: Optional[Any],
model: str,
healthy_deployments: list,
messages: Optional[List[Dict[str, str]]],
input: Optional[Union[str, List]],
request_kwargs: Optional[Dict],
) -> Optional[Any]:
"""
Sync sibling of `_select_deployment_async`. Caller handles
`simple-shuffle` separately.
"""
if selector is None:
return None
# `cost-based-routing` is intentionally omitted —
# `LowestCostLoggingHandler` only implements
# `async_get_available_deployments`
match strategy:
case "least-busy":
return selector.get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments,
)
case "usage-based-routing" | "usage-based-routing-v2":
return selector.get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
)
case "latency-based-routing":
return selector.get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
request_kwargs=request_kwargs,
)
case _:
return None
def initialize_assistants_endpoint(self):
## INITIALIZE PASS THROUGH ASSISTANTS ENDPOINT ##
@ -9016,8 +9294,13 @@ class Router:
if (
var == "routing_strategy_args"
and self.routing_strategy == "latency-based-routing"
and self.lowestlatency_logger is not None
):
_settings_to_return[var] = self.lowestlatency_logger.routing_args.json()
_settings_to_return["routing_groups"] = [
group.model_dump() for group in self._routing_groups.values()
]
return _settings_to_return
def update_settings(self, **kwargs):
@ -9028,6 +9311,7 @@ class Router:
_allowed_settings = [
"routing_strategy_args",
"routing_strategy",
"routing_groups",
"allowed_fails",
"cooldown_time",
"num_retries",
@ -9049,26 +9333,34 @@ class Router:
]
_existing_router_settings = self.get_settings()
rebuild_routing_groups = False
for var in kwargs:
if var in _allowed_settings:
if var in _int_settings:
_casted_value = int(kwargs[var])
setattr(self, var, _casted_value)
elif var == "routing_groups":
self._routing_groups_input = kwargs[var]
rebuild_routing_groups = True
else:
value = kwargs[var]
# only run routing strategy init if it has changed
if (
var == "routing_strategy"
and _existing_router_settings["routing_strategy"] != kwargs[var]
):
self.routing_strategy_init(
routing_strategy=kwargs[var],
routing_strategy_args=kwargs.get(
"routing_strategy_args", {}
),
)
setattr(self, var, kwargs[var])
if var == "routing_strategy":
value = self._normalize_strategy(value)
if _existing_router_settings["routing_strategy"] != value:
self.routing_strategy_init(
routing_strategy=value,
routing_strategy_args=kwargs.get(
"routing_strategy_args", {}
),
)
rebuild_routing_groups = True
setattr(self, var, value)
else:
verbose_router_logger.debug("Setting {} is not allowed".format(var))
if rebuild_routing_groups:
self._init_routing_groups(self._routing_groups_input)
verbose_router_logger.debug(f"Updated Router settings: {self.get_settings()}")
def _get_client(self, deployment, kwargs, client_type=None):
@ -9779,6 +10071,11 @@ class Router:
messages = pre_routing_hook_response.messages
#########################################################
# Resolve the strategy and logger AFTER the pre-routing hook, since
# the hook can replace `model` and routing-group lookup must key
# off the final model name.
strategy, strategy_selector = self._get_routing_context(model)
healthy_deployments = await self.async_get_healthy_deployments(
model=model,
request_kwargs=request_kwargs,
@ -9798,61 +10095,21 @@ class Router:
return healthy_deployments[0]
start_time = time.time()
if (
self.routing_strategy == "usage-based-routing-v2"
and self.lowesttpm_logger_v2 is not None
):
deployment = (
await self.lowesttpm_logger_v2.async_get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments, # type: ignore
messages=messages,
input=input,
)
)
elif (
self.routing_strategy == "cost-based-routing"
and self.lowestcost_logger is not None
):
deployment = (
await self.lowestcost_logger.async_get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments, # type: ignore
messages=messages,
input=input,
)
)
elif (
self.routing_strategy == "latency-based-routing"
and self.lowestlatency_logger is not None
):
deployment = (
await self.lowestlatency_logger.async_get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments, # type: ignore
messages=messages,
input=input,
request_kwargs=request_kwargs,
)
)
elif self.routing_strategy == "simple-shuffle":
if strategy == "simple-shuffle":
return simple_shuffle(
llm_router_instance=self,
healthy_deployments=healthy_deployments,
model=model,
)
elif (
self.routing_strategy == "least-busy"
and self.leastbusy_logger is not None
):
deployment = (
await self.leastbusy_logger.async_get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments, # type: ignore
)
)
else:
deployment = None
deployment = await self._select_deployment_async(
strategy=strategy,
selector=strategy_selector,
model=model,
healthy_deployments=healthy_deployments, # type: ignore
messages=messages,
input=input,
request_kwargs=request_kwargs,
)
if deployment is None:
exception = await async_raise_no_deployment_exception(
litellm_router_instance=self,
@ -9960,49 +10217,22 @@ class Router:
# 5. Apply load balancing strategy
start_time = time.perf_counter()
if (
self.routing_strategy == "usage-based-routing-v2"
and self.lowesttpm_logger_v2 is not None
):
deployment = (
await self.lowesttpm_logger_v2.async_get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
messages=messages,
input=input,
)
)
elif (
self.routing_strategy == "latency-based-routing"
and self.lowestlatency_logger is not None
):
deployment = (
await self.lowestlatency_logger.async_get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
messages=messages,
input=input,
request_kwargs=request_kwargs,
)
)
elif self.routing_strategy == "simple-shuffle":
strategy, strategy_selector = self._get_routing_context(model)
if strategy == "simple-shuffle":
return simple_shuffle(
llm_router_instance=self,
healthy_deployments=pass_through_deployments,
model=model,
)
elif (
self.routing_strategy == "least-busy"
and self.leastbusy_logger is not None
):
deployment = (
await self.leastbusy_logger.async_get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
)
)
else:
deployment = None
deployment = await self._select_deployment_async(
strategy=strategy,
selector=strategy_selector,
model=model,
healthy_deployments=pass_through_deployments, # type: ignore
messages=messages,
input=input,
request_kwargs=request_kwargs,
)
if deployment is None:
exception = await async_raise_no_deployment_exception(
@ -10191,11 +10421,8 @@ class Router:
cooldown_list=_cooldown_list,
)
if self.routing_strategy == "least-busy" and self.leastbusy_logger is not None:
deployment = self.leastbusy_logger.get_available_deployments(
model_group=model, healthy_deployments=healthy_deployments # type: ignore
)
elif self.routing_strategy == "simple-shuffle":
strategy, strategy_selector = self._get_routing_context(model)
if strategy == "simple-shuffle":
# if users pass rpm or tpm, we do a random weighted pick - based on rpm/tpm
############## Check 'weight' param set for weighted pick #################
return simple_shuffle(
@ -10203,37 +10430,15 @@ class Router:
healthy_deployments=healthy_deployments,
model=model,
)
elif (
self.routing_strategy == "latency-based-routing"
and self.lowestlatency_logger is not None
):
deployment = self.lowestlatency_logger.get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments, # type: ignore
request_kwargs=request_kwargs,
)
elif (
self.routing_strategy == "usage-based-routing"
and self.lowesttpm_logger is not None
):
deployment = self.lowesttpm_logger.get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments, # type: ignore
messages=messages,
input=input,
)
elif (
self.routing_strategy == "usage-based-routing-v2"
and self.lowesttpm_logger_v2 is not None
):
deployment = self.lowesttpm_logger_v2.get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments, # type: ignore
messages=messages,
input=input,
)
else:
deployment = None
deployment = self._select_deployment_sync(
strategy=strategy,
selector=strategy_selector,
model=model,
healthy_deployments=healthy_deployments, # type: ignore
messages=messages,
input=input,
request_kwargs=request_kwargs,
)
if deployment is None:
verbose_router_logger.info(
@ -10359,47 +10564,22 @@ class Router:
)
# 6. Apply load balancing strategy
if self.routing_strategy == "least-busy" and self.leastbusy_logger is not None:
deployment = self.leastbusy_logger.get_available_deployments(
model_group=model, healthy_deployments=pass_through_deployments # type: ignore
)
elif self.routing_strategy == "simple-shuffle":
strategy, strategy_selector = self._get_routing_context(model)
if strategy == "simple-shuffle":
return simple_shuffle(
llm_router_instance=self,
healthy_deployments=pass_through_deployments,
model=model,
)
elif (
self.routing_strategy == "latency-based-routing"
and self.lowestlatency_logger is not None
):
deployment = self.lowestlatency_logger.get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
request_kwargs=request_kwargs,
)
elif (
self.routing_strategy == "usage-based-routing"
and self.lowesttpm_logger is not None
):
deployment = self.lowesttpm_logger.get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
messages=messages,
input=input,
)
elif (
self.routing_strategy == "usage-based-routing-v2"
and self.lowesttpm_logger_v2 is not None
):
deployment = self.lowesttpm_logger_v2.get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
messages=messages,
input=input,
)
else:
deployment = None
deployment = self._select_deployment_sync(
strategy=strategy,
selector=strategy_selector,
model=model,
healthy_deployments=pass_through_deployments, # type: ignore
messages=messages,
input=input,
request_kwargs=request_kwargs,
)
if deployment is None:
verbose_router_logger.info(

View file

@ -393,6 +393,7 @@ class AnthropicMessagesRequestOptionalParams(TypedDict, total=False):
AnthropicOutputConfig
] # Configuration for Claude's output behavior
cache_control: Optional[Dict[str, Any]] # Automatic prompt caching
reasoning_effort: Optional[str]
class AnthropicMessagesRequest(AnthropicMessagesRequestOptionalParams, total=False):

View file

@ -1041,3 +1041,4 @@ class BedrockInvokeAnthropicMessagesRequest(TypedDict, total=False):
# `metadata` is part of the common Anthropic Messages API shape.
thinking: dict
metadata: dict
output_config: dict

View file

@ -112,6 +112,14 @@ ROUTER_SETTINGS_FIELDS: List[RouterSettingsField] = [
field_default={},
ui_field_name="Routing Strategy Args",
),
RouterSettingsField(
field_name="routing_groups",
field_type="List",
field_value=None,
field_description="Named subsets of model_names that share a routing strategy. Models not claimed by an explicit group fall through to the top-level routing_strategy.",
field_default=[],
ui_field_name="Routing Groups",
),
RouterSettingsField(
field_name="num_retries",
field_type="Integer",

View file

@ -38,6 +38,19 @@ class ModelConfig(BaseModel):
model_config = ConfigDict(protected_namespaces=())
class RoutingGroup(BaseModel):
"""
A group of models that share a routing strategy.
"""
group_name: str
models: List[str]
routing_strategy: str
routing_strategy_args: Optional[dict] = None
model_config = ConfigDict(protected_namespaces=())
class RouterConfig(BaseModel):
model_list: List[ModelConfig]
@ -65,6 +78,7 @@ class RouterConfig(BaseModel):
"usage-based-routing",
"latency-based-routing",
] = "simple-shuffle"
routing_groups: Optional[List[RoutingGroup]] = None
model_config = ConfigDict(protected_namespaces=())
@ -76,6 +90,7 @@ class UpdateRouterConfig(BaseModel):
routing_strategy_args: Optional[dict] = None
routing_strategy: Optional[str] = None
routing_groups: Optional[List[RoutingGroup]] = None
model_group_retry_policy: Optional[dict] = None
model_group_affinity_config: Optional[Dict[str, List[str]]] = None
allowed_fails: Optional[int] = None

View file

@ -977,6 +977,7 @@
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
@ -1174,6 +1175,7 @@
"supports_vision": true,
"supports_prompt_caching": false,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_tool_choice": true
},
"global.anthropic.claude-opus-4-7": {
@ -1321,6 +1323,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -1350,6 +1353,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -1379,6 +1383,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -1407,6 +1412,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -1435,6 +1441,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -1929,6 +1936,7 @@
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
@ -2052,6 +2060,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -9219,6 +9228,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -9226,6 +9236,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -9361,6 +9372,7 @@
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
@ -9388,6 +9400,7 @@
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
@ -9409,6 +9422,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -9442,6 +9456,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -9475,6 +9490,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -9491,7 +9507,6 @@
"us": 1.1,
"fast": 6.0
},
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"claude-opus-4-7-20260416": {
@ -9510,6 +9525,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -9526,7 +9542,6 @@
"us": 1.1,
"fast": 6.0
},
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"claude-sonnet-4-20250514": {
@ -10804,6 +10819,7 @@
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_reasoning": true,
"supports_minimal_reasoning_effort": true,
"supports_tool_choice": true
},
"databricks/databricks-claude-sonnet-4": {
@ -17164,7 +17180,8 @@
],
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": true
"supports_vision": true,
"supports_minimal_reasoning_effort": true
},
"github_copilot/claude-opus-4.6-fast": {
"litellm_provider": "github_copilot",
@ -17677,7 +17694,8 @@
"mode": "chat",
"output_cost_per_token": 2.5e-05,
"supports_function_calling": true,
"supports_vision": true
"supports_vision": true,
"supports_minimal_reasoning_effort": true
},
"gmi/anthropic/claude-sonnet-4.5": {
"input_cost_per_token": 3e-06,
@ -26095,6 +26113,7 @@
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 159,
@ -26113,6 +26132,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true,
@ -26134,6 +26154,7 @@
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -26199,6 +26220,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
@ -28078,7 +28100,8 @@
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false,
"supports_function_calling": true
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true
},
"perplexity/anthropic/claude-sonnet-4-5": {
"litellm_provider": "perplexity",
@ -30557,6 +30580,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -30585,6 +30609,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -30612,6 +30637,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -31192,6 +31218,7 @@
"output_cost_per_token": 2.5e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_minimal_reasoning_effort": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -32384,6 +32411,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -32410,6 +32438,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@ -32576,6 +32605,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
@ -39587,6 +39617,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_max_reasoning_effort": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,

View file

@ -289,11 +289,10 @@ async def test_increment_remaining_budget_metrics(prometheus_logger):
future_reset_time_team = datetime.now() + timedelta(hours=10)
future_reset_time_key = datetime.now() + timedelta(hours=12)
# Mock the get_team_object and get_key_object functions to return objects with budget reset times
with patch(
"litellm.proxy.auth.auth_checks.get_team_object"
) as mock_get_team, patch(
"litellm.proxy.auth.auth_checks.get_key_object"
) as mock_get_key:
with (
patch("litellm.proxy.auth.auth_checks.get_team_object") as mock_get_team,
patch("litellm.proxy.auth.auth_checks.get_key_object") as mock_get_key,
):
mock_get_team.return_value = MagicMock(budget_reset_at=future_reset_time_team)
mock_get_key.return_value = MagicMock(budget_reset_at=future_reset_time_key)
@ -732,6 +731,51 @@ async def test_async_log_failure_event(prometheus_logger):
prometheus_logger.litellm_deployment_total_requests.labels().inc.assert_called_once()
@pytest.mark.asyncio
async def test_async_log_failure_event_litellm_side_rate_limit(prometheus_logger):
"""LiteLLM-side reject (no deployment picked) routes the requested model
into `requested_model` and skips the partial-outage flag."""
standard_logging_object = create_standard_logging_payload()
standard_logging_object["model_id"] = ""
standard_logging_object["model_group"] = ""
standard_logging_object["api_base"] = ""
rate_limit_exc = Exception("LiteLLM rate limit exceeded")
rate_limit_exc.status_code = 429
kwargs = {
"model": "us/azure/openai/gpt-5-mini",
"litellm_params": {},
"start_time": datetime.now(),
"completion_start_time": datetime.now(),
"api_call_start_time": datetime.now(),
"end_time": datetime.now() + timedelta(seconds=1),
"standard_logging_object": standard_logging_object,
"exception": rate_limit_exc,
}
prometheus_logger.litellm_llm_api_failed_requests_metric = MagicMock()
prometheus_logger.litellm_deployment_failure_responses = MagicMock()
prometheus_logger.litellm_deployment_total_requests = MagicMock()
prometheus_logger.set_deployment_partial_outage = MagicMock()
await prometheus_logger.async_log_failure_event(
kwargs, MagicMock(), kwargs["start_time"], kwargs["end_time"]
)
prometheus_logger.set_deployment_partial_outage.assert_not_called()
prometheus_logger.litellm_deployment_failure_responses.labels.assert_called_once()
actual_failure_labels = (
prometheus_logger.litellm_deployment_failure_responses.labels.call_args.kwargs
)
assert actual_failure_labels["requested_model"] == "us/azure/openai/gpt-5-mini"
assert actual_failure_labels["litellm_model_name"] == ""
assert actual_failure_labels["model_id"] == ""
assert actual_failure_labels["api_base"] == ""
assert actual_failure_labels["api_provider"] == ""
assert actual_failure_labels["exception_status"] == "429"
@pytest.mark.asyncio
async def test_async_post_call_failure_hook(prometheus_logger):
"""
@ -1518,9 +1562,12 @@ async def test_initialize_remaining_budget_metrics(prometheus_logger):
"""
litellm.prometheus_initialize_budget_metrics = True
# Mock the prisma client and get_paginated_teams function
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams"
) as mock_get_teams:
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams"
) as mock_get_teams,
):
# Create mock team data with proper datetime objects for budget_reset_at
future_reset = datetime.now() + timedelta(hours=24) # Reset 24 hours from now
mock_teams = [
@ -1613,11 +1660,15 @@ async def test_initialize_remaining_budget_metrics_exception_handling(
"""
litellm.prometheus_initialize_budget_metrics = True
# Mock the prisma client and get_paginated_teams function to raise an exception
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams"
) as mock_get_teams, patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper"
) as mock_list_keys:
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams"
) as mock_get_teams,
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper"
) as mock_list_keys,
):
# Make get_paginated_teams raise an exception
mock_get_teams.side_effect = Exception("Database error")
mock_list_keys.side_effect = Exception("Key listing error")
@ -1636,9 +1687,7 @@ async def test_initialize_remaining_budget_metrics_exception_handling(
# Mock litellm_organizationtable to raise an exception for org budget metrics
mock_orgtable = MagicMock()
mock_orgtable.find_many = MagicMock(
side_effect=Exception("Org database error")
)
mock_orgtable.find_many = MagicMock(side_effect=Exception("Org database error"))
mock_orgtable.count = MagicMock(side_effect=Exception("Org count error"))
mock_db = MagicMock()
@ -1699,9 +1748,12 @@ async def test_initialize_api_key_budget_metrics(prometheus_logger):
"""
litellm.prometheus_initialize_budget_metrics = True
# Mock the prisma client and _list_key_helper function
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper"
) as mock_list_keys:
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper"
) as mock_list_keys,
):
# Create mock key data with proper datetime objects for budget_reset_at
future_reset = datetime.now() + timedelta(hours=24) # Reset 24 hours from now
key1 = UserAPIKeyAuth(

View file

@ -171,6 +171,25 @@ class TestHelperFunctions:
mock_response.choices[0].message.content = "Hello from LLM"
assert _extract_response_text(mock_response) == "Hello from LLM"
def test_extract_response_text_combines_all_choices(self):
from litellm.proxy.guardrails.guardrail_hooks.semantic_guard.semantic_guard import (
_extract_response_text,
)
first_choice = MagicMock()
first_choice.message.content = "first response"
second_choice = MagicMock()
second_choice.message.content = [
{"type": "text", "text": "second"},
{"type": "text", "text": "response"},
]
mock_response = MagicMock()
mock_response.choices = [first_choice, second_choice]
assert (
_extract_response_text(mock_response) == "first response\nsecond response"
)
def test_extract_response_text_empty(self):
from litellm.proxy.guardrails.guardrail_hooks.semantic_guard.semantic_guard import (
_extract_response_text,

View file

@ -6,10 +6,13 @@ Tests the model-based routing ID encoding/decoding used by the batch
and file proxy endpoints.
"""
from types import MappingProxyType
from litellm.proxy.openai_files_endpoints.common_utils import (
decode_model_from_file_id,
encode_file_id_with_model,
get_original_file_id,
prepare_data_with_credentials,
)
@ -50,6 +53,42 @@ class TestEncodeFileIdWithModel:
assert result.startswith("file-")
class TestPrepareDataWithCredentials:
def test_preserves_trusted_internal_credentials_snapshot(self):
data = {"file_id": "file-abc"}
credentials = {
"custom_llm_provider": "bedrock",
"s3_bucket_name": "safe-bucket",
}
prepare_data_with_credentials(
data=data,
credentials=credentials,
include_internal_credentials=True,
)
assert data["s3_bucket_name"] == "safe-bucket"
assert "custom_llm_provider" not in data
assert isinstance(
data["_litellm_internal_model_credentials"], type(MappingProxyType({}))
)
assert (
data["_litellm_internal_model_credentials"]["s3_bucket_name"]
== "safe-bucket"
)
def test_does_not_add_internal_credentials_by_default(self):
data = {"file_id": "file-abc"}
credentials = {
"custom_llm_provider": "bedrock",
"s3_bucket_name": "safe-bucket",
}
prepare_data_with_credentials(data=data, credentials=credentials)
assert "_litellm_internal_model_credentials" not in data
class TestRoundTrip:
"""Tests for encode -> decode round-trip integrity."""

View file

@ -18,6 +18,7 @@ from litellm import Router
# this tests debug logs from litellm router and litellm proxy server
from litellm._logging import verbose_logger, verbose_proxy_logger, verbose_router_logger
from litellm.llms.custom_httpx.async_client_cleanup import close_litellm_async_clients
# this tests debug logs from litellm router and litellm proxy server
@ -74,6 +75,9 @@ def test_async_fallbacks(caplog):
pytest.fail(f"An exception occurred: {e}")
finally:
router.reset()
# Close cached aiohttp/httpx clients before the event loop ends
# to prevent "Unclosed client session" / "Unclosed connector" warnings.
await close_litellm_async_clients()
asyncio.run(_make_request())
captured_logs = [rec.message for rec in caplog.records]

View file

@ -128,8 +128,8 @@ async def get_spend_info(session, entity_type: str, entity_id: str):
async def get_proxy_readiness(session):
"""Fetch /health/readiness. Used both as a fail-fast gate and as a diagnostic on poll timeout."""
url = "http://0.0.0.0:4000/health/readiness"
"""Fetch authenticated readiness details. Used both as a fail-fast gate and as a diagnostic on poll timeout."""
url = "http://0.0.0.0:4000/health/readiness/details"
headers = {"Authorization": "Bearer sk-1234"}
async with session.get(url, headers=headers) as response:
return response.status, await response.json()
@ -140,7 +140,7 @@ async def assert_proxy_healthy(session):
status, body = await get_proxy_readiness(session)
if status != 200 or body.get("db") != "connected":
pytest.fail(
f"Proxy /health/readiness unhealthy (status={status}). "
f"Proxy /health/readiness/details unhealthy (status={status}). "
f"Cannot run spend accuracy test. Response: {body}"
)
print(f"Proxy readiness OK: {body}")

View file

@ -73,13 +73,32 @@ async def test_health_readiness():
response_json = await response.json()
print(response_json)
assert "litellm_version" in response_json
assert "status" in response_json
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
@pytest.mark.asyncio
async def test_health_readiness_details():
"""
Check if authenticated readiness diagnostics expose version metadata.
"""
async with aiohttp.ClientSession() as session:
url = "http://0.0.0.0:4000/health/readiness/details"
headers = {"Authorization": "Bearer sk-1234"}
async with session.get(url, headers=headers) as response:
status = response.status
response_json = await response.json()
print(response_json)
assert "status" in response_json
assert "litellm_version" in response_json
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
@pytest.mark.asyncio
async def test_health_liveliness():
"""

View file

@ -19,9 +19,7 @@ def test_qdrant_semantic_cache_initialization(monkeypatch):
patch(
"litellm.llms.custom_httpx.http_handler._get_httpx_client"
) as mock_sync_client,
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
) as mock_async_client,
patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"),
):
# Mock the collection exists check
@ -31,6 +29,9 @@ def test_qdrant_semantic_cache_initialization(monkeypatch):
mock_sync_client_instance = MagicMock()
mock_sync_client_instance.get.return_value = mock_response
mock_index_response = MagicMock()
mock_index_response.status_code = 200
mock_sync_client_instance.put.return_value = mock_index_response
mock_sync_client.return_value = mock_sync_client_instance
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
@ -48,6 +49,17 @@ 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
mock_sync_client_instance.put.assert_called_once_with(
url="http://test.qdrant.local/collections/test_collection/index",
headers={
"Content-Type": "application/json",
"api-key": "test_key",
},
json={
"field_name": QdrantSemanticCache.CACHE_KEY_FIELD_NAME,
"field_schema": "keyword",
},
)
# Test initialization with missing similarity_threshold
with pytest.raises(Exception, match="similarity_threshold must be provided"):
@ -67,9 +79,7 @@ def test_qdrant_semantic_cache_get_cache_hit():
patch(
"litellm.llms.custom_httpx.http_handler._get_httpx_client"
) as mock_sync_client,
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
) as mock_async_client,
patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"),
):
# Mock the collection exists check
@ -98,6 +108,7 @@ def test_qdrant_semantic_cache_get_cache_hit():
"result": [
{
"payload": {
QdrantSemanticCache.CACHE_KEY_FIELD_NAME: "test_key",
"text": "What is the capital of France?", # Original prompt
"response": '{"id": "test-123", "choices": [{"message": {"content": "Paris is the capital of France."}}]}',
},
@ -127,6 +138,177 @@ def test_qdrant_semantic_cache_get_cache_hit():
# Verify search was called
qdrant_cache.sync_client.post.assert_called()
assert qdrant_cache.sync_client.post.call_args.kwargs["json"]["filter"] == {
"must": [
{
"key": QdrantSemanticCache.CACHE_KEY_FIELD_NAME,
"match": {"value": "test_key"},
}
]
}
def test_qdrant_semantic_cache_rejects_unscoped_cache_hit():
"""
Test QDRANT semantic cache rejects old or unscoped cache hits.
Legacy points have only text and response payloads, so they cannot be
safely migrated to a generated LiteLLM cache key.
"""
with (
patch(
"litellm.llms.custom_httpx.http_handler._get_httpx_client"
) as mock_sync_client,
patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"),
):
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {"result": {"exists": True}}
mock_sync_client_instance = MagicMock()
mock_sync_client_instance.get.return_value = mock_response
mock_sync_client.return_value = mock_sync_client_instance
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
qdrant_cache = QdrantSemanticCache(
collection_name="test_collection",
qdrant_api_base="http://test.qdrant.local",
qdrant_api_key="test_key",
similarity_threshold=0.8,
)
mock_search_response = MagicMock()
mock_search_response.status_code = 200
mock_search_response.json.return_value = {
"result": [
{
"payload": {
"text": "What is the capital of France?",
"response": '{"id": "test-123"}',
},
"score": 0.9,
}
]
}
qdrant_cache.sync_client.post = MagicMock(return_value=mock_search_response)
with patch(
"litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2, 0.3]}]}
):
metadata = {}
result = qdrant_cache.get_cache(
key="test_key",
messages=[{"content": "What is the capital of France?"}],
metadata=metadata,
)
assert result is None
assert metadata["semantic-similarity"] == 0.0
def test_qdrant_semantic_cache_payload_index_failure_is_non_blocking():
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
qdrant_cache = QdrantSemanticCache.__new__(QdrantSemanticCache)
qdrant_cache.qdrant_api_base = "http://test.qdrant.local"
qdrant_cache.collection_name = "test_collection"
qdrant_cache.headers = {"Content-Type": "application/json"}
qdrant_cache.sync_client = MagicMock()
response = MagicMock()
response.status_code = 400
response.text = "bad index"
qdrant_cache.sync_client.put.return_value = response
qdrant_cache._ensure_cache_key_payload_index()
qdrant_cache.sync_client.put.assert_called_once()
def test_qdrant_semantic_cache_payload_index_exception_is_non_blocking():
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
qdrant_cache = QdrantSemanticCache.__new__(QdrantSemanticCache)
qdrant_cache.qdrant_api_base = "http://test.qdrant.local"
qdrant_cache.collection_name = "test_collection"
qdrant_cache.headers = {"Content-Type": "application/json"}
qdrant_cache.sync_client = MagicMock()
qdrant_cache.sync_client.put.side_effect = Exception("boom")
qdrant_cache._ensure_cache_key_payload_index()
qdrant_cache.sync_client.put.assert_called_once()
def _mock_qdrant_get_cache_result(qdrant_result):
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
qdrant_cache = QdrantSemanticCache.__new__(QdrantSemanticCache)
qdrant_cache.embedding_model = "text-embedding-ada-002"
qdrant_cache.qdrant_api_base = "http://test.qdrant.local"
qdrant_cache.collection_name = "test_collection"
qdrant_cache.headers = {
"Content-Type": "application/json",
"api-key": "test_key",
}
qdrant_cache.similarity_threshold = 0.8
qdrant_cache.sync_client = MagicMock()
mock_search_response = MagicMock()
mock_search_response.status_code = 200
mock_search_response.json.return_value = {"result": qdrant_result}
qdrant_cache.sync_client.post.return_value = mock_search_response
return qdrant_cache, QdrantSemanticCache
@pytest.mark.parametrize("qdrant_result", [None, []])
def test_qdrant_semantic_cache_get_cache_sets_metadata_on_empty_miss(qdrant_result):
qdrant_cache, _ = _mock_qdrant_get_cache_result(qdrant_result)
metadata = {}
with patch(
"litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2, 0.3]}]}
):
result = qdrant_cache.get_cache(
key="test_key",
messages=[{"content": "What is the capital of Spain?"}],
metadata=metadata,
)
assert result is None
assert metadata["semantic-similarity"] == 0.0
def test_qdrant_semantic_cache_get_cache_sets_metadata_on_below_threshold_miss():
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
qdrant_cache, _ = _mock_qdrant_get_cache_result(
[
{
"payload": {
QdrantSemanticCache.CACHE_KEY_FIELD_NAME: "test_key",
"text": "What is the capital of Spain?",
"response": '{"id": "test-456"}',
},
"score": 0.7,
}
]
)
metadata = {}
with patch(
"litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2, 0.3]}]}
):
result = qdrant_cache.get_cache(
key="test_key",
messages=[{"content": "What is the capital of Spain?"}],
metadata=metadata,
)
assert result is None
assert metadata["semantic-similarity"] == 0.7
def test_qdrant_semantic_cache_get_cache_miss():
@ -138,9 +320,7 @@ def test_qdrant_semantic_cache_get_cache_miss():
patch(
"litellm.llms.custom_httpx.http_handler._get_httpx_client"
) as mock_sync_client,
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
) as mock_async_client,
patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"),
):
# Mock the collection exists check
@ -230,6 +410,7 @@ async def test_qdrant_semantic_cache_async_get_cache_hit():
"result": [
{
"payload": {
QdrantSemanticCache.CACHE_KEY_FIELD_NAME: "test_key",
"text": "What is the capital of Spain?", # Original prompt
"response": '{"id": "test-456", "choices": [{"message": {"content": "Madrid is the capital of Spain."}}]}',
},
@ -262,6 +443,16 @@ async def test_qdrant_semantic_cache_async_get_cache_hit():
# Verify async search was called
qdrant_cache.async_client.post.assert_called()
assert qdrant_cache.async_client.post.call_args.kwargs["json"][
"filter"
] == {
"must": [
{
"key": QdrantSemanticCache.CACHE_KEY_FIELD_NAME,
"match": {"value": "test_key"},
}
]
}
@pytest.mark.asyncio
@ -336,9 +527,7 @@ def test_qdrant_semantic_cache_set_cache():
patch(
"litellm.llms.custom_httpx.http_handler._get_httpx_client"
) as mock_sync_client,
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
) as mock_async_client,
patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"),
):
# Mock the collection exists check
@ -384,6 +573,12 @@ def test_qdrant_semantic_cache_set_cache():
# Verify upsert was called
qdrant_cache.sync_client.put.assert_called()
upsert_payload = qdrant_cache.sync_client.put.call_args.kwargs["json"][
"points"
][0]["payload"]
assert (
upsert_payload[QdrantSemanticCache.CACHE_KEY_FIELD_NAME] == "test_key"
)
@pytest.mark.asyncio
@ -450,6 +645,12 @@ async def test_qdrant_semantic_cache_async_set_cache():
# Verify async upsert was called
qdrant_cache.async_client.put.assert_called()
upsert_payload = qdrant_cache.async_client.put.call_args.kwargs["json"][
"points"
][0]["payload"]
assert (
upsert_payload[QdrantSemanticCache.CACHE_KEY_FIELD_NAME] == "test_key"
)
def test_qdrant_semantic_cache_custom_vector_size():
@ -462,9 +663,7 @@ def test_qdrant_semantic_cache_custom_vector_size():
patch(
"litellm.llms.custom_httpx.http_handler._get_httpx_client"
) as mock_sync_client,
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
) as mock_async_client,
patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"),
):
# Mock the collection does NOT exist (so it will be created)
@ -505,9 +704,13 @@ def test_qdrant_semantic_cache_custom_vector_size():
assert qdrant_cache.vector_size == 768
# Verify the PUT call to create the collection used vector_size=768
put_call = mock_sync_client_instance.put.call_args
assert put_call is not None
create_payload = put_call.kwargs.get("json") or put_call[1].get("json")
put_call = next(
call
for call in mock_sync_client_instance.put.call_args_list
if call.kwargs["url"]
== "http://test.qdrant.local/collections/test_collection_768"
)
create_payload = put_call.kwargs["json"]
assert create_payload["vectors"]["size"] == 768
assert create_payload["vectors"]["distance"] == "Cosine"
@ -521,9 +724,7 @@ def test_qdrant_semantic_cache_default_vector_size():
patch(
"litellm.llms.custom_httpx.http_handler._get_httpx_client"
) as mock_sync_client,
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
) as mock_async_client,
patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"),
):
# Mock the collection exists check
@ -559,9 +760,7 @@ def test_qdrant_semantic_cache_large_vector_size():
patch(
"litellm.llms.custom_httpx.http_handler._get_httpx_client"
) as mock_sync_client,
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
) as mock_async_client,
patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"),
):
# Mock the collection does NOT exist (so it will be created)
@ -599,6 +798,11 @@ def test_qdrant_semantic_cache_large_vector_size():
assert qdrant_cache.vector_size == 4096
# Verify the collection was created with 4096
put_call = mock_sync_client_instance.put.call_args
create_payload = put_call.kwargs.get("json") or put_call[1].get("json")
put_call = next(
call
for call in mock_sync_client_instance.put.call_args_list
if call.kwargs["url"]
== "http://test.qdrant.local/collections/test_collection_4096"
)
create_payload = put_call.kwargs["json"]
assert create_payload["vectors"]["size"] == 4096

View file

@ -72,24 +72,301 @@ def test_redis_semantic_cache_get_cache(monkeypatch):
"prompt": "What is the capital of France?",
"response": '{"content": "Paris is the capital of France."}',
"vector_distance": 0.1, # Distance of 0.1 means similarity of 0.9
RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key",
}
]
redis_semantic_cache.llmcache.check = MagicMock(return_value=mock_result)
# Mock the embedding function
with patch(
"litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2, 0.3]}]}
with (
patch(
"litellm.embedding",
return_value={"data": [{"embedding": [0.1, 0.2, 0.3]}]},
),
patch.object(
redis_semantic_cache,
"_get_cache_key_filter_expression",
return_value="cache-key-filter",
),
):
# Test get_cache with a message
metadata = {}
result = redis_semantic_cache.get_cache(
key="test_key", messages=[{"content": "What is the capital of France?"}]
key="test_key",
messages=[{"content": "What is the capital of France?"}],
metadata=metadata,
)
# Verify result is properly parsed
assert result == {"content": "Paris is the capital of France."}
assert metadata["semantic-similarity"] == pytest.approx(0.9)
# Verify llmcache.check was called
redis_semantic_cache.llmcache.check.assert_called_once()
redis_semantic_cache.llmcache.check.assert_called_once_with(
prompt="What is the capital of France?",
filter_expression="cache-key-filter",
)
def test_redis_semantic_cache_rejects_unscoped_cache_hit(monkeypatch):
semantic_cache_mock = MagicMock()
custom_vectorizer_mock = MagicMock()
with patch.dict(
"sys.modules",
{
"redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock),
"redisvl.utils.vectorize": MagicMock(
CustomTextVectorizer=custom_vectorizer_mock
),
},
):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_PORT", "6379")
monkeypatch.setenv("REDIS_PASSWORD", "test_password")
redis_semantic_cache = RedisSemanticCache(similarity_threshold=0.8)
redis_semantic_cache.llmcache.check = MagicMock(
return_value=[
{
"prompt": "What is the capital of France?",
"response": '{"content": "Paris"}',
"vector_distance": 0.1,
}
]
)
with patch.object(
redis_semantic_cache,
"_get_cache_key_filter_expression",
return_value="cache-key-filter",
):
metadata = {}
result = redis_semantic_cache.get_cache(
key="test_key",
messages=[{"content": "What is the capital of France?"}],
metadata=metadata,
)
assert result is None
assert metadata["semantic-similarity"] == 0.0
def test_redis_semantic_cache_set_cache_stores_cache_key_filter(monkeypatch):
semantic_cache_mock = MagicMock()
custom_vectorizer_mock = MagicMock()
with patch.dict(
"sys.modules",
{
"redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock),
"redisvl.utils.vectorize": MagicMock(
CustomTextVectorizer=custom_vectorizer_mock
),
},
):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_PORT", "6379")
monkeypatch.setenv("REDIS_PASSWORD", "test_password")
redis_semantic_cache = RedisSemanticCache(similarity_threshold=0.8)
redis_semantic_cache.llmcache.store = MagicMock()
redis_semantic_cache.set_cache(
key="test_key",
value={"content": "Paris"},
messages=[{"content": "What is the capital of France?"}],
ttl=60,
)
redis_semantic_cache.llmcache.store.assert_called_once_with(
"What is the capital of France?",
"{'content': 'Paris'}",
filters={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"},
ttl=60,
)
def test_redis_semantic_cache_uses_isolated_index_for_old_schema(monkeypatch):
fallback_cache_mock = MagicMock()
semantic_cache_mock = MagicMock(
side_effect=[
ValueError("stored index schema differs from requested fields"),
fallback_cache_mock,
]
)
custom_vectorizer_mock = MagicMock()
with patch.dict(
"sys.modules",
{
"redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock),
"redisvl.utils.vectorize": MagicMock(
CustomTextVectorizer=custom_vectorizer_mock
),
},
):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_PORT", "6379")
monkeypatch.setenv("REDIS_PASSWORD", "test_password")
redis_semantic_cache = RedisSemanticCache(
similarity_threshold=0.8,
index_name="existing_index",
)
assert redis_semantic_cache.llmcache is fallback_cache_mock
assert semantic_cache_mock.call_args_list[0].kwargs["name"] == "existing_index"
assert (
semantic_cache_mock.call_args_list[1].kwargs["name"]
== "existing_index_isolated"
)
assert semantic_cache_mock.call_args_list[1].kwargs["filterable_fields"] == [
RedisSemanticCache._cache_key_filterable_field()
]
def test_redis_semantic_cache_overwrites_stale_isolated_index(monkeypatch):
fallback_cache_mock = MagicMock()
semantic_cache_mock = MagicMock(
side_effect=[
ValueError("Existing index schema does not match"),
ValueError("Existing index schema does not match"),
fallback_cache_mock,
]
)
custom_vectorizer_mock = MagicMock()
with patch.dict(
"sys.modules",
{
"redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock),
"redisvl.utils.vectorize": MagicMock(
CustomTextVectorizer=custom_vectorizer_mock
),
},
):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_PORT", "6379")
monkeypatch.setenv("REDIS_PASSWORD", "test_password")
redis_semantic_cache = RedisSemanticCache(
similarity_threshold=0.8,
index_name="existing_index",
)
assert redis_semantic_cache.llmcache is fallback_cache_mock
assert (
semantic_cache_mock.call_args_list[2].kwargs["name"]
== "existing_index_isolated"
)
assert semantic_cache_mock.call_args_list[2].kwargs["overwrite"] is True
assert semantic_cache_mock.call_args_list[2].kwargs["filterable_fields"] == [
RedisSemanticCache._cache_key_filterable_field()
]
def test_redis_semantic_cache_reraises_unexpected_isolated_index_error(monkeypatch):
semantic_cache_mock = MagicMock(
side_effect=[
ValueError("Existing index schema does not match"),
ValueError("connection failed"),
]
)
custom_vectorizer_mock = MagicMock()
with patch.dict(
"sys.modules",
{
"redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock),
"redisvl.utils.vectorize": MagicMock(
CustomTextVectorizer=custom_vectorizer_mock
),
},
):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_PORT", "6379")
monkeypatch.setenv("REDIS_PASSWORD", "test_password")
with pytest.raises(ValueError, match="connection failed"):
RedisSemanticCache(
similarity_threshold=0.8,
index_name="existing_index",
)
def test_redis_semantic_cache_reraises_unexpected_index_error():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
redis_semantic_cache.distance_threshold = 0.2
semantic_cache_mock = MagicMock(side_effect=ValueError("connection failed"))
with pytest.raises(ValueError, match="connection failed"):
redis_semantic_cache._init_semantic_cache(
semantic_cache_cls=semantic_cache_mock,
index_name="existing_index",
redis_url="redis://localhost:6379",
cache_vectorizer=MagicMock(),
)
def test_redis_semantic_cache_matches_bytes_cache_key():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
assert redis_semantic_cache._cache_hit_matches_key(
cache_hit={RedisSemanticCache.CACHE_KEY_FIELD_NAME: b"test_key"},
key="test_key",
)
def test_redis_semantic_cache_rejects_pre_isolation_unscoped_hit():
"""Pre-isolation entries with no cache-key field cannot be safely
reassigned to a caller's scope and are treated as misses."""
from litellm.caching.redis_semantic_cache import RedisSemanticCache
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
cache_hit = {
"prompt": "What is the capital of France?",
"response": '{"content": "Paris"}',
"vector_distance": 0.1,
}
assert not redis_semantic_cache._cache_hit_matches_key(
cache_hit=cache_hit,
key="test_key",
)
def test_redis_semantic_cache_builds_filter_expression(monkeypatch):
class FakeTag:
def __init__(self, field_name):
self.field_name = field_name
def __eq__(self, value):
return (self.field_name, value)
with patch.dict("sys.modules", {"redisvl.query.filter": MagicMock(Tag=FakeTag)}):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
assert redis_semantic_cache._get_cache_key_filter_expression("test_key") == (
RedisSemanticCache.CACHE_KEY_FIELD_NAME,
"test_key",
)
@pytest.mark.asyncio
@ -123,6 +400,7 @@ async def test_redis_semantic_cache_async_get_cache(monkeypatch):
"prompt": "What is the capital of France?",
"response": '{"content": "Paris is the capital of France."}',
"vector_distance": 0.1, # Distance of 0.1 means similarity of 0.9
RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key",
}
]
@ -131,16 +409,117 @@ async def test_redis_semantic_cache_async_get_cache(monkeypatch):
return_value=[0.1, 0.2, 0.3]
)
# Test async_get_cache with a message
result = await redis_semantic_cache.async_get_cache(
key="test_key",
messages=[{"content": "What is the capital of France?"}],
metadata={},
)
with patch.object(
redis_semantic_cache,
"_get_cache_key_filter_expression",
return_value="cache-key-filter",
):
# Test async_get_cache with a message
result = await redis_semantic_cache.async_get_cache(
key="test_key",
messages=[{"content": "What is the capital of France?"}],
metadata={},
)
# Verify result is properly parsed
assert result == {"content": "Paris is the capital of France."}
# Verify methods were called
redis_semantic_cache._get_async_embedding.assert_called_once()
redis_semantic_cache.llmcache.acheck.assert_called_once()
redis_semantic_cache.llmcache.acheck.assert_called_once_with(
prompt="What is the capital of France?",
vector=[0.1, 0.2, 0.3],
filter_expression="cache-key-filter",
)
@pytest.mark.asyncio
async def test_redis_semantic_cache_async_get_cache_rejects_unscoped_hit(monkeypatch):
semantic_cache_mock = MagicMock()
custom_vectorizer_mock = MagicMock()
with patch.dict(
"sys.modules",
{
"redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock),
"redisvl.utils.vectorize": MagicMock(
CustomTextVectorizer=custom_vectorizer_mock
),
},
):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_PORT", "6379")
monkeypatch.setenv("REDIS_PASSWORD", "test_password")
redis_semantic_cache = RedisSemanticCache(similarity_threshold=0.8)
redis_semantic_cache.llmcache.acheck = AsyncMock(
return_value=[
{
"prompt": "What is the capital of France?",
"response": '{"content": "Paris"}',
"vector_distance": 0.1,
}
]
)
redis_semantic_cache._get_async_embedding = AsyncMock(
return_value=[0.1, 0.2, 0.3]
)
with patch.object(
redis_semantic_cache,
"_get_cache_key_filter_expression",
return_value="cache-key-filter",
):
result = await redis_semantic_cache.async_get_cache(
key="test_key",
messages=[{"content": "What is the capital of France?"}],
metadata={},
)
assert result is None
@pytest.mark.asyncio
async def test_redis_semantic_cache_async_set_cache_stores_cache_key_filter(
monkeypatch,
):
semantic_cache_mock = MagicMock()
custom_vectorizer_mock = MagicMock()
with patch.dict(
"sys.modules",
{
"redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock),
"redisvl.utils.vectorize": MagicMock(
CustomTextVectorizer=custom_vectorizer_mock
),
},
):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_PORT", "6379")
monkeypatch.setenv("REDIS_PASSWORD", "test_password")
redis_semantic_cache = RedisSemanticCache(similarity_threshold=0.8)
redis_semantic_cache.llmcache.astore = AsyncMock()
redis_semantic_cache._get_async_embedding = AsyncMock(
return_value=[0.1, 0.2, 0.3]
)
await redis_semantic_cache.async_set_cache(
key="test_key",
value={"content": "Paris"},
messages=[{"content": "What is the capital of France?"}],
ttl=60,
)
redis_semantic_cache.llmcache.astore.assert_called_once_with(
"What is the capital of France?",
"{'content': 'Paris'}",
vector=[0.1, 0.2, 0.3],
filters={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"},
ttl=60,
)

View file

@ -1,5 +1,9 @@
import os
from unittest.mock import patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.integrations.gcs_bucket.gcs_bucket import GCSBucketLogger
from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase
@ -86,3 +90,41 @@ class TestGCSBucketBase:
project_id=None, # Should be None when no env var is set
custom_llm_provider="vertex_ai",
)
@pytest.mark.asyncio
async def test_log_json_data_on_gcs_url_encodes_object_name(self):
handler = GCSBucketBase(bucket_name="test-bucket")
handler.async_httpx_client = AsyncMock()
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {"name": "logs/object"}
handler.async_httpx_client.post.return_value = mock_response
await handler._log_json_data_on_gcs(
headers={"Authorization": "Bearer token"},
bucket_name="test-bucket",
object_name="logs/object?uploadType=media&name=evil",
logging_payload={"ok": True},
)
post_url = handler.async_httpx_client.post.call_args.kwargs["url"]
assert "name=logs%2Fobject%3FuploadType%3Dmedia%26name%3Devil" in post_url
assert "name=logs/object?" not in post_url
def test_gcs_log_id_is_only_used_as_sanitized_hint(self):
logger = GCSBucketLogger.__new__(GCSBucketLogger)
object_name = logger._get_object_name(
kwargs={
"litellm_params": {
"metadata": {"gcs_log_id": "../../target?uploadType=media"}
}
},
logging_payload={"id": "payload"},
response_obj={"id": "response-id"},
)
assert "/custom-" in object_name
assert object_name.endswith("-target_uploadType_media")
assert ".." not in object_name
assert "?" not in object_name

View file

@ -64,7 +64,7 @@ def test_env_reference_in_metadata_raises_with_guidance():
assert "metadata" in message
def test_env_reference_in_litellm_params_metadata_raises():
def test_gcs_bucket_name_in_litellm_params_metadata_is_ignored():
kwargs = {
"litellm_params": {
"metadata": {
@ -73,10 +73,21 @@ def test_env_reference_in_litellm_params_metadata_raises():
}
}
with pytest.raises(ValueError) as exc_info:
initialize_standard_callback_dynamic_params(kwargs)
params = initialize_standard_callback_dynamic_params(kwargs)
assert "gcs_bucket_name" in str(exc_info.value)
assert params.get("gcs_bucket_name") is None
def test_gcs_callback_params_are_not_extracted_from_request_kwargs():
kwargs = {
"gcs_bucket_name": "server-bucket",
"gcs_path_service_account": "/path/to/service-account.json",
}
params = initialize_standard_callback_dynamic_params(kwargs)
assert params.get("gcs_bucket_name") is None
assert params.get("gcs_path_service_account") is None
def test_non_string_values_are_not_flagged():

View file

@ -7,6 +7,7 @@ from litellm.litellm_core_utils import url_utils
from litellm.litellm_core_utils.url_utils import (
SSRFError,
_is_blocked_ip,
assert_same_origin,
encode_url_path_segment,
encode_url_path_segments,
validate_url,
@ -424,12 +425,51 @@ class TestHostAllowlist:
validate_url("http://internal.corp/")
class TestProviderUrlDestinationAllowlist:
def test_host_entry_matches_any_scheme_and_port(self):
assert url_utils.is_url_destination_allowed_by_host(
"https://trusted.example/v1/chat/completions",
["trusted.example"],
)
assert url_utils.is_url_destination_allowed_by_host(
"http://trusted.example:8080/v1/chat/completions",
["trusted.example"],
)
def test_origin_entry_matches_scheme_and_default_port(self):
assert url_utils.is_url_destination_allowed_by_host(
"https://trusted.example/v1/chat/completions",
["https://trusted.example"],
)
assert not url_utils.is_url_destination_allowed_by_host(
"http://trusted.example/v1/chat/completions",
["https://trusted.example"],
)
def test_port_entry_only_matches_same_effective_port(self):
assert url_utils.is_url_destination_allowed_by_host(
"https://trusted.example/v1/chat/completions",
["trusted.example:443"],
)
assert not url_utils.is_url_destination_allowed_by_host(
"https://trusted.example:8443/v1/chat/completions",
["trusted.example:443"],
)
def test_rejects_userinfo_and_invalid_port(self):
assert not url_utils.is_url_destination_allowed_by_host(
"https://user:pass@trusted.example/v1/chat/completions",
["trusted.example"],
)
assert not url_utils.is_url_destination_allowed_by_host(
"https://trusted.example:99999/v1/chat/completions",
["trusted.example"],
)
# ── assert_same_origin ────────────────────────────────────────────────────────
from litellm.litellm_core_utils.url_utils import assert_same_origin
def test_assert_same_origin_matches_scheme_host_port():
"""A polling URL on the same scheme + host + port as the api_base
passes — the upstream is trusted; the URL it returned points back at

View file

@ -0,0 +1,23 @@
"""Force ``litellm.model_cost`` to load from the PR-local JSON for these tests.
By default ``litellm.model_cost`` is fetched from the main branch on GitHub,
which lags behind PR-branch flag additions. This fixture loads the local
file so per-model flag tests pass in CI as well as locally.
See https://github.com/BerriAI/litellm/issues/27122.
"""
import pytest
import litellm
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
@pytest.fixture(autouse=True)
def _use_pr_local_model_cost_map(monkeypatch):
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(
litellm,
"model_cost",
get_model_cost_map(url=litellm.model_cost_map_url),
)
yield

View file

@ -8,6 +8,7 @@ sys.path.insert(
) # Adds the parent directory to the system path
from unittest.mock import MagicMock, patch
import litellm
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
@ -1631,8 +1632,9 @@ def test_effort_validation():
)
assert result["output_config"]["effort"] == effort
# Invalid value should raise error
with pytest.raises(ValueError, match="Invalid effort value"):
with pytest.raises(
litellm.exceptions.BadRequestError, match="Invalid effort value"
):
optional_params = {"output_config": {"effort": "invalid"}}
config.transform_request(
model="claude-opus-4-5-20251101",
@ -1687,7 +1689,10 @@ def test_max_effort_rejected_for_opus_45():
messages = [{"role": "user", "content": "Test"}]
with pytest.raises(ValueError, match="effort='max' is not supported by this model"):
with pytest.raises(
litellm.exceptions.BadRequestError,
match="effort='max' is not supported by this model",
):
optional_params = {"output_config": {"effort": "max"}}
config.transform_request(
model="claude-opus-4-5-20251101",
@ -1739,6 +1744,97 @@ def test_effort_with_other_features():
assert "thinking" in result
def test_anthropic_drop_params_strips_output_config_for_pre_4_5_models():
"""``drop_params=True`` strips unsupported ``output_config`` for pre-4.5 models."""
config = AnthropicConfig()
messages = [{"role": "user", "content": "Hello"}]
original = litellm.drop_params
litellm.drop_params = True
try:
result = config.transform_request(
model="claude-3-haiku-20240307",
messages=messages,
optional_params={"output_config": {"effort": "low"}},
litellm_params={},
headers={},
)
finally:
litellm.drop_params = original
assert "output_config" not in result
def test_anthropic_drop_params_keeps_output_config_for_supporting_models():
"""``drop_params=True`` must not strip on models that support effort."""
config = AnthropicConfig()
messages = [{"role": "user", "content": "Hello"}]
original = litellm.drop_params
litellm.drop_params = True
try:
result = config.transform_request(
model="claude-opus-4-7",
messages=messages,
optional_params={"output_config": {"effort": "high"}},
litellm_params={},
headers={},
)
finally:
litellm.drop_params = original
assert result.get("output_config") == {"effort": "high"}
def test_anthropic_drop_params_false_forwards_to_unsupported_model():
"""Default ``drop_params=False`` forwards ``output_config`` and lets the provider 400."""
config = AnthropicConfig()
messages = [{"role": "user", "content": "Hello"}]
original = litellm.drop_params
litellm.drop_params = False
try:
result = config.transform_request(
model="claude-3-haiku-20240307",
messages=messages,
optional_params={"output_config": {"effort": "low"}},
litellm_params={},
headers={},
)
finally:
litellm.drop_params = original
assert result.get("output_config") == {"effort": "low"}
@pytest.mark.parametrize(
"model",
[
"claude-opus-4-5-20251101",
"claude-opus-4-6",
"claude-opus-4-7",
"claude-sonnet-4-6",
"anthropic.claude-mythos-preview",
"bedrock/anthropic.claude-mythos-preview",
],
)
def test_anthropic_model_supports_effort_param_recognizes_supporting_models(model):
assert AnthropicConfig._model_supports_effort_param(model) is True
@pytest.mark.parametrize(
"model",
[
"claude-3-haiku-20240307",
"claude-3-5-sonnet-20241022",
"claude-3-opus-20240229",
"claude-sonnet-4-20250514",
],
)
def test_anthropic_model_supports_effort_param_rejects_non_supporting_models(model):
assert AnthropicConfig._model_supports_effort_param(model) is False
def test_translate_system_message_skips_empty_string_content():
"""
Test that translate_system_message skips system messages with empty string content.
@ -1950,6 +2046,67 @@ def test_get_config_without_model_uses_fallback():
assert config["max_tokens"] == 4096
def test_get_config_does_not_leak_module_constants():
"""``get_config`` must not leak the reasoning-effort lookup table onto the wire."""
cfg = AnthropicConfig.get_config(model="claude-opus-4-7")
for forbidden in (
"REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT",
"_REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT",
):
assert forbidden not in cfg
@pytest.mark.parametrize(
"model,level,expected",
[
("claude-opus-4-7", "max", True),
("claude-opus-4-7", "xhigh", True),
("claude-opus-4-6", "max", True),
("claude-opus-4-6", "xhigh", False),
("claude-sonnet-4-6", "max", True),
("claude-sonnet-4-6", "xhigh", False),
("bedrock/invoke/us.anthropic.claude-opus-4-7", "max", True),
("bedrock/invoke/us.anthropic.claude-opus-4-7", "xhigh", True),
("bedrock/invoke/us.anthropic.claude-opus-4-6-v1", "max", True),
("bedrock/invoke/us.anthropic.claude-opus-4-6-v1", "xhigh", False),
("bedrock/invoke/us.anthropic.claude-sonnet-4-6", "max", True),
("vertex_ai/claude-opus-4-7", "xhigh", True),
("azure_ai/claude-opus-4-7", "xhigh", True),
],
)
def test_supports_effort_level_handles_provider_prefixes(model, level, expected):
"""``_supports_effort_level`` resolves bedrock/vertex/azure-prefixed model ids."""
assert AnthropicConfig._supports_effort_level(model, level) is expected
@pytest.mark.parametrize(
"model,effort,expect_error",
[
("claude-opus-4-6", "max", False),
("claude-sonnet-4-6", "max", False),
("claude-opus-4-7", "max", False),
("claude-opus-4-5-20251101", "max", True),
("claude-sonnet-4-5", "max", True),
("claude-opus-4-7", "xhigh", False),
("claude-opus-4-6", "xhigh", True),
("claude-sonnet-4-6", "xhigh", True),
("claude-opus-4-5-20251101", "high", False),
("claude-haiku-4-5", "low", False),
("claude-opus-4-5-20251101", None, False),
],
)
def test_validate_effort_for_model_centralises_per_model_gating(
model, effort, expect_error
):
err = AnthropicConfig._validate_effort_for_model(model, effort)
if expect_error:
assert err is not None
assert effort in err
assert model in err
else:
assert err is None
def test_transform_request_uses_dynamic_max_tokens():
"""
Test that transform_request uses dynamic max_tokens based on model
@ -2153,12 +2310,12 @@ def test_reasoning_effort_maps_to_budget_thinking_for_non_opus_4_6():
"""
config = AnthropicConfig()
# Test with Claude Sonnet 4.5 (non-Opus 4.6 model)
# ``minimal`` floors at ANTHROPIC_MIN_THINKING_BUDGET_TOKENS (1024).
test_cases = [
("low", 1024), # DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET
("medium", 2048), # DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET
("high", 4096), # DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET
("minimal", 128), # DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET
("low", 1024),
("medium", 2048),
("high", 4096),
("minimal", 1024),
]
for effort, expected_budget in test_cases:
@ -2244,19 +2401,31 @@ def test_reasoning_effort_does_not_set_output_config_for_older_models():
), f"output_config should not be set for {model}"
def test_max_effort_rejected_for_sonnet_46():
"""Test that effort='max' is rejected for Sonnet 4.6 (Opus-only effort level)."""
@pytest.mark.parametrize(
"model",
[
"claude-sonnet-4-6",
"claude-sonnet-4-6-20260219",
"us.anthropic.claude-sonnet-4-6",
"bedrock/converse/us.anthropic.claude-sonnet-4-6",
"vertex_ai/claude-sonnet-4-6",
"openrouter/anthropic/claude-sonnet-4.6",
],
)
def test_max_effort_accepted_for_sonnet_46_variants(model):
"""``effort='max'`` is supported on Claude 4.6 (Opus + Sonnet) and 4.7."""
config = AnthropicConfig()
messages = [{"role": "user", "content": "Test"}]
with pytest.raises(ValueError, match="effort='max' is not supported by this model"):
config.transform_request(
model="claude-sonnet-4-6-20260219",
messages=messages,
optional_params={"output_config": {"effort": "max"}},
litellm_params={},
headers={},
)
result = config.transform_request(
model=model,
messages=messages,
optional_params={"output_config": {"effort": "max"}},
litellm_params={},
headers={},
)
assert result["output_config"]["effort"] == "max"
def test_max_effort_accepted_for_opus_46():
@ -2335,6 +2504,74 @@ def test_reasoning_effort_none_omits_thinking_and_output_config(model):
assert "output_config" not in result
@pytest.mark.parametrize(
"effort",
["disabled", "invalid", ""],
)
def test_reasoning_effort_garbage_raises_bad_request(effort):
"""Unmapped reasoning_effort raises BadRequestError (clean 400, not a 500)."""
config = AnthropicConfig()
with pytest.raises(litellm.exceptions.BadRequestError):
config.map_openai_params(
non_default_params={"reasoning_effort": effort},
optional_params={},
model="claude-sonnet-4-5-20250929",
drop_params=False,
)
@pytest.mark.parametrize(
"effort,expected_budget",
[("xhigh", 8192), ("max", 16384)],
)
def test_reasoning_effort_xhigh_max_maps_to_budget_on_budget_model(
effort, expected_budget
):
"""``xhigh`` / ``max`` extend the budget_tokens progression on budget-mode models."""
config = AnthropicConfig()
result = config.map_openai_params(
non_default_params={"reasoning_effort": effort},
optional_params={},
model="claude-sonnet-4-5-20250929",
drop_params=False,
)
assert result["thinking"]["type"] == "enabled"
assert result["thinking"]["budget_tokens"] == expected_budget
assert "output_config" not in result
def test_output_config_effort_empty_string_raises_bad_request():
"""``output_config={"effort": ""}`` is rejected with a 400."""
config = AnthropicConfig()
with pytest.raises(litellm.exceptions.BadRequestError, match="Invalid effort"):
config.transform_request(
model="claude-opus-4-7",
messages=[{"role": "user", "content": "hi"}],
optional_params={"output_config": {"effort": ""}, "max_tokens": 32},
litellm_params={},
headers={},
)
def test_reasoning_effort_minimal_floors_at_anthropic_provider_minimum():
"""``minimal`` floors at the Anthropic provider minimum (1024)."""
config = AnthropicConfig()
result = config.map_openai_params(
non_default_params={"reasoning_effort": "minimal"},
optional_params={},
model="claude-sonnet-4-5-20250929",
drop_params=False,
)
assert result["thinking"]["type"] == "enabled"
assert result["thinking"]["budget_tokens"] >= 1024
def test_effort_beta_header_still_injected_for_older_models():
"""
Test that is_effort_used still returns True for pre-4.6 models

View file

@ -0,0 +1,193 @@
"""Tests for ``reasoning_effort`` translation on the Anthropic /v1/messages route."""
import pytest
from litellm.llms.anthropic.common_utils import AnthropicError
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
@pytest.mark.parametrize(
"reasoning_effort,expected_effort",
[
("minimal", "low"),
("low", "low"),
("medium", "medium"),
("high", "high"),
("xhigh", "xhigh"),
("max", "max"),
],
)
def test_reasoning_effort_maps_to_output_config_for_adaptive_model(
reasoning_effort, expected_effort
):
config = AnthropicMessagesConfig()
optional_params = {"max_tokens": 1024, "reasoning_effort": reasoning_effort}
result = config.transform_anthropic_messages_request(
model="claude-opus-4-7",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert "reasoning_effort" not in result
assert result.get("thinking") == {"type": "adaptive"}
assert result.get("output_config") == {"effort": expected_effort}
def test_reasoning_effort_none_clears_thinking_and_output_config():
config = AnthropicMessagesConfig()
optional_params = {
"max_tokens": 1024,
"reasoning_effort": "none",
"thinking": {"type": "adaptive"},
"output_config": {"effort": "high"},
}
result = config.transform_anthropic_messages_request(
model="claude-opus-4-7",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert "reasoning_effort" not in result
assert "thinking" not in result
assert "output_config" not in result
def test_reasoning_effort_on_non_adaptive_model_uses_thinking_budget():
config = AnthropicMessagesConfig()
optional_params = {"max_tokens": 1024, "reasoning_effort": "high"}
result = config.transform_anthropic_messages_request(
model="claude-opus-4-5",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert "reasoning_effort" not in result
assert "output_config" not in result
thinking = result.get("thinking")
assert isinstance(thinking, dict)
assert thinking.get("type") == "enabled"
assert isinstance(thinking.get("budget_tokens"), int)
assert thinking["budget_tokens"] >= 1024
@pytest.mark.parametrize("bad_effort", ["invalid", "disabled", ""])
def test_invalid_reasoning_effort_raises_400(bad_effort):
config = AnthropicMessagesConfig()
optional_params = {"max_tokens": 1024, "reasoning_effort": bad_effort}
with pytest.raises(AnthropicError) as exc_info:
config.transform_anthropic_messages_request(
model="claude-opus-4-7",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert exc_info.value.status_code == 400
@pytest.mark.parametrize(
"model,bad_effort",
[
("claude-opus-4-6", "xhigh"),
("bedrock/invoke/us.anthropic.claude-opus-4-6-v1", "xhigh"),
("claude-sonnet-4-6", "xhigh"),
],
)
def test_reasoning_effort_unsupported_tier_raises_400_messages(model, bad_effort):
config = AnthropicMessagesConfig()
optional_params = {"max_tokens": 1024, "reasoning_effort": bad_effort}
with pytest.raises(AnthropicError) as exc_info:
config.transform_anthropic_messages_request(
model=model,
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert exc_info.value.status_code == 400
assert "not supported by this model" in str(exc_info.value)
@pytest.mark.parametrize(
"model",
[
"claude-sonnet-4-6",
"bedrock/invoke/us.anthropic.claude-sonnet-4-6",
],
)
def test_reasoning_effort_max_accepted_on_sonnet_46_messages(model):
config = AnthropicMessagesConfig()
optional_params = {"max_tokens": 1024, "reasoning_effort": "max"}
result = config.transform_anthropic_messages_request(
model=model,
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
output_config = result.get("output_config")
assert isinstance(output_config, dict) and output_config.get("effort") == "max"
def test_explicit_output_config_wins_over_reasoning_effort():
config = AnthropicMessagesConfig()
optional_params = {
"max_tokens": 1024,
"reasoning_effort": "low",
"output_config": {"effort": "max"},
}
result = config.transform_anthropic_messages_request(
model="claude-opus-4-7",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert "reasoning_effort" not in result
assert result.get("output_config") == {"effort": "max"}
def test_explicit_thinking_wins_over_reasoning_effort():
config = AnthropicMessagesConfig()
optional_params = {
"max_tokens": 1024,
"reasoning_effort": "low",
"thinking": {"type": "enabled", "budget_tokens": 8000},
}
result = config.transform_anthropic_messages_request(
model="claude-opus-4-5",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert "reasoning_effort" not in result
assert result.get("thinking") == {"type": "enabled", "budget_tokens": 8000}
def test_reasoning_effort_in_supported_params():
config = AnthropicMessagesConfig()
assert "reasoning_effort" in config.get_supported_anthropic_messages_params(
"claude-opus-4-7"
)

View file

@ -291,6 +291,100 @@ class TestAzureAnthropicConfig:
assert "anthropic-beta" in headers
assert "compact-2026-01-12" in headers["anthropic-beta"]
def test_output_config_promoted_from_extra_body(self):
"""Anthropic ``output_config`` routed via ``extra_body`` reaches the request body."""
config = AzureAnthropicConfig()
messages = [{"role": "user", "content": "Hello"}]
optional_params = {
"max_tokens": 100,
"extra_body": {"output_config": {"effort": "low"}},
}
litellm_params = {"api_key": "test-key"}
headers = {"api-key": "test-key", "anthropic-version": "2023-06-01"}
result = config.transform_request(
model="claude-opus-4-6",
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
assert result["output_config"] == {"effort": "low"}
assert "extra_body" not in result
def test_invalid_output_config_effort_raises_via_extra_body(self):
"""Invalid ``effort`` via ``extra_body`` raises BadRequestError."""
import litellm
config = AzureAnthropicConfig()
messages = [{"role": "user", "content": "Hello"}]
optional_params = {
"max_tokens": 100,
"extra_body": {"output_config": {"effort": "invalid"}},
}
litellm_params = {"api_key": "test-key"}
headers = {"api-key": "test-key", "anthropic-version": "2023-06-01"}
with pytest.raises(litellm.exceptions.BadRequestError) as exc_info:
config.transform_request(
model="claude-opus-4-6",
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
assert "Invalid effort value" in str(exc_info.value)
def test_unsupported_effort_xhigh_raises_via_extra_body(self):
"""Unsupported ``effort='xhigh'`` via ``extra_body`` raises BadRequestError."""
import litellm
config = AzureAnthropicConfig()
messages = [{"role": "user", "content": "Hello"}]
optional_params = {
"max_tokens": 100,
"extra_body": {"output_config": {"effort": "xhigh"}},
}
litellm_params = {"api_key": "test-key"}
headers = {"api-key": "test-key", "anthropic-version": "2023-06-01"}
with pytest.raises(litellm.exceptions.BadRequestError) as exc_info:
config.transform_request(
model="claude-sonnet-4-6",
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
assert "xhigh" in str(exc_info.value)
def test_extra_body_promotion_does_not_clobber_top_level(self):
"""Top-level ``optional_params`` wins over duplicates in ``extra_body``."""
config = AzureAnthropicConfig()
messages = [{"role": "user", "content": "Hello"}]
optional_params = {
"max_tokens": 100,
"output_config": {"effort": "low"},
"extra_body": {"output_config": {"effort": "high"}},
}
litellm_params = {"api_key": "test-key"}
headers = {"api-key": "test-key", "anthropic-version": "2023-06-01"}
result = config.transform_request(
model="claude-opus-4-6",
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
assert result["output_config"] == {"effort": "low"}
def test_context_management_mixed_edits_beta_headers(self):
"""Test that context_management with both compact and other edits adds both beta headers"""
config = AzureAnthropicConfig()

View file

@ -406,37 +406,26 @@ def test_opus_4_5_model_detection():
# f"computer-use beta should be kept, got: {anthropic_beta}"
def test_output_config_removed_from_bedrock_chat_invoke_request():
"""
Test that output_config parameter is stripped from Bedrock Chat Invoke requests.
Bedrock Invoke API doesn't support the output_config parameter (Anthropic-only).
Ensures the chat/invoke path mirrors the messages/invoke path fix.
Fixes: https://github.com/BerriAI/litellm/issues/22797
"""
def test_output_config_forwarded_for_bedrock_chat_invoke_request():
"""Bedrock Invoke chat path forwards ``output_config`` for adaptive Claude models."""
config = AmazonAnthropicClaudeConfig()
messages = [{"role": "user", "content": "test"}]
# Inject output_config into optional_params (simulates Anthropic SDK forwarding it)
optional_params = {
"max_tokens": 100,
"output_config": {"effort": "high"},
}
result = config.transform_request(
model="anthropic.claude-sonnet-4-20250514-v1:0",
model="anthropic.claude-opus-4-7",
messages=messages,
optional_params=optional_params,
litellm_params={},
headers={},
)
assert (
"output_config" not in result
), f"output_config should be stripped for Bedrock Chat Invoke, got keys: {list(result.keys())}"
# Verify normal params survive
assert result.get("output_config") == {"effort": "high"}
assert result["max_tokens"] == 100

View file

@ -310,6 +310,111 @@ def test_reasoning_effort_none_omits_thinking_for_anthropic_converse(model):
assert "thinking" not in optional_params
@pytest.mark.parametrize(
"model,effort,expected_effort",
[
("bedrock/converse/us.anthropic.claude-opus-4-7", "low", "low"),
("bedrock/converse/us.anthropic.claude-opus-4-7", "medium", "medium"),
("bedrock/converse/us.anthropic.claude-opus-4-7", "high", "high"),
("bedrock/converse/us.anthropic.claude-opus-4-7", "xhigh", "xhigh"),
("bedrock/converse/us.anthropic.claude-opus-4-7", "max", "max"),
("bedrock/converse/us.anthropic.claude-opus-4-6-v1", "max", "max"),
("bedrock/converse/us.anthropic.claude-sonnet-4-6", "high", "high"),
("bedrock/converse/us.anthropic.claude-sonnet-4-6", "minimal", "low"),
],
)
def test_reasoning_effort_sets_output_config_for_adaptive_models_converse(
model, effort, expected_effort
):
"""Adaptive Claude 4.6 / 4.7 on Bedrock Converse routes the tier via ``output_config.effort``."""
config = AmazonConverseConfig()
optional_params = config.map_openai_params(
non_default_params={"reasoning_effort": effort},
optional_params={},
model=model,
drop_params=False,
)
assert optional_params["thinking"]["type"] == "adaptive"
assert optional_params["output_config"] == {"effort": expected_effort}
@pytest.mark.parametrize(
"model",
[
"bedrock/converse/us.anthropic.claude-opus-4-7",
"bedrock/converse/us.anthropic.claude-opus-4-6-v1",
"bedrock/converse/us.anthropic.claude-sonnet-4-6",
],
)
def test_output_config_effort_forwarded_into_additional_request_fields(model):
"""``output_config`` rides along inside ``additionalModelRequestFields``."""
config = AmazonConverseConfig()
messages = [{"role": "user", "content": "hi"}]
result = config._transform_request(
model=model,
messages=messages,
optional_params={
"maxTokens": 256,
"thinking": {"type": "adaptive"},
"output_config": {"effort": "high"},
},
litellm_params={},
headers={},
)
additional = result.get("additionalModelRequestFields", {})
assert additional.get("output_config") == {"effort": "high"}
@pytest.mark.parametrize(
"effort",
["disabled", "invalid", ""],
)
def test_reasoning_effort_garbage_raises_bad_request_converse(effort):
"""Unmapped reasoning_effort on Bedrock Converse Anthropic raises BadRequestError."""
config = AmazonConverseConfig()
with pytest.raises(litellm.exceptions.BadRequestError):
config.map_openai_params(
non_default_params={"reasoning_effort": effort},
optional_params={},
model="bedrock/converse/us.anthropic.claude-opus-4-7",
drop_params=False,
)
@pytest.mark.parametrize(
"model",
[
"bedrock/converse/us.anthropic.claude-sonnet-4-6",
"bedrock/converse/global.anthropic.claude-sonnet-4-6",
"bedrock/converse/eu.anthropic.claude-sonnet-4-6",
"bedrock/converse/au.anthropic.claude-sonnet-4-6",
],
)
def test_output_config_effort_max_passes_through_on_sonnet_46_variants(model):
"""``effort='max'`` flows through for every Bedrock Converse Sonnet 4.6 id."""
config = AmazonConverseConfig()
messages = [{"role": "user", "content": "hi"}]
result = config._transform_request(
model=model,
messages=messages,
optional_params={
"maxTokens": 256,
"output_config": {"effort": "max"},
},
litellm_params={},
headers={},
)
additional = result.get("additionalModelRequestFields", {})
assert additional.get("output_config") == {"effort": "max"}
def test_get_supported_openai_params():
config = AmazonConverseConfig()
supported_params = config.get_supported_openai_params(
@ -3329,6 +3434,57 @@ def test_transform_request_strips_anthropic_output_config():
assert "output_config" not in additional_fields
def test_converse_drop_params_strips_output_config_for_pre_4_5_anthropic():
"""``drop_params=True`` strips unsupported ``output_config`` on Bedrock Converse."""
config = AmazonConverseConfig()
messages = [{"role": "user", "content": "hi"}]
original = litellm.drop_params
litellm.drop_params = True
try:
result = config._transform_request(
model="bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
messages=messages,
optional_params={
"maxTokens": 256,
"output_config": {"effort": "low"},
},
litellm_params={},
headers={},
)
finally:
litellm.drop_params = original
additional = result.get("additionalModelRequestFields", {})
assert "output_config" not in additional
def test_converse_drop_params_keeps_output_config_for_supporting_anthropic():
"""``drop_params=True`` does not strip on models that support ``output_config``."""
config = AmazonConverseConfig()
messages = [{"role": "user", "content": "hi"}]
original = litellm.drop_params
litellm.drop_params = True
try:
result = config._transform_request(
model="bedrock/converse/us.anthropic.claude-opus-4-7",
messages=messages,
optional_params={
"maxTokens": 256,
"thinking": {"type": "adaptive"},
"output_config": {"effort": "high"},
},
litellm_params={},
headers={},
)
finally:
litellm.drop_params = original
additional = result.get("additionalModelRequestFields", {})
assert additional.get("output_config") == {"effort": "high"}
def test_transform_response_native_structured_output():
"""Test response handling when model returns JSON as text content (native structured output)."""
response_json = {

View file

@ -0,0 +1,206 @@
import base64
import os
from types import MappingProxyType
from unittest.mock import MagicMock, patch
import pytest
import litellm.files.main as files_main
from litellm.llms.bedrock.files.handler import BedrockFilesHandler
from litellm.types.utils import SpecialEnums
def _encode_unified_file_id(s3_uri: str) -> str:
unified_file_id = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
"application/json",
"unified-id",
"",
s3_uri,
"model-id",
)
return base64.urlsafe_b64encode(unified_file_id.encode()).decode().rstrip("=")
class TestBedrockFilesHandler:
def setup_method(self):
self.handler = BedrockFilesHandler()
def test_should_parse_direct_managed_s3_uri(self):
bucket, key = self.handler._parse_s3_uri(
s3_uri="s3://safe-bucket/litellm-bedrock-files-model-id-abc.jsonl",
configured_bucket_name="safe-bucket",
)
assert bucket == "safe-bucket"
assert key == "litellm-bedrock-files-model-id-abc.jsonl"
def test_should_parse_managed_batch_output_uri(self):
bucket, key = self.handler._parse_s3_uri(
s3_uri="s3://safe-bucket/litellm-batch-outputs/job/",
configured_bucket_name="safe-bucket",
)
assert bucket == "safe-bucket"
assert key == "litellm-batch-outputs/job/"
def test_should_reject_arbitrary_bucket(self):
with pytest.raises(ValueError, match="configured storage bucket"):
self.handler._parse_s3_uri(
s3_uri="s3://other-bucket/litellm-bedrock-files-model-id-abc.jsonl",
configured_bucket_name="safe-bucket",
)
def test_should_reject_unmanaged_same_bucket_key(self):
with pytest.raises(ValueError, match="LiteLLM-managed"):
self.handler._parse_s3_uri(
s3_uri="s3://safe-bucket/private/output.jsonl",
configured_bucket_name="safe-bucket",
)
def test_should_allow_legacy_same_bucket_key_when_server_flag_enabled(self):
bucket, key = self.handler._parse_s3_uri(
s3_uri="s3://safe-bucket/private/output.jsonl",
configured_bucket_name="safe-bucket",
allow_legacy_cloud_file_ids=True,
)
assert bucket == "safe-bucket"
assert key == "private/output.jsonl"
def test_should_keep_configured_prefix_for_legacy_keys(self):
bucket, key = self.handler._parse_s3_uri(
s3_uri="s3://safe-bucket/team-a/private/output.jsonl",
configured_bucket_name="safe-bucket/team-a",
allow_legacy_cloud_file_ids=True,
)
assert bucket == "safe-bucket"
assert key == "team-a/private/output.jsonl"
def test_should_reject_legacy_key_outside_configured_prefix(self):
with pytest.raises(ValueError, match="configured storage prefix"):
self.handler._parse_s3_uri(
s3_uri="s3://safe-bucket/team-b/private/output.jsonl",
configured_bucket_name="safe-bucket/team-a",
allow_legacy_cloud_file_ids=True,
)
def test_should_reject_dot_segment_key(self):
with pytest.raises(ValueError, match="invalid path segment"):
self.handler._parse_s3_uri(
s3_uri="s3://safe-bucket/litellm-bedrock-files/../secret.jsonl",
configured_bucket_name="safe-bucket",
)
def test_should_reject_empty_middle_path_segment(self):
with pytest.raises(ValueError, match="invalid path segment"):
self.handler._parse_s3_uri(
s3_uri="s3://safe-bucket/litellm-bedrock-files//secret.jsonl",
configured_bucket_name="safe-bucket",
)
def test_should_extract_unified_managed_s3_uri(self):
file_id = _encode_unified_file_id(
"s3://safe-bucket/litellm-batch-outputs/job/output.jsonl"
)
assert (
self.handler._extract_s3_uri_from_file_id(file_id)
== "s3://safe-bucket/litellm-batch-outputs/job/output.jsonl"
)
def test_should_reject_file_id_without_s3_scheme(self):
with pytest.raises(ValueError, match="managed LiteLLM S3 file id"):
self.handler._extract_s3_uri_from_file_id("safe-bucket/private.jsonl")
def test_should_reject_unified_unmanaged_s3_uri(self):
file_id = _encode_unified_file_id("s3://safe-bucket/private/output.jsonl")
s3_uri = self.handler._extract_s3_uri_from_file_id(file_id)
with pytest.raises(ValueError, match="LiteLLM-managed"):
self.handler._parse_s3_uri(
s3_uri=s3_uri,
configured_bucket_name="safe-bucket",
)
def test_should_not_trust_request_s3_bucket_name_for_expected_bucket(self):
with patch.dict(os.environ, {"AWS_S3_BUCKET_NAME": "safe-bucket"}):
assert (
self.handler._get_configured_s3_bucket_name(
{"s3_bucket_name": "attacker-bucket"}
)
== "safe-bucket"
)
def test_should_trust_proxy_config_s3_bucket_name_for_expected_bucket(self):
trusted_credentials = MappingProxyType({"s3_bucket_name": "safe-bucket"})
with patch.dict(os.environ, {}, clear=True):
assert (
self.handler._get_configured_s3_bucket_name(
{
"s3_bucket_name": "attacker-bucket",
"_litellm_internal_model_credentials": trusted_credentials,
}
)
== "safe-bucket"
)
def test_should_not_trust_user_supplied_internal_credentials_dict(self):
with patch.dict(os.environ, {}, clear=True):
with pytest.raises(ValueError, match="S3 bucket_name is required"):
self.handler._get_configured_s3_bucket_name(
{
"_litellm_internal_model_credentials": {
"s3_bucket_name": "attacker-bucket"
}
}
)
def test_should_require_server_s3_bucket_name(self):
with patch.dict(os.environ, {}, clear=True):
with pytest.raises(ValueError, match="S3 bucket_name is required"):
self.handler._get_configured_s3_bucket_name(
{"s3_bucket_name": "attacker-bucket"}
)
def test_should_forward_trusted_model_credentials_to_bedrock_provider_config():
trusted_credentials = MappingProxyType({"s3_bucket_name": "safe-bucket"})
mock_response = MagicMock()
with patch.object(
files_main.base_llm_http_handler,
"retrieve_file_content",
return_value=mock_response,
) as mock_retrieve_file_content:
response = files_main.file_content(
file_id="s3://safe-bucket/litellm-bedrock-files/file.jsonl",
custom_llm_provider="bedrock",
_litellm_internal_model_credentials=trusted_credentials,
)
assert response is mock_response
litellm_params = mock_retrieve_file_content.call_args.kwargs["litellm_params"]
assert litellm_params["_litellm_internal_model_credentials"] is trusted_credentials
assert "s3_bucket_name" not in litellm_params
def test_should_forward_trusted_model_credentials_to_retrieve_provider_config():
trusted_credentials = MappingProxyType({"allow_legacy_cloud_file_ids": True})
mock_response = MagicMock()
with patch.object(
files_main.base_llm_http_handler,
"retrieve_file",
return_value=mock_response,
) as mock_retrieve_file:
response = files_main.file_retrieve(
file_id="gs://safe-bucket/private/file.jsonl",
custom_llm_provider="vertex_ai",
_litellm_internal_model_credentials=trusted_credentials,
)
assert response is mock_response
litellm_params = mock_retrieve_file.call_args.kwargs["litellm_params"]
assert litellm_params["_litellm_internal_model_credentials"] is trusted_credentials

View file

@ -4,9 +4,7 @@ Test bedrock files transformation functionality
import json
import os
from typing import Any, Dict, List
import pytest
from urllib.parse import unquote, urlparse
from litellm.llms.bedrock.files.transformation import BedrockJsonlFilesTransformation
@ -43,19 +41,6 @@ class TestBedrockFilesTransformation:
)
)
# Print the transformation results for validation
print("\n=== INPUT (OpenAI format) ===")
for i, content in enumerate(openai_jsonl_content):
print(f"Record {i+1}:")
print(json.dumps(content, indent=2))
print()
print("\n=== OUTPUT (Bedrock format) ===")
for i, content in enumerate(bedrock_jsonl_content):
print(f"Record {i+1}:")
print(json.dumps(content, indent=2))
print()
# Basic validation
assert len(bedrock_jsonl_content) == len(
openai_jsonl_content
@ -88,17 +73,6 @@ class TestBedrockFilesTransformation:
"max_tokens" in model_input
), f"Record {i+1} should have max_tokens"
# Write expected output to file for reference
expected_output_path = os.path.join(
os.path.dirname(__file__), "expected_bedrock_batch_completions.jsonl"
)
with open(expected_output_path, "w") as f:
for record in bedrock_jsonl_content:
f.write(json.dumps(record) + "\n")
print(f"\n=== Expected output written to: {expected_output_path} ===")
def test_nova_text_only_uses_converse_format(self):
"""
Test that Nova models produce Converse API format in batch modelInput.
@ -327,6 +301,44 @@ class TestBedrockFilesTransformation:
), f"us-west-2 must not appear when s3_region_name is set, got: {url}"
assert "litellm-batch-352026" in url
def test_get_complete_file_url_sanitizes_untrusted_filename(self):
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
create_file_data = {
"file": ("../../owned.jsonl?acl=public", b"hello", "application/jsonl"),
"purpose": "assistants",
}
url = config.get_complete_file_url(
api_base=None,
api_key=None,
model="amazon.nova-pro-v1:0",
optional_params={"aws_region_name": "us-west-2"},
litellm_params={"s3_bucket_name": "safe-bucket"},
data=create_file_data,
)
parsed_url = urlparse(url)
object_key = unquote(parsed_url.path).split("/safe-bucket/", 1)[1]
assert object_key.startswith("litellm-bedrock-files/")
assert object_key.endswith("-owned.jsonl_acl_public")
assert ".." not in object_key
assert parsed_url.query == ""
def test_batch_object_name_sanitizes_model_path(self):
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
object_name = config._get_s3_object_name_from_batch_jsonl(
[{"body": {"model": "bedrock/../../secret:model"}}]
)
assert object_name.startswith("litellm-bedrock-files-")
assert object_name.endswith(".jsonl")
assert "/" not in object_name
assert ".." not in object_name
def test_transform_create_file_request_injects_s3_region_for_signing(self):
"""
When s3_region_name is provided, transform_create_file_request must pass

View file

@ -592,13 +592,8 @@ def test_remove_scope_from_cache_control():
assert request["messages"][0]["content"][0]["cache_control"]["type"] == "ephemeral"
def test_bedrock_messages_strips_output_config():
"""
Ensure output_config is stripped from the request before sending to
Bedrock Invoke, which doesn't support this Anthropic-specific parameter.
Regression test for: https://github.com/BerriAI/litellm/issues/22797
"""
def test_bedrock_messages_forwards_output_config():
"""Bedrock Invoke /v1/messages forwards ``output_config`` for adaptive Claude models."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
@ -611,26 +606,20 @@ def test_bedrock_messages_strips_output_config():
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-3-haiku-20240307-v1:0",
model="anthropic.claude-opus-4-7",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert (
"output_config" not in result
), "output_config should be stripped — Bedrock Invoke rejects it"
assert result.get("output_config") == {"effort": "high"}
# Other params should be preserved
assert result.get("max_tokens") == 4096
def test_bedrock_messages_strips_output_config_with_output_format():
"""
When both output_config and output_format are present, both should be
stripped (output_format is converted to inline schema, output_config
is simply dropped).
"""
def test_bedrock_messages_forwards_output_config_with_output_format():
"""``output_config`` is forwarded; ``output_format`` is converted to inline schema."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
@ -647,6 +636,29 @@ def test_bedrock_messages_strips_output_config_with_output_format():
},
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-opus-4-7",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert result.get("output_config") == {"effort": "low"}
assert "output_format" not in result
def test_bedrock_messages_forwards_output_config_for_non_adaptive_model():
"""``output_config`` is forwarded for non-adaptive models so the provider's error surfaces."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
optional_params = {
"max_tokens": 4096,
"output_config": {"effort": "high"},
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-3-haiku-20240307-v1:0",
messages=messages,
@ -655,8 +667,201 @@ def test_bedrock_messages_strips_output_config_with_output_format():
headers={},
)
assert result.get("output_config") == {"effort": "high"}
assert result.get("max_tokens") == 4096
def test_bedrock_messages_drop_params_strips_output_config_for_pre_4_5():
"""``drop_params=True`` strips ``output_config`` for pre-4.5 Anthropic on /v1/messages."""
import litellm
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
optional_params = {
"max_tokens": 4096,
"output_config": {"effort": "low"},
}
original = litellm.drop_params
litellm.drop_params = True
try:
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-3-haiku-20240307-v1:0",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
finally:
litellm.drop_params = original
assert "output_config" not in result
assert "output_format" not in result
def test_bedrock_messages_drop_params_keeps_output_config_for_4_7():
"""``drop_params=True`` does not strip on opus-4-7 (supports effort)."""
import litellm
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
optional_params = {
"max_tokens": 4096,
"output_config": {"effort": "high"},
}
original = litellm.drop_params
litellm.drop_params = True
try:
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-opus-4-7",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
finally:
litellm.drop_params = original
assert result.get("output_config") == {"effort": "high"}
@pytest.mark.parametrize(
"reasoning_effort,expected_effort",
[
("minimal", "low"),
("low", "low"),
("medium", "medium"),
("high", "high"),
("xhigh", "xhigh"),
("max", "max"),
],
)
def test_bedrock_messages_maps_reasoning_effort_for_adaptive_model(
reasoning_effort, expected_effort
):
"""``reasoning_effort`` maps to ``thinking`` + ``output_config.effort`` on /v1/messages."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
optional_params = {
"max_tokens": 4096,
"reasoning_effort": reasoning_effort,
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-opus-4-7",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert "reasoning_effort" not in result
assert result.get("thinking") == {"type": "adaptive"}
assert result.get("output_config") == {"effort": expected_effort}
def test_bedrock_messages_reasoning_effort_on_non_adaptive_uses_thinking_budget():
"""Non-adaptive models map ``reasoning_effort`` to ``thinking.budget_tokens``."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
optional_params = {
"max_tokens": 4096,
"reasoning_effort": "medium",
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-opus-4-5-20251101-v1:0",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert "reasoning_effort" not in result
assert "output_config" not in result
thinking = result.get("thinking")
assert isinstance(thinking, dict)
assert thinking.get("type") == "enabled"
assert isinstance(thinking.get("budget_tokens"), int)
assert thinking["budget_tokens"] >= 1024
def test_bedrock_messages_reasoning_effort_none_clears_thinking():
"""``reasoning_effort='none'`` clears both ``thinking`` and ``output_config``."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
optional_params = {
"max_tokens": 4096,
"reasoning_effort": "none",
"output_config": {"effort": "high"},
"thinking": {"type": "adaptive"},
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-opus-4-7",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert "reasoning_effort" not in result
assert "thinking" not in result
assert "output_config" not in result
def test_bedrock_messages_invalid_reasoning_effort_raises_400():
"""Garbage ``reasoning_effort`` raises AnthropicError (400)."""
from litellm.llms.anthropic.common_utils import AnthropicError
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
for bad_effort in ("invalid", "disabled", ""):
with pytest.raises(AnthropicError):
cfg.transform_anthropic_messages_request(
model="anthropic.claude-opus-4-7",
messages=messages,
anthropic_messages_optional_request_params={
"max_tokens": 4096,
"reasoning_effort": bad_effort,
},
litellm_params=GenericLiteLLMParams(),
headers={},
)
def test_bedrock_messages_explicit_output_config_wins_over_reasoning_effort():
"""Explicit ``output_config.effort`` wins over the ``reasoning_effort`` alias."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
optional_params = {
"max_tokens": 4096,
"reasoning_effort": "low",
"output_config": {"effort": "max"},
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-opus-4-7",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert "reasoning_effort" not in result
assert result.get("output_config") == {"effort": "max"}
def test_bedrock_messages_strips_context_management():
@ -728,13 +933,13 @@ def test_bedrock_messages_allowlist_filters_anthropic_only_fields():
"mcp_servers",
"container",
"inference_geo",
"output_config",
"context_management",
"model",
"stream",
):
assert bad not in result, f"{bad} should be stripped by the allowlist"
assert result.get("output_config") == {"effort": "low"}
# Supported fields pass through.
assert result["max_tokens"] == 4096
assert result["temperature"] == 0.5
@ -882,7 +1087,10 @@ async def test_promote_message_start_cache_when_message_stop_omits_cache_fields(
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"input_tokens": 10, "output_tokens": 181},
}
yield {"type": "message_stop", "usage": {"input_tokens": 10, "output_tokens": 181}}
yield {
"type": "message_stop",
"usage": {"input_tokens": 10, "output_tokens": 181},
}
merged: list[dict] = []
async for chunk in cfg._promote_message_stop_usage(_stream()):
@ -936,7 +1144,11 @@ async def test_unified_bedrock_messages_cache_on_start_only_never_negative_cost(
},
},
}
yield {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}
yield {
"type": "content_block_start",
"index": 0,
"content_block": {"type": "text", "text": ""},
}
yield {
"type": "content_block_delta",
"index": 0,
@ -948,7 +1160,10 @@ async def test_unified_bedrock_messages_cache_on_start_only_never_negative_cost(
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 181, "input_tokens": 10},
}
yield {"type": "message_stop", "usage": {"input_tokens": 10, "output_tokens": 181}}
yield {
"type": "message_stop",
"usage": {"input_tokens": 10, "output_tokens": 181},
}
logging_obj = LiteLLMLoggingObj(
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",

View file

@ -3,6 +3,7 @@ Test Vertex AI files handler functionality
"""
import asyncio
from types import MappingProxyType
import pytest
from unittest.mock import AsyncMock, patch
@ -12,6 +13,14 @@ from litellm.llms.vertex_ai.files.handler import VertexAIFilesHandler
from litellm.types.llms.openai import FileContentRequest, HttpxBinaryResponseContent
def _mock_gcs_logging_config(bucket_name: str = "test-bucket"):
return {
"bucket_name": bucket_name,
"path_service_account": None,
"vertex_instance": None,
}
class TestVertexAIFilesHandler:
"""Test Vertex AI files handler"""
@ -22,57 +31,84 @@ class TestVertexAIFilesHandler:
def test_extract_bucket_and_object_from_file_id_standard_path(self):
"""Test extraction of bucket and object from URL-encoded file_id with standard path"""
# Sample file_id with nested folder structure
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-folder" "%2Fsub-folder%2Ftest-file.txt"
file_id = (
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
"%2Ftest-folder%2Fsub-folder%2Ftest-file.txt"
)
bucket_name, encoded_object_path = (
self.handler._extract_bucket_and_object_from_file_id(file_id)
bucket_name, object_path = self.handler._extract_bucket_and_object_from_file_id(
file_id=file_id,
configured_bucket_name="test-bucket",
)
# Verify bucket name extraction
assert bucket_name == "test-bucket"
# Verify object path encoding
expected_encoded_object = "test-folder%2Fsub-folder%2Ftest-file.txt"
assert encoded_object_path == expected_encoded_object
expected_object = "litellm-vertex-files/test-folder/sub-folder/test-file.txt"
assert object_path == expected_object
def test_extract_bucket_and_object_from_file_id_bucket_only(self):
def test_extract_bucket_and_object_from_file_id_rejects_bucket_only(self):
"""Test extraction when only bucket name is provided"""
file_id = "gs%3A%2F%2Ftest-bucket"
bucket_name, encoded_object_path = (
self.handler._extract_bucket_and_object_from_file_id(file_id)
)
with pytest.raises(ValueError, match="object name"):
self.handler._extract_bucket_and_object_from_file_id(
file_id=file_id,
configured_bucket_name="test-bucket",
)
assert bucket_name == "test-bucket"
assert encoded_object_path == ""
def test_extract_bucket_and_object_from_file_id_simple_path(self):
def test_extract_bucket_and_object_from_file_id_rejects_unmanaged_path(self):
"""Test extraction with simple path"""
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
bucket_name, encoded_object_path = (
self.handler._extract_bucket_and_object_from_file_id(file_id)
with pytest.raises(ValueError, match="LiteLLM-managed"):
self.handler._extract_bucket_and_object_from_file_id(
file_id=file_id,
configured_bucket_name="test-bucket",
)
def test_extract_bucket_and_object_from_file_id_allows_trusted_legacy_flag(self):
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
trusted_credentials = MappingProxyType({"allow_legacy_cloud_file_ids": True})
bucket_name, object_path = self.handler._extract_bucket_and_object_from_file_id(
file_id=file_id,
configured_bucket_name="test-bucket",
litellm_params={
"_litellm_internal_model_credentials": trusted_credentials,
},
)
assert bucket_name == "test-bucket"
assert encoded_object_path == "test-file.txt"
assert object_path == "test-file.txt"
def test_extract_bucket_and_object_from_file_id_no_gs_prefix(self):
def test_extract_bucket_and_object_from_file_id_rejects_no_gs_prefix(self):
"""Test extraction when gs:// prefix is missing"""
file_id = "test-bucket%2Ftest-file.txt"
file_id = "test-bucket%2Flitellm-vertex-files%2Ftest-file.txt"
bucket_name, encoded_object_path = (
self.handler._extract_bucket_and_object_from_file_id(file_id)
)
with pytest.raises(ValueError, match="gs://"):
self.handler._extract_bucket_and_object_from_file_id(
file_id=file_id,
configured_bucket_name="test-bucket",
)
assert bucket_name == "test-bucket"
assert encoded_object_path == "test-file.txt"
def test_extract_bucket_and_object_from_file_id_rejects_wrong_bucket(self):
file_id = "gs%3A%2F%2Fother-bucket%2Flitellm-vertex-files%2Ftest-file.txt"
with pytest.raises(ValueError, match="configured storage bucket"):
self.handler._extract_bucket_and_object_from_file_id(
file_id=file_id,
configured_bucket_name="test-bucket",
)
@pytest.mark.asyncio
async def test_afile_content_success(self):
"""Test successful async file content retrieval"""
# Setup test data
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
file_id = (
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
"%2Fuploads%2Fabc-test-file.txt"
)
expected_content = b"test file content"
file_content_request = FileContentRequest(
@ -80,9 +116,17 @@ class TestVertexAIFilesHandler:
)
# Mock the download_gcs_object method
with patch.object(
self.handler, "download_gcs_object", new_callable=AsyncMock
) as mock_download:
with (
patch.object(
self.handler, "download_gcs_object", new_callable=AsyncMock
) as mock_download,
patch.object(
self.handler,
"get_gcs_logging_config",
new_callable=AsyncMock,
return_value=_mock_gcs_logging_config(),
),
):
mock_download.return_value = expected_content
# Call the method
@ -104,7 +148,10 @@ class TestVertexAIFilesHandler:
# Verify the download was called with correct parameters
mock_download.assert_called_once()
call_args = mock_download.call_args
assert call_args.kwargs["object_name"] == "test-file.txt"
assert (
call_args.kwargs["object_name"]
== "litellm-vertex-files/uploads/abc-test-file.txt"
)
assert "standard_callback_dynamic_params" in call_args.kwargs
assert (
call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"]
@ -132,22 +179,33 @@ class TestVertexAIFilesHandler:
@pytest.mark.asyncio
async def test_afile_content_download_failure(self):
"""Test async file content retrieval when download fails"""
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
file_id = (
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
"%2Fuploads%2Fabc-test-file.txt"
)
file_content_request = FileContentRequest(
file_id=file_id, extra_headers=None, extra_body=None
)
# Mock download to return None (failure)
with patch.object(
self.handler, "download_gcs_object", new_callable=AsyncMock
) as mock_download:
with (
patch.object(
self.handler, "download_gcs_object", new_callable=AsyncMock
) as mock_download,
patch.object(
self.handler,
"get_gcs_logging_config",
new_callable=AsyncMock,
return_value=_mock_gcs_logging_config(),
),
):
mock_download.return_value = None
# Should raise ValueError for failed download
with pytest.raises(
ValueError,
match="Failed to download file from GCS: gs://test-bucket/test-file.txt",
match="Failed to download file from GCS: gs://test-bucket/litellm-vertex-files/uploads/abc-test-file.txt",
):
await self.handler.afile_content(
file_content_request=file_content_request,

View file

@ -5,6 +5,8 @@ Includes tests for Vertex AI batch output transformation to OpenAI format.
import json
import urllib.parse
from types import MappingProxyType
from urllib.parse import parse_qs, urlparse
import httpx
import pytest
@ -29,13 +31,20 @@ class TestParseGcsUri:
"""Tests for the _parse_gcs_uri helper used by retrieve / content / delete."""
def test_should_parse_standard_gs_uri(self, config):
bucket, encoded = config._parse_gcs_uri("gs://my-bucket/path/to/object.jsonl")
file_id = "gs://my-bucket/litellm-vertex-files/path/to/object.jsonl"
bucket, encoded = config._parse_gcs_uri(
file_id, litellm_params={"bucket_name": "my-bucket"}
)
assert bucket == "my-bucket"
assert encoded == urllib.parse.quote("path/to/object.jsonl", safe="")
assert encoded == urllib.parse.quote(
"litellm-vertex-files/path/to/object.jsonl", safe=""
)
def test_should_parse_uri_with_nested_publisher_path(self, config):
uri = "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc-123"
bucket, encoded = config._parse_gcs_uri(uri)
bucket, encoded = config._parse_gcs_uri(
uri, litellm_params={"bucket_name": "litellm-local"}
)
assert bucket == "litellm-local"
expected_path = (
"litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc-123"
@ -43,30 +52,141 @@ class TestParseGcsUri:
assert encoded == urllib.parse.quote(expected_path, safe="")
def test_should_handle_url_encoded_input(self, config):
encoded_uri = urllib.parse.quote("gs://my-bucket/some/path", safe="")
bucket, encoded = config._parse_gcs_uri(encoded_uri)
encoded_uri = urllib.parse.quote(
"gs://my-bucket/litellm-vertex-files/some/path", safe=""
)
bucket, encoded = config._parse_gcs_uri(
encoded_uri, litellm_params={"bucket_name": "my-bucket"}
)
assert bucket == "my-bucket"
assert encoded == urllib.parse.quote("some/path", safe="")
assert encoded == urllib.parse.quote("litellm-vertex-files/some/path", safe="")
def test_should_handle_bucket_only(self, config):
bucket, encoded = config._parse_gcs_uri("gs://my-bucket")
assert bucket == "my-bucket"
assert encoded == ""
def test_should_reject_bucket_only(self, config):
with pytest.raises(ValueError, match="object name"):
config._parse_gcs_uri(
"gs://my-bucket", litellm_params={"bucket_name": "my-bucket"}
)
def test_should_reject_no_gs_prefix(self, config):
with pytest.raises(ValueError, match="gs://"):
config._parse_gcs_uri(
"my-bucket/litellm-vertex-files/object.txt",
litellm_params={"bucket_name": "my-bucket"},
)
def test_should_reject_unmanaged_object_path(self, config):
with pytest.raises(ValueError, match="LiteLLM-managed"):
config._parse_gcs_uri(
"gs://my-bucket/private/object.txt",
litellm_params={"bucket_name": "my-bucket"},
)
def test_should_reject_request_supplied_legacy_flag(self, config):
with pytest.raises(ValueError, match="LiteLLM-managed"):
config._parse_gcs_uri(
"gs://my-bucket/private/object.txt",
litellm_params={
"bucket_name": "my-bucket",
"allow_legacy_cloud_file_ids": True,
},
)
def test_should_allow_legacy_object_path_with_trusted_server_flag(self, config):
trusted_credentials = MappingProxyType({"allow_legacy_cloud_file_ids": True})
bucket, encoded = config._parse_gcs_uri(
"gs://my-bucket/private/object.txt",
litellm_params={
"bucket_name": "my-bucket",
"_litellm_internal_model_credentials": trusted_credentials,
},
)
def test_should_handle_no_gs_prefix(self, config):
bucket, encoded = config._parse_gcs_uri("my-bucket/object.txt")
assert bucket == "my-bucket"
assert encoded == "object.txt"
assert encoded == urllib.parse.quote("private/object.txt", safe="")
def test_should_reject_user_supplied_legacy_flag_snapshot(self, config):
with pytest.raises(ValueError, match="LiteLLM-managed"):
config._parse_gcs_uri(
"gs://my-bucket/private/object.txt",
litellm_params={
"bucket_name": "my-bucket",
"_litellm_internal_model_credentials": {
"allow_legacy_cloud_file_ids": True
},
},
)
def test_should_keep_configured_prefix_for_legacy_object_path(self, config):
trusted_credentials = MappingProxyType({"allow_legacy_cloud_file_ids": True})
bucket, encoded = config._parse_gcs_uri(
"gs://my-bucket/team-a/private/object.txt",
litellm_params={
"bucket_name": "my-bucket/team-a",
"_litellm_internal_model_credentials": trusted_credentials,
},
)
assert bucket == "my-bucket"
assert encoded == urllib.parse.quote("team-a/private/object.txt", safe="")
def test_should_reject_legacy_object_outside_configured_prefix(self, config):
trusted_credentials = MappingProxyType({"allow_legacy_cloud_file_ids": True})
with pytest.raises(ValueError, match="configured storage prefix"):
config._parse_gcs_uri(
"gs://my-bucket/team-b/private/object.txt",
litellm_params={
"bucket_name": "my-bucket/team-a",
"_litellm_internal_model_credentials": trusted_credentials,
},
)
def test_should_reject_unconfigured_bucket(self, config):
with pytest.raises(ValueError, match="configured storage bucket"):
config._parse_gcs_uri(
"gs://other-bucket/litellm-vertex-files/object.txt",
litellm_params={"bucket_name": "my-bucket"},
)
class TestCreateFileUrl:
def test_should_ignore_request_metadata_bucket_and_sanitize_filename(self, config):
url = config.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params={
"bucket_name": "safe-bucket",
"litellm_metadata": {"gcs_bucket_name": "attacker-bucket"},
},
data={
"file": ("../../owned.jsonl?alt=media", b"{}", "application/jsonl"),
"purpose": "assistants",
},
)
parsed_url = urlparse(url)
object_name = parse_qs(parsed_url.query)["name"][0]
assert "/b/safe-bucket/" in parsed_url.path
assert "attacker-bucket" not in url
assert object_name.startswith("litellm-vertex-files/uploads/")
assert object_name.endswith("-owned.jsonl_alt_media")
assert ".." not in object_name
assert "?" not in object_name
class TestTransformRetrieveFile:
def test_should_build_correct_gcs_metadata_url(self, config):
file_id = "gs://my-bucket/path/to/file.jsonl"
file_id = "gs://my-bucket/litellm-vertex-files/path/to/file.jsonl"
url, params = config.transform_retrieve_file_request(
file_id=file_id, optional_params={}, litellm_params={}
file_id=file_id,
optional_params={},
litellm_params={"bucket_name": "my-bucket"},
)
expected_encoded = urllib.parse.quote(
"litellm-vertex-files/path/to/file.jsonl", safe=""
)
expected_encoded = urllib.parse.quote("path/to/file.jsonl", safe="")
assert (
url
== f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{expected_encoded}"
@ -119,13 +239,13 @@ class TestTransformRetrieveFile:
class TestTransformFileContent:
def test_should_build_gcs_media_download_url(self, config):
file_id = "gs://my-bucket/path/to/file.jsonl"
file_id = "gs://my-bucket/litellm-vertex-files/path/to/file.jsonl"
url, params = config.transform_file_content_request(
file_content_request={"file_id": file_id},
optional_params={},
litellm_params={},
litellm_params={"bucket_name": "my-bucket"},
)
encoded = urllib.parse.quote("path/to/file.jsonl", safe="")
encoded = urllib.parse.quote("litellm-vertex-files/path/to/file.jsonl", safe="")
assert (
url
== f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded}?alt=media"
@ -254,11 +374,13 @@ class TestTransformFileContent:
class TestTransformDeleteFile:
def test_should_build_correct_gcs_delete_url(self, config):
file_id = "gs://my-bucket/path/to/file.jsonl"
file_id = "gs://my-bucket/litellm-vertex-files/path/to/file.jsonl"
url, params = config.transform_delete_file_request(
file_id=file_id, optional_params={}, litellm_params={}
file_id=file_id,
optional_params={},
litellm_params={"bucket_name": "my-bucket"},
)
encoded = urllib.parse.quote("path/to/file.jsonl", safe="")
encoded = urllib.parse.quote("litellm-vertex-files/path/to/file.jsonl", safe="")
assert (
url == f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded}"
)

View file

@ -498,12 +498,8 @@ def test_vertex_ai_partner_models_anthropic_remove_prompt_caching_scope_beta_hea
), "Header should be removed if no supported values remain"
def test_vertex_ai_anthropic_output_config_effort_only_dropped():
"""
``output_config`` containing only ``effort`` (an Anthropic-only key Vertex
rejects with "Extra inputs are not permitted") is dropped entirely so the
request body has no empty dict.
"""
def test_vertex_ai_anthropic_output_config_effort_only_forwarded():
"""Vertex AI Claude 4.6/4.7 accept ``output_config.effort`` on rawPredict."""
config = VertexAIAnthropicConfig()
messages = [{"role": "user", "content": "What is 2+2?"}]
@ -515,16 +511,14 @@ def test_vertex_ai_anthropic_output_config_effort_only_dropped():
}
result = config.transform_request(
model="claude-3-5-sonnet-20241022",
model="claude-opus-4-6",
messages=messages,
optional_params=optional_params,
litellm_params={},
headers=headers,
)
assert (
"output_config" not in result
), "output_config containing only effort must be dropped"
assert result["output_config"] == {"effort": "high"}
assert result["max_tokens"] == 1024
assert "messages" in result
@ -566,14 +560,8 @@ def test_vertex_ai_anthropic_output_config_format_passes_through():
assert result["output_config"] == output_config
def test_vertex_ai_anthropic_output_config_format_plus_effort_strips_only_effort():
"""
Greptile P1 on PR #23396: when ``output_config`` contains BOTH ``format``
and ``effort``, the prior conditional-passthrough logic forwarded the
full dict including the unsupported ``effort`` key, reproducing the
400 error the fix was meant to resolve. Only ``effort`` (and any future
Vertex-unsupported keys) should be filtered; ``format`` must survive.
"""
def test_vertex_ai_anthropic_output_config_format_plus_effort_preserved():
"""Both ``format`` and ``effort`` ride along on Vertex Claude 4.6/4.7."""
config = VertexAIAnthropicConfig()
messages = [{"role": "user", "content": "Return a person object."}]
@ -591,7 +579,7 @@ def test_vertex_ai_anthropic_output_config_format_plus_effort_strips_only_effort
optional_params = {"max_tokens": 1024, "output_config": output_config}
result = config.transform_request(
model="claude-3-5-sonnet-20241022",
model="claude-opus-4-6",
messages=messages,
optional_params=optional_params,
litellm_params={},
@ -599,9 +587,7 @@ def test_vertex_ai_anthropic_output_config_format_plus_effort_strips_only_effort
)
assert "output_config" in result
assert (
"effort" not in result["output_config"]
), "effort must be stripped — Vertex returns 400 on unknown keys"
assert result["output_config"]["effort"] == "high"
assert result["output_config"]["format"] == output_config["format"]
@ -623,15 +609,8 @@ def test_vertex_ai_anthropic_output_config_non_dict_dropped():
assert "output_config" not in result
def test_vertex_ai_anthropic_output_format_preserved_output_config_effort_dropped():
"""
When the request carries both ``output_format`` (top-level structured
outputs) AND an ``output_config`` whose only useful key for Vertex is
``effort``: ``output_format`` must be forwarded (Vertex accepts it),
while ``output_config`` is dropped because Vertex returns 400 on
``effort``. This replaces the old "drop both" behavior, which was the
silent strip the bug report flagged.
"""
def test_vertex_ai_anthropic_output_format_and_output_config_effort_preserved():
"""Both ``output_format`` and ``output_config.effort`` are forwarded on Vertex 4.6/4.7."""
config = VertexAIAnthropicConfig()
messages = [{"role": "user", "content": "Extract structured data"}]
@ -653,7 +632,7 @@ def test_vertex_ai_anthropic_output_format_preserved_output_config_effort_droppe
}
test_data = {
"model": "claude-3-5-sonnet-20241022",
"model": "claude-opus-4-6",
"messages": messages,
"max_tokens": 2048,
"output_format": output_format,
@ -671,7 +650,7 @@ def test_vertex_ai_anthropic_output_format_preserved_output_config_effort_droppe
try:
result = config.transform_request(
model="claude-3-5-sonnet-20241022",
model="claude-opus-4-6",
messages=messages,
optional_params=optional_params,
litellm_params={},
@ -680,9 +659,8 @@ def test_vertex_ai_anthropic_output_format_preserved_output_config_effort_droppe
# output_format flows through unchanged — Vertex AI Claude accepts it.
assert result["output_format"] == output_format
# output_config containing only ``effort`` is dropped to avoid the
# 400 "Extra inputs are not permitted" the silent strip used to mask.
assert "output_config" not in result
# output_config.effort now flows through (Vertex accepts it on 4.6/4.7).
assert result["output_config"] == {"effort": "high"}
assert result["max_tokens"] == 2048
assert "model" not in result, "model is still stripped (Vertex routes by URL)"
finally:
@ -702,10 +680,10 @@ def test_sanitize_vertex_anthropic_output_params_unit():
sanitize_vertex_anthropic_output_params(data)
assert data == {"max_tokens": 8}
# Effort-only → dropped entirely.
# Effort-only → preserved (Vertex 4.6/4.7 accept it on rawPredict).
data = {"output_config": {"effort": "high"}}
sanitize_vertex_anthropic_output_params(data)
assert "output_config" not in data
assert data["output_config"] == {"effort": "high"}
# Format-only → preserved unchanged.
fmt = {"format": {"type": "json_schema", "schema": {"type": "object"}}}
@ -713,10 +691,10 @@ def test_sanitize_vertex_anthropic_output_params_unit():
sanitize_vertex_anthropic_output_params(data)
assert data["output_config"] == fmt
# Mixed → effort filtered, format kept.
# Mixed → both effort and format kept (no current Vertex-unsupported keys).
data = {"output_config": {"format": fmt["format"], "effort": "high"}}
sanitize_vertex_anthropic_output_params(data)
assert data["output_config"] == fmt
assert data["output_config"] == {"format": fmt["format"], "effort": "high"}
# Non-dict → dropped defensively.
data = {"output_config": "garbage"}

View file

@ -6,6 +6,90 @@ sys.path.insert(
) # Adds the parent directory to the system path
from litellm.llms.xai.chat.transformation import XAIChatConfig
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
ModelResponse,
Usage,
)
class TestXAIReasoningTokenFolding:
"""``_fold_reasoning_tokens_into_completion`` re-aligns xAI Usage to the OpenAI invariant."""
@staticmethod
def _make_response(
prompt_tokens: int,
completion_tokens: int,
total_tokens: int,
reasoning_tokens: int = 0,
) -> ModelResponse:
details = (
CompletionTokensDetailsWrapper(reasoning_tokens=reasoning_tokens)
if reasoning_tokens
else None
)
usage = Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=total_tokens,
completion_tokens_details=details,
)
response = ModelResponse()
setattr(response, "usage", usage)
return response
def test_should_fold_when_total_explained_by_reasoning_gap(self):
# xAI live shape: 14 + 10 + 312 == 336.
response = self._make_response(
prompt_tokens=14,
completion_tokens=10,
total_tokens=336,
reasoning_tokens=312,
)
XAIChatConfig._fold_reasoning_tokens_into_completion(response)
usage = response.usage
assert usage.completion_tokens == 322
assert usage.total_tokens == usage.prompt_tokens + usage.completion_tokens
def test_should_not_fold_when_already_normalised(self):
response = self._make_response(
prompt_tokens=14,
completion_tokens=322,
total_tokens=336,
reasoning_tokens=312,
)
XAIChatConfig._fold_reasoning_tokens_into_completion(response)
assert response.usage.completion_tokens == 322
def test_should_skip_when_no_reasoning_tokens(self):
response = self._make_response(
prompt_tokens=14,
completion_tokens=10,
total_tokens=24,
reasoning_tokens=0,
)
XAIChatConfig._fold_reasoning_tokens_into_completion(response)
assert response.usage.completion_tokens == 10
def test_should_skip_when_gap_does_not_match_reasoning(self):
# Refuse to fold if xAI accounting changes (gap != reasoning_tokens).
response = self._make_response(
prompt_tokens=14,
completion_tokens=10,
total_tokens=999,
reasoning_tokens=312,
)
XAIChatConfig._fold_reasoning_tokens_into_completion(response)
assert response.usage.completion_tokens == 10
assert response.usage.total_tokens == 999
class TestXAIParallelToolCalls:
@ -14,9 +98,7 @@ class TestXAIParallelToolCalls:
def test_get_supported_openai_params_includes_parallel_tool_calls(self):
"""Test that parallel_tool_calls is in supported parameters."""
config = XAIChatConfig()
supported_params = config.get_supported_openai_params(
"xai/grok-4.20"
)
supported_params = config.get_supported_openai_params("xai/grok-4.20")
assert "parallel_tool_calls" in supported_params
def test_transform_request_preserves_parallel_tool_calls(self):

View file

@ -110,7 +110,7 @@ class TestXAICostCalculator:
usage = Usage(
prompt_tokens=10,
completion_tokens=200,
total_tokens=210,
total_tokens=360,
completion_tokens_details=CompletionTokensDetailsWrapper(
accepted_prediction_tokens=0,
audio_tokens=0,
@ -136,7 +136,7 @@ class TestXAICostCalculator:
usage = Usage(
prompt_tokens=20,
completion_tokens=300,
total_tokens=320,
total_tokens=520,
completion_tokens_details=CompletionTokensDetailsWrapper(
accepted_prediction_tokens=0,
audio_tokens=0,
@ -177,7 +177,7 @@ class TestXAICostCalculator:
usage = Usage(
prompt_tokens=12,
completion_tokens=50, # Less than reasoning_tokens
total_tokens=62,
total_tokens=162,
completion_tokens_details=CompletionTokensDetailsWrapper(
accepted_prediction_tokens=0,
audio_tokens=0,
@ -204,7 +204,7 @@ class TestXAICostCalculator:
usage = Usage(
prompt_tokens=150000, # Above 128k threshold
completion_tokens=100000, # Above 128k threshold
total_tokens=250000,
total_tokens=300000,
completion_tokens_details=CompletionTokensDetailsWrapper(
accepted_prediction_tokens=0,
audio_tokens=0,
@ -233,7 +233,7 @@ class TestXAICostCalculator:
usage = Usage(
prompt_tokens=100000, # Below 128k threshold
completion_tokens=50000,
total_tokens=150000,
total_tokens=160000,
completion_tokens_details=CompletionTokensDetailsWrapper(
accepted_prediction_tokens=0,
audio_tokens=0,
@ -261,7 +261,7 @@ class TestXAICostCalculator:
usage = Usage(
prompt_tokens=200000, # Above 128k threshold
completion_tokens=100000,
total_tokens=300000,
total_tokens=350000,
completion_tokens_details=CompletionTokensDetailsWrapper(
accepted_prediction_tokens=0,
audio_tokens=0,
@ -289,7 +289,7 @@ class TestXAICostCalculator:
usage = Usage(
prompt_tokens=150000, # Above 128k threshold
completion_tokens=50000, # Below 128k threshold
total_tokens=200000,
total_tokens=210000,
completion_tokens_details=CompletionTokensDetailsWrapper(
accepted_prediction_tokens=0,
audio_tokens=0,
@ -331,6 +331,29 @@ class TestXAICostCalculator:
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
def test_already_normalised_usage_does_not_double_count_reasoning(self):
"""Cost calc must not double-bill when Usage is already OpenAI-normalised."""
usage = Usage(
prompt_tokens=12,
completion_tokens=200,
total_tokens=212,
completion_tokens_details=CompletionTokensDetailsWrapper(
accepted_prediction_tokens=0,
audio_tokens=0,
reasoning_tokens=100,
rejected_prediction_tokens=0,
text_tokens=None,
),
)
prompt_cost, completion_cost = cost_per_token(model="grok-3-mini", usage=usage)
expected_prompt_cost = 12 * 3e-7
expected_completion_cost = 200 * 5e-7
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
def test_web_search_cost_calculation(self):
"""Test web search cost calculation for X.AI models."""
# Test with web_search_requests in prompt_tokens_details (primary path)

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