mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_yj_may4
This commit is contained in:
commit
e35cd5af76
135 changed files with 11315 additions and 1524 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
175
litellm/litellm_core_utils/cloud_storage_security.py
Normal file
175
litellm/litellm_core_utils/cloud_storage_security.py
Normal 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")
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)}
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
236
litellm/proxy/guardrails/_content_utils.py
Normal file
236
litellm/proxy/guardrails/_content_utils.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 ""
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
23
tests/test_litellm/llms/anthropic/chat/conftest.py
Normal file
23
tests/test_litellm/llms/anthropic/chat/conftest.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue