mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
Merge branch 'BerriAI:litellm_internal_staging' into litellm_internal_staging
This commit is contained in:
commit
00915f15fe
128 changed files with 7531 additions and 904 deletions
|
|
@ -309,9 +309,13 @@ class Cache:
|
|||
param_value = kwargs[param]
|
||||
cache_key += f"{str(param)}: {str(param_value)}"
|
||||
|
||||
verbose_logger.debug("\nCreated cache key: %s", cache_key)
|
||||
hashed_cache_key = Cache._get_hashed_cache_key(cache_key)
|
||||
hashed_cache_key = self._add_namespace_to_cache_key(hashed_cache_key, **kwargs)
|
||||
verbose_logger.debug(
|
||||
"\nCreated cache key: %s (source material length: %d)",
|
||||
hashed_cache_key,
|
||||
len(cache_key),
|
||||
)
|
||||
# Remove preset_cache_key from kwargs to avoid "got multiple values" TypeError
|
||||
# when kwargs already contains preset_cache_key from upstream callers
|
||||
kwargs_for_preset = {k: v for k, v in kwargs.items() if k != "preset_cache_key"}
|
||||
|
|
@ -497,6 +501,34 @@ class Cache:
|
|||
return cached_response
|
||||
return cached_result
|
||||
|
||||
@staticmethod
|
||||
def _get_safe_cache_lookup_kwargs(kwargs: Dict[str, Any]) -> Dict[str, Any]:
|
||||
cache_lookup_kwargs: Dict[str, Any] = {}
|
||||
for prompt_kwarg in ("messages", "input"):
|
||||
if prompt_kwarg in kwargs:
|
||||
cache_lookup_kwargs[prompt_kwarg] = kwargs[prompt_kwarg]
|
||||
|
||||
if isinstance(kwargs.get("metadata"), dict):
|
||||
cache_lookup_kwargs["metadata"] = {}
|
||||
|
||||
return cache_lookup_kwargs
|
||||
|
||||
@staticmethod
|
||||
def _update_metadata_from_cache_lookup_kwargs(
|
||||
original_kwargs: Dict[str, Any], cache_lookup_kwargs: Dict[str, Any]
|
||||
) -> None:
|
||||
original_metadata = original_kwargs.get("metadata")
|
||||
cache_lookup_metadata = cache_lookup_kwargs.get("metadata")
|
||||
if not isinstance(original_metadata, dict) or not isinstance(
|
||||
cache_lookup_metadata, dict
|
||||
):
|
||||
return
|
||||
|
||||
if "semantic-similarity" in cache_lookup_metadata:
|
||||
original_metadata["semantic-similarity"] = cache_lookup_metadata[
|
||||
"semantic-similarity"
|
||||
]
|
||||
|
||||
def get_cache(self, dynamic_cache_object: Optional[BaseCache] = None, **kwargs):
|
||||
"""
|
||||
Retrieves the cached result for the given arguments.
|
||||
|
|
@ -511,7 +543,6 @@ class Cache:
|
|||
try: # never block execution
|
||||
if self.should_use_cache(**kwargs) is not True:
|
||||
return
|
||||
messages = kwargs.get("messages", [])
|
||||
if "cache_key" in kwargs:
|
||||
cache_key = kwargs["cache_key"]
|
||||
else:
|
||||
|
|
@ -523,12 +554,19 @@ class Cache:
|
|||
or cache_control_args.get("s-max-age")
|
||||
or float("inf")
|
||||
)
|
||||
cache_lookup_kwargs = self._get_safe_cache_lookup_kwargs(kwargs)
|
||||
if dynamic_cache_object is not None:
|
||||
cached_result = dynamic_cache_object.get_cache(
|
||||
cache_key, messages=messages
|
||||
cache_key, **cache_lookup_kwargs
|
||||
)
|
||||
else:
|
||||
cached_result = self.cache.get_cache(cache_key, messages=messages)
|
||||
cached_result = self.cache.get_cache(
|
||||
cache_key, **cache_lookup_kwargs
|
||||
)
|
||||
self._update_metadata_from_cache_lookup_kwargs(
|
||||
original_kwargs=kwargs,
|
||||
cache_lookup_kwargs=cache_lookup_kwargs,
|
||||
)
|
||||
return self._get_cache_logic(
|
||||
cached_result=cached_result, max_age=max_age
|
||||
)
|
||||
|
|
@ -549,7 +587,6 @@ class Cache:
|
|||
if self.should_use_cache(**kwargs) is not True:
|
||||
return
|
||||
|
||||
kwargs.get("messages", [])
|
||||
if "cache_key" in kwargs:
|
||||
cache_key = kwargs["cache_key"]
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -213,6 +213,78 @@ class RedisSemanticCache(BaseCache):
|
|||
ttl = int(ttl)
|
||||
return ttl
|
||||
|
||||
@classmethod
|
||||
def _get_prompt_from_kwargs(cls, **kwargs) -> Optional[str]:
|
||||
"""
|
||||
Extract a semantic-cache prompt from chat or Responses API request kwargs.
|
||||
"""
|
||||
messages = kwargs.get("messages")
|
||||
if messages:
|
||||
return get_str_from_messages(messages)
|
||||
|
||||
if "input" not in kwargs:
|
||||
return None
|
||||
|
||||
prompt_parts: List[str] = []
|
||||
cls._collect_responses_input_text(kwargs.get("input"), prompt_parts)
|
||||
prompt = "\n".join(prompt_parts).strip()
|
||||
return prompt or None
|
||||
|
||||
@classmethod
|
||||
def _collect_responses_input_text(cls, value: Any, prompt_parts: List[str]) -> None:
|
||||
value = cls._coerce_response_input_value(value)
|
||||
if value is None:
|
||||
return
|
||||
|
||||
if isinstance(value, str):
|
||||
stripped_value = value.strip()
|
||||
if stripped_value:
|
||||
prompt_parts.append(stripped_value)
|
||||
return
|
||||
|
||||
if isinstance(value, (list, tuple)):
|
||||
for item in value:
|
||||
cls._collect_responses_input_text(item, prompt_parts)
|
||||
return
|
||||
|
||||
if isinstance(value, dict):
|
||||
content = value.get("content")
|
||||
if content is not None:
|
||||
cls._collect_responses_input_text(content, prompt_parts)
|
||||
return
|
||||
|
||||
for text_key in ("text", "output", "input_text", "output_text"):
|
||||
text_value = value.get(text_key)
|
||||
if isinstance(text_value, str):
|
||||
stripped_text = text_value.strip()
|
||||
if stripped_text:
|
||||
prompt_parts.append(stripped_text)
|
||||
return
|
||||
return
|
||||
|
||||
content = getattr(value, "content", None)
|
||||
if content is not None:
|
||||
cls._collect_responses_input_text(content, prompt_parts)
|
||||
return
|
||||
|
||||
for text_key in ("text", "output", "input_text", "output_text"):
|
||||
text_value = getattr(value, text_key, None)
|
||||
if isinstance(text_value, str):
|
||||
stripped_text = text_value.strip()
|
||||
if stripped_text:
|
||||
prompt_parts.append(stripped_text)
|
||||
return
|
||||
|
||||
@staticmethod
|
||||
def _coerce_response_input_value(value: Any) -> Any:
|
||||
model_dump = getattr(value, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
return model_dump()
|
||||
dict_method = getattr(value, "dict", None)
|
||||
if callable(dict_method):
|
||||
return dict_method()
|
||||
return value
|
||||
|
||||
def _get_embedding(self, prompt: str) -> List[float]:
|
||||
"""
|
||||
Generate an embedding vector for the given prompt using the configured embedding model.
|
||||
|
|
@ -278,13 +350,11 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
value_str: Optional[str] = None
|
||||
try:
|
||||
# Extract the prompt from messages
|
||||
messages = kwargs.get("messages", [])
|
||||
if not messages:
|
||||
print_verbose("No messages provided for semantic caching")
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
prompt = get_str_from_messages(messages)
|
||||
value_str = str(value)
|
||||
|
||||
store_kwargs: Dict[str, Any] = {
|
||||
|
|
@ -315,14 +385,12 @@ class RedisSemanticCache(BaseCache):
|
|||
print_verbose(f"Redis semantic-cache get_cache, kwargs: {kwargs}")
|
||||
|
||||
try:
|
||||
# Extract the prompt from messages
|
||||
messages = kwargs.get("messages", [])
|
||||
if not messages:
|
||||
print_verbose("No messages provided for semantic cache lookup")
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt 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 in this exact
|
||||
# LiteLLM cache-key scope.
|
||||
check_kwargs: Dict[str, Any] = {
|
||||
|
|
@ -428,13 +496,11 @@ class RedisSemanticCache(BaseCache):
|
|||
print_verbose(f"Async Redis semantic-cache set_cache, kwargs: {kwargs}")
|
||||
|
||||
try:
|
||||
# Extract the prompt from messages
|
||||
messages = kwargs.get("messages", [])
|
||||
if not messages:
|
||||
print_verbose("No messages provided for semantic caching")
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
prompt = get_str_from_messages(messages)
|
||||
value_str = str(value)
|
||||
|
||||
# Generate embedding for the value (response) to cache
|
||||
|
|
@ -471,15 +537,12 @@ class RedisSemanticCache(BaseCache):
|
|||
print_verbose(f"Async Redis semantic-cache get_cache, kwargs: {kwargs}")
|
||||
|
||||
try:
|
||||
# Extract the prompt from messages
|
||||
messages = kwargs.get("messages", [])
|
||||
if not messages:
|
||||
print_verbose("No messages provided for semantic cache lookup")
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic cache lookup")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
prompt = get_str_from_messages(messages)
|
||||
|
||||
# Generate embedding for the prompt
|
||||
prompt_embedding = await self._get_async_embedding(prompt, **kwargs)
|
||||
|
||||
|
|
|
|||
|
|
@ -831,6 +831,7 @@ openai_compatible_providers: List = [
|
|||
"nano-gpt", # Nano-GPT - JSON-configured provider
|
||||
"poe", # Poe - JSON-configured provider
|
||||
"chutes", # Chutes - JSON-configured provider
|
||||
"parasail", # Parasail - JSON-configured provider
|
||||
"featherless_ai",
|
||||
"nscale",
|
||||
"nebius",
|
||||
|
|
|
|||
|
|
@ -80,11 +80,15 @@ class FocusLiteLLMDatabase:
|
|||
vt.team_id,
|
||||
vt.key_alias as api_key_alias,
|
||||
tt.team_alias,
|
||||
ut.user_email as user_email
|
||||
ut.user_email as user_email,
|
||||
COALESCE(vt.organization_id, tt.organization_id) as organization_id,
|
||||
ot.organization_alias as organization_alias
|
||||
FROM "LiteLLM_DailyUserSpend" dus
|
||||
LEFT JOIN "LiteLLM_VerificationToken" vt ON dus.api_key = vt.token
|
||||
LEFT JOIN "LiteLLM_TeamTable" tt ON vt.team_id = tt.team_id
|
||||
LEFT JOIN "LiteLLM_UserTable" ut ON dus.user_id = ut.user_id
|
||||
LEFT JOIN "LiteLLM_OrganizationTable" ot
|
||||
ON ot.organization_id = COALESCE(vt.organization_id, tt.organization_id)
|
||||
{where_clause}
|
||||
ORDER BY dus.date DESC, dus.created_at DESC
|
||||
{limit_clause}
|
||||
|
|
|
|||
|
|
@ -2,12 +2,14 @@
|
|||
|
||||
from .base import FocusDestination, FocusTimeWindow
|
||||
from .factory import FocusDestinationFactory
|
||||
from .gcs_destination import FocusGCSDestination
|
||||
from .s3_destination import FocusS3Destination
|
||||
from .vantage_destination import FocusVantageDestination
|
||||
|
||||
__all__ = [
|
||||
"FocusDestination",
|
||||
"FocusDestinationFactory",
|
||||
"FocusGCSDestination",
|
||||
"FocusTimeWindow",
|
||||
"FocusS3Destination",
|
||||
"FocusVantageDestination",
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import os
|
|||
from typing import Any, Dict, Optional
|
||||
|
||||
from .base import FocusDestination
|
||||
from .gcs_destination import FocusGCSDestination
|
||||
from .s3_destination import FocusS3Destination
|
||||
from .vantage_destination import FocusVantageDestination
|
||||
|
||||
|
|
@ -29,6 +30,8 @@ class FocusDestinationFactory:
|
|||
return FocusS3Destination(prefix=prefix, config=normalized_config)
|
||||
if provider_lower == "vantage":
|
||||
return FocusVantageDestination(prefix=prefix, config=normalized_config)
|
||||
if provider_lower == "gcs":
|
||||
return FocusGCSDestination(prefix=prefix, config=normalized_config)
|
||||
raise NotImplementedError(
|
||||
f"Provider '{provider}' not supported for Focus export"
|
||||
)
|
||||
|
|
@ -72,6 +75,18 @@ class FocusDestinationFactory:
|
|||
"VANTAGE_INTEGRATION_TOKEN must be provided for Vantage exports"
|
||||
)
|
||||
return {k: v for k, v in resolved.items() if v is not None}
|
||||
if provider == "gcs":
|
||||
resolved = {
|
||||
"bucket_name": overrides.get("bucket_name")
|
||||
or os.getenv("FOCUS_GCS_BUCKET_NAME"),
|
||||
"service_account_json": overrides.get("service_account_json")
|
||||
or os.getenv("FOCUS_GCS_PATH_SERVICE_ACCOUNT"),
|
||||
}
|
||||
if not resolved.get("bucket_name"):
|
||||
raise ValueError(
|
||||
"FOCUS_GCS_BUCKET_NAME must be provided for GCS exports"
|
||||
)
|
||||
return {k: v for k, v in resolved.items() if v is not None}
|
||||
raise NotImplementedError(
|
||||
f"Provider '{provider}' not supported for Focus export configuration"
|
||||
)
|
||||
|
|
|
|||
74
litellm/integrations/focus/destinations/gcs_destination.py
Normal file
74
litellm/integrations/focus/destinations/gcs_destination.py
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
"""GCS destination for Focus export — reuses GCSBucketBase auth and httpx client."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import timezone
|
||||
from typing import Any, Optional
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase
|
||||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
encode_gcs_object_name_for_url,
|
||||
)
|
||||
|
||||
from .base import FocusDestination, FocusTimeWindow
|
||||
|
||||
|
||||
class FocusGCSDestination(GCSBucketBase, FocusDestination):
|
||||
"""Upload serialized Focus exports to GCS using the GCS JSON API."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
prefix: str,
|
||||
config: Optional[dict[str, Any]] = None,
|
||||
) -> None:
|
||||
config = config or {}
|
||||
bucket_name = config.get("bucket_name")
|
||||
if not bucket_name:
|
||||
raise ValueError("bucket_name must be provided for GCS destination")
|
||||
super().__init__(bucket_name=bucket_name)
|
||||
service_account_json = config.get("service_account_json")
|
||||
if service_account_json is not None:
|
||||
self.path_service_account_json = service_account_json
|
||||
self.prefix = prefix.rstrip("/")
|
||||
|
||||
async def deliver(
|
||||
self,
|
||||
*,
|
||||
content: bytes,
|
||||
time_window: FocusTimeWindow,
|
||||
filename: str,
|
||||
) -> None:
|
||||
object_name = self._build_object_key(time_window=time_window, filename=filename)
|
||||
headers = await self.construct_request_headers(
|
||||
service_account_json=self.path_service_account_json
|
||||
)
|
||||
headers["Content-Type"] = "application/octet-stream"
|
||||
encoded_name = encode_gcs_object_name_for_url(object_name)
|
||||
url = (
|
||||
f"https://storage.googleapis.com/upload/storage/v1/b/"
|
||||
f"{self.BUCKET_NAME}/o?uploadType=media&name={encoded_name}"
|
||||
)
|
||||
response = await self.async_httpx_client.post(
|
||||
url=url, headers=headers, data=content
|
||||
)
|
||||
if response.status_code != 200:
|
||||
raise RuntimeError(
|
||||
f"GCS upload failed: status={response.status_code} body={response.text}"
|
||||
)
|
||||
verbose_logger.debug(
|
||||
"Focus GCS: uploaded %d bytes to gs://%s/%s",
|
||||
len(content),
|
||||
self.BUCKET_NAME,
|
||||
object_name,
|
||||
)
|
||||
|
||||
def _build_object_key(self, *, time_window: FocusTimeWindow, filename: str) -> str:
|
||||
start_utc = time_window.start_time.astimezone(timezone.utc)
|
||||
date_component = f"date={start_utc.strftime('%Y-%m-%d')}"
|
||||
parts = [self.prefix, date_component]
|
||||
if time_window.frequency == "hourly":
|
||||
parts.append(f"hour={start_utc.strftime('%H')}")
|
||||
key_prefix = "/".join(filter(None, parts))
|
||||
return f"{key_prefix}/{filename}" if key_prefix else filename
|
||||
|
|
@ -12,6 +12,8 @@ from .schema import FOCUS_NORMALIZED_SCHEMA
|
|||
_TAG_KEYS = (
|
||||
"team_id",
|
||||
"team_alias",
|
||||
"organization_id",
|
||||
"organization_alias",
|
||||
"user_id",
|
||||
"user_email",
|
||||
"api_key_alias",
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus
|
||||
|
||||
GALILEO_CLOUD_API_BASE_URL = "https://api.galileo.ai"
|
||||
# Cap the in-memory buffer so persistent flush failures (e.g. Galileo
|
||||
|
|
@ -89,6 +90,52 @@ class GalileoObserve(CustomLogger):
|
|||
return bool(self.api_key)
|
||||
return bool(self.username and self.password)
|
||||
|
||||
async def async_health_check(self) -> IntegrationHealthCheckStatus:
|
||||
try:
|
||||
if not self.project_id:
|
||||
return IntegrationHealthCheckStatus(
|
||||
status="unhealthy",
|
||||
error_message="GALILEO_PROJECT_ID environment variable not set",
|
||||
)
|
||||
|
||||
if not self.base_url:
|
||||
return IntegrationHealthCheckStatus(
|
||||
status="unhealthy",
|
||||
error_message="GALILEO_BASE_URL environment variable not set",
|
||||
)
|
||||
|
||||
if not self.use_v2_api and (not self.username or not self.password):
|
||||
return IntegrationHealthCheckStatus(
|
||||
status="unhealthy",
|
||||
error_message=(
|
||||
"GALILEO_API_KEY or GALILEO_USERNAME and GALILEO_PASSWORD "
|
||||
"environment variables must be set"
|
||||
),
|
||||
)
|
||||
|
||||
if not await self._ensure_headers():
|
||||
return IntegrationHealthCheckStatus(
|
||||
status="unhealthy",
|
||||
error_message="Galileo authentication failed",
|
||||
)
|
||||
|
||||
response = await self.async_httpx_handler.get(
|
||||
url=f"{self.base_url}/current_user",
|
||||
headers=self.headers,
|
||||
)
|
||||
if response.status_code >= 400:
|
||||
return IntegrationHealthCheckStatus(
|
||||
status="unhealthy",
|
||||
error_message=(f"Galileo API returned HTTP {response.status_code}"),
|
||||
)
|
||||
|
||||
return IntegrationHealthCheckStatus(status="healthy", error_message=None)
|
||||
except Exception as e:
|
||||
return IntegrationHealthCheckStatus(
|
||||
status="unhealthy",
|
||||
error_message=f"Galileo health check failed: {str(e)}",
|
||||
)
|
||||
|
||||
async def async_set_galileo_headers(self) -> None:
|
||||
galileo_login_response = await self.async_httpx_handler.post(
|
||||
url=f"{self.base_url}/login",
|
||||
|
|
@ -399,9 +446,9 @@ class GalileoObserve(CustomLogger):
|
|||
return prompt
|
||||
|
||||
@staticmethod
|
||||
def _serialize_galileo_output(value: Any) -> Optional[str]:
|
||||
def _serialize_galileo_output(value: Any) -> str:
|
||||
if value is None:
|
||||
return None
|
||||
return ""
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
|
||||
|
|
@ -460,11 +507,11 @@ class GalileoObserve(CustomLogger):
|
|||
response_obj: Any,
|
||||
level: str = "DEFAULT",
|
||||
status_message: Optional[str] = None,
|
||||
) -> Tuple[str, Optional[str], Any]:
|
||||
) -> Tuple[str, str, Any]:
|
||||
"""
|
||||
Mirror Langfuse _get_langfuse_input_output_content for Galileo ingest.
|
||||
|
||||
Returns (input_text, output_text, messages_for_span). output_text None skips ingest.
|
||||
Returns (input_text, output_text, messages_for_span).
|
||||
"""
|
||||
call_type = kwargs.get("call_type")
|
||||
prompt = self._build_prompt(kwargs)
|
||||
|
|
@ -477,10 +524,11 @@ class GalileoObserve(CustomLogger):
|
|||
return self._prompt_to_input_text(prompt), status_message, prompt
|
||||
|
||||
if response_obj is not None and (
|
||||
call_type == "embedding"
|
||||
call_type in ("embedding", "aembedding")
|
||||
or isinstance(response_obj, litellm.EmbeddingResponse)
|
||||
):
|
||||
return self._prompt_to_input_text(prompt), None, prompt
|
||||
# Match Langfuse OTEL: log embeddings without serializing vectors.
|
||||
return self._prompt_to_input_text(prompt), "embedding-output", prompt
|
||||
|
||||
if response_obj is not None and isinstance(response_obj, litellm.ModelResponse):
|
||||
output = self._get_chat_content_for_galileo(response_obj)
|
||||
|
|
@ -549,7 +597,7 @@ class GalileoObserve(CustomLogger):
|
|||
):
|
||||
input_val = kwargs.get("input")
|
||||
return (
|
||||
self._serialize_galileo_output(input_val) or "",
|
||||
self._serialize_galileo_output(input_val),
|
||||
self._serialize_galileo_output(response_obj),
|
||||
input_val,
|
||||
)
|
||||
|
|
@ -574,11 +622,11 @@ class GalileoObserve(CustomLogger):
|
|||
kwargs.get("messages") or [],
|
||||
)
|
||||
|
||||
return self._prompt_to_input_text(prompt), None, kwargs.get("messages") or []
|
||||
return self._prompt_to_input_text(prompt), "", kwargs.get("messages") or []
|
||||
|
||||
def get_output_str_from_response(
|
||||
self, response_obj: Any, kwargs: Dict[str, Any]
|
||||
) -> Optional[str]:
|
||||
) -> str:
|
||||
_, output_text, _ = self._get_galileo_input_output_content(
|
||||
kwargs=kwargs, response_obj=response_obj
|
||||
)
|
||||
|
|
@ -659,11 +707,6 @@ class GalileoObserve(CustomLogger):
|
|||
input_text, output_text, messages = self._get_galileo_input_output_content(
|
||||
kwargs=kwargs, response_obj=response_obj
|
||||
)
|
||||
if output_text is None:
|
||||
verbose_logger.debug(
|
||||
"Galileo Logger: skipping %s — no text output to log", _call_type
|
||||
)
|
||||
return
|
||||
|
||||
raw_start = slo.get("startTime")
|
||||
raw_end = slo.get("endTime")
|
||||
|
|
|
|||
|
|
@ -17,6 +17,10 @@ from litellm.integrations.custom_logger import CustomLogger
|
|||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
)
|
||||
from litellm.llms.vertex_ai.common_utils import (
|
||||
redact_vertex_ai_metadata_from_litellm_params,
|
||||
redact_vertex_ai_metadata_from_logged_object,
|
||||
)
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
|
|
@ -119,10 +123,12 @@ def _redact_standard_logging_object(model_call_details: dict):
|
|||
# ResponsesAPIResponse format - redact content in output items
|
||||
if isinstance(response.get("output"), list):
|
||||
_redact_responses_api_output_dict(response["output"], redacted_str)
|
||||
redact_vertex_ai_metadata_from_logged_object(response)
|
||||
elif isinstance(response, dict) and "choices" in response:
|
||||
# ModelResponse dict format - redact content in choices
|
||||
if isinstance(response.get("choices"), list):
|
||||
_redact_model_response_dict_choices(response["choices"], redacted_str)
|
||||
redact_vertex_ai_metadata_from_logged_object(response)
|
||||
elif isinstance(response, str):
|
||||
standard_logging_object["response"] = redacted_str
|
||||
else:
|
||||
|
|
@ -164,6 +170,7 @@ def perform_redaction(model_call_details: dict, result):
|
|||
model_call_details["prompt"] = ""
|
||||
model_call_details["input"] = ""
|
||||
_redact_standard_logging_object(model_call_details)
|
||||
redact_vertex_ai_metadata_from_litellm_params(model_call_details)
|
||||
|
||||
# Redact streaming response
|
||||
if (
|
||||
|
|
@ -174,6 +181,7 @@ def perform_redaction(model_call_details: dict, result):
|
|||
if hasattr(_streaming_response, "choices"):
|
||||
for choice in _streaming_response.choices:
|
||||
_redact_choice_content(choice)
|
||||
redact_vertex_ai_metadata_from_logged_object(_streaming_response)
|
||||
elif hasattr(_streaming_response, "output"):
|
||||
_redact_responses_api_output(_streaming_response.output)
|
||||
# Redact reasoning field in ResponsesAPIResponse
|
||||
|
|
@ -200,12 +208,14 @@ def perform_redaction(model_call_details: dict, result):
|
|||
if hasattr(_result, "choices") and _result.choices is not None:
|
||||
for choice in _result.choices:
|
||||
_redact_choice_content(choice)
|
||||
redact_vertex_ai_metadata_from_logged_object(_result)
|
||||
elif isinstance(_result, dict) and "choices" in _result:
|
||||
# Handle dict representation of ModelResponse (e.g., from model_dump())
|
||||
if _result.get("choices") is not None:
|
||||
_redact_model_response_dict_choices(
|
||||
_result["choices"], "redacted-by-litellm"
|
||||
)
|
||||
redact_vertex_ai_metadata_from_logged_object(_result)
|
||||
elif isinstance(_result, dict) and "output" in _result:
|
||||
if isinstance(_result.get("output"), list):
|
||||
_redact_responses_api_output_dict(
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from litellm.types.utils import (
|
|||
ServerToolUse,
|
||||
Usage,
|
||||
)
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.utils import print_verbose, token_counter
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -79,6 +80,54 @@ class ChunkProcessor:
|
|||
model_response._hidden_params = chunk.get("_hidden_params", {})
|
||||
return model_response
|
||||
|
||||
@staticmethod
|
||||
def apply_provider_assembled_streaming_metadata(
|
||||
response: ModelResponse,
|
||||
chunks: List[Any],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> None:
|
||||
if not chunks:
|
||||
return
|
||||
|
||||
model = getattr(response, "model", None)
|
||||
if not model:
|
||||
return
|
||||
|
||||
custom_llm_provider = None
|
||||
if logging_obj is not None:
|
||||
custom_llm_provider = logging_obj.model_call_details.get(
|
||||
"custom_llm_provider"
|
||||
)
|
||||
|
||||
try:
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import (
|
||||
get_llm_provider,
|
||||
)
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
if custom_llm_provider:
|
||||
provider = LlmProviders(custom_llm_provider)
|
||||
else:
|
||||
_, provider_str, _, _ = get_llm_provider(model)
|
||||
provider = LlmProviders(provider_str)
|
||||
|
||||
provider_config = ProviderConfigManager.get_provider_chat_config(
|
||||
model=model,
|
||||
provider=provider,
|
||||
)
|
||||
if provider_config is not None:
|
||||
provider_config.apply_assembled_streaming_response_metadata(
|
||||
response=response,
|
||||
chunks=chunks,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"apply_provider_assembled_streaming_metadata failed for model=%s: %s",
|
||||
model,
|
||||
e,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_chunk_id(chunks: List[Dict[str, Any]]) -> str:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1607,6 +1607,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
)
|
||||
return _tool
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
"""
|
||||
Whether to drop x-anthropic-billing-header system blocks before sending upstream.
|
||||
|
||||
The first-party Anthropic API uses these blocks for Claude Code attribution, so the
|
||||
base config keeps them. Providers that reject them (e.g. Bedrock) override this to True.
|
||||
"""
|
||||
return False
|
||||
|
||||
def translate_system_message(
|
||||
self, messages: List[AllMessageValues]
|
||||
) -> List[AnthropicSystemMessageContent]:
|
||||
|
|
@ -1614,7 +1623,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
Translate system message to anthropic format.
|
||||
|
||||
Removes system message from the original list and returns a new list of anthropic system message content.
|
||||
Filters out system messages containing x-anthropic-billing-header metadata.
|
||||
When should_strip_billing_metadata() is True, x-anthropic-billing-header system blocks are dropped.
|
||||
"""
|
||||
system_prompt_indices = []
|
||||
anthropic_system_message_list: List[AnthropicSystemMessageContent] = []
|
||||
|
|
@ -1626,10 +1635,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
# Skip empty text blocks - Anthropic API raises errors for empty text
|
||||
if not system_message_block["content"]:
|
||||
continue
|
||||
# Skip system messages containing x-anthropic-billing-header metadata
|
||||
if system_message_block["content"].startswith(
|
||||
"x-anthropic-billing-header:"
|
||||
):
|
||||
if self.should_strip_billing_metadata() and system_message_block[
|
||||
"content"
|
||||
].startswith("x-anthropic-billing-header:"):
|
||||
continue
|
||||
anthropic_system_message_content = AnthropicSystemMessageContent(
|
||||
type="text",
|
||||
|
|
@ -1648,9 +1656,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
text_value = _content.get("text")
|
||||
if _content.get("type") == "text" and not text_value:
|
||||
continue
|
||||
# Skip system messages containing x-anthropic-billing-header metadata
|
||||
if (
|
||||
_content.get("type") == "text"
|
||||
self.should_strip_billing_metadata()
|
||||
and _content.get("type") == "text"
|
||||
and text_value
|
||||
and text_value.startswith("x-anthropic-billing-header:")
|
||||
):
|
||||
|
|
|
|||
|
|
@ -84,6 +84,15 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
if isinstance(content, list):
|
||||
_process_content_list(content)
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
"""
|
||||
Whether to drop x-anthropic-billing-header system blocks before sending upstream.
|
||||
|
||||
The first-party Anthropic API uses these blocks for Claude Code attribution, so the
|
||||
base config keeps them. Providers that reject them override this to True.
|
||||
"""
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _filter_billing_headers_from_system(system_param):
|
||||
"""
|
||||
|
|
@ -286,14 +295,12 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
optional_params=anthropic_messages_optional_request_params,
|
||||
)
|
||||
|
||||
# Filter out x-anthropic-billing-header from system messages
|
||||
system_param = anthropic_messages_optional_request_params.get("system")
|
||||
if system_param is not None:
|
||||
if self.should_strip_billing_metadata() and system_param is not None:
|
||||
filtered_system = self._filter_billing_headers_from_system(system_param)
|
||||
if filtered_system is not None and len(filtered_system) > 0:
|
||||
anthropic_messages_optional_request_params["system"] = filtered_system
|
||||
else:
|
||||
# Remove system parameter if all content was filtered out
|
||||
anthropic_messages_optional_request_params.pop("system", None)
|
||||
|
||||
# Transform context_management from OpenAI format to Anthropic format if needed
|
||||
|
|
|
|||
|
|
@ -43,7 +43,10 @@ from .common_utils import (
|
|||
process_azure_headers,
|
||||
select_azure_base_url_or_endpoint,
|
||||
)
|
||||
from .image_generation import get_azure_image_generation_config
|
||||
from .image_generation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
get_azure_image_generation_config,
|
||||
)
|
||||
from .image_generation.http_utils import azure_deployment_image_generation_json_body
|
||||
|
||||
|
||||
|
|
@ -1097,10 +1100,14 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
)
|
||||
|
||||
def create_azure_base_url(
|
||||
self, azure_client_params: dict, model: Optional[str]
|
||||
self,
|
||||
azure_client_params: dict,
|
||||
model: Optional[str],
|
||||
base_model: Optional[str] = None,
|
||||
) -> str:
|
||||
from litellm.llms.azure_ai.image_generation import (
|
||||
AzureFoundryFluxImageGenerationConfig,
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
)
|
||||
|
||||
api_base: str = azure_client_params.get(
|
||||
|
|
@ -1112,6 +1119,12 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
if model is None:
|
||||
model = ""
|
||||
|
||||
if AzureFoundryMAIImageGenerationConfig.is_mai_model(base_model or model):
|
||||
return AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url(
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
)
|
||||
|
||||
# Handle FLUX 2 models on Azure AI which use a different URL pattern
|
||||
# e.g., /providers/blackforestlabs/v1/flux-2-pro instead of /openai/deployments/{model}/images/generations
|
||||
if AzureFoundryFluxImageGenerationConfig.is_flux2_model(model):
|
||||
|
|
@ -1153,10 +1166,10 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
if api_base.endswith("/"):
|
||||
api_base = api_base.rstrip("/")
|
||||
api_version: str = azure_client_params.get("api_version", "")
|
||||
# Use the deployment name (model) for URL construction, not the base_model from data
|
||||
img_gen_api_base = self.create_azure_base_url(
|
||||
azure_client_params=azure_client_params,
|
||||
model=model or data.get("model", ""),
|
||||
base_model=data.get("model", ""),
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
|
|
@ -1285,9 +1298,10 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
if aimg_generation is True:
|
||||
return self.aimage_generation(data=data, input=input, logging_obj=logging_obj, model_response=model_response, api_key=api_key, client=client, azure_client_params=azure_client_params, timeout=timeout, headers=headers, model=model) # type: ignore
|
||||
|
||||
# Use the deployment name (model) for URL construction, not the base_model from data
|
||||
img_gen_api_base = self.create_azure_base_url(
|
||||
azure_client_params=azure_client_params, model=model
|
||||
azure_client_params=azure_client_params,
|
||||
model=model,
|
||||
base_model=base_model,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
|
|
@ -1309,6 +1323,21 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
data=data,
|
||||
headers=headers,
|
||||
)
|
||||
provider_config = get_azure_image_generation_config(
|
||||
data.get("model", "dall-e-2")
|
||||
)
|
||||
if isinstance(provider_config, AzureFoundryMAIImageGenerationConfig):
|
||||
return provider_config.transform_image_generation_response(
|
||||
model=data.get("model", "dall-e-2"),
|
||||
raw_response=httpx_response,
|
||||
model_response=model_response or ImageResponse(),
|
||||
logging_obj=logging_obj,
|
||||
request_data=data,
|
||||
optional_params=data,
|
||||
litellm_params=data,
|
||||
encoding=litellm.encoding,
|
||||
)
|
||||
|
||||
response = httpx_response.json()
|
||||
|
||||
## LOGGING
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.azure_ai.image_generation import AzureFoundryMAIImageGenerationConfig
|
||||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
|
|
@ -24,6 +25,8 @@ def get_azure_image_generation_config(model: str) -> BaseImageGenerationConfig:
|
|||
return AzureDallE2ImageGenerationConfig()
|
||||
elif "dalle3" in model:
|
||||
return AzureDallE3ImageGenerationConfig()
|
||||
elif AzureFoundryMAIImageGenerationConfig.is_mai_model(model):
|
||||
return AzureFoundryMAIImageGenerationConfig()
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
f"Using AzureGPTImageGenerationConfig for model: {model}. This follows the gpt-image model format."
|
||||
|
|
|
|||
|
|
@ -21,6 +21,9 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig):
|
|||
and Azure endpoint format.
|
||||
"""
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
def validate_anthropic_messages_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -40,6 +40,9 @@ class AzureAnthropicConfig(AnthropicConfig):
|
|||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "azure_ai"
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -1,21 +1,33 @@
|
|||
from litellm.llms.azure_ai.image_generation.flux_transformation import (
|
||||
AzureFoundryFluxImageGenerationConfig,
|
||||
)
|
||||
from litellm.llms.azure_ai.image_generation.mai_transformation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
|
||||
|
||||
from .flux2_transformation import AzureFoundryFlux2ImageEditConfig
|
||||
from .mai_transformation import AzureFoundryMAIImageEditConfig
|
||||
from .transformation import AzureFoundryFluxImageEditConfig
|
||||
|
||||
__all__ = ["AzureFoundryFluxImageEditConfig", "AzureFoundryFlux2ImageEditConfig"]
|
||||
__all__ = [
|
||||
"AzureFoundryFluxImageEditConfig",
|
||||
"AzureFoundryFlux2ImageEditConfig",
|
||||
"AzureFoundryMAIImageEditConfig",
|
||||
]
|
||||
|
||||
|
||||
def get_azure_ai_image_edit_config(model: str) -> BaseImageEditConfig:
|
||||
"""
|
||||
Get the appropriate image edit config for an Azure AI model.
|
||||
|
||||
- MAI models use /mai/v1/images/edits with multipart form data and size
|
||||
- FLUX 2 models use JSON with base64 image
|
||||
- FLUX 1 models use multipart/form-data
|
||||
"""
|
||||
if AzureFoundryMAIImageGenerationConfig.is_mai_model(model):
|
||||
return AzureFoundryMAIImageEditConfig()
|
||||
|
||||
# Check if it's a FLUX 2 model
|
||||
if AzureFoundryFluxImageGenerationConfig.is_flux2_model(model):
|
||||
return AzureFoundryFlux2ImageEditConfig()
|
||||
|
|
|
|||
199
litellm/llms/azure_ai/image_edit/mai_transformation.py
Normal file
199
litellm/llms/azure_ai/image_edit/mai_transformation.py
Normal file
|
|
@ -0,0 +1,199 @@
|
|||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast
|
||||
|
||||
import httpx
|
||||
from httpx._types import RequestFiles
|
||||
|
||||
from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo
|
||||
from litellm.llms.azure_ai.image_generation.mai_transformation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
)
|
||||
from litellm.llms.openai.common_utils import OpenAIError
|
||||
from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.images.main import ImageEditOptionalRequestParams
|
||||
from litellm.types.llms.openai import FileTypes
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import ImageResponse
|
||||
from litellm.utils import convert_to_model_response_object
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class AzureFoundryMAIImageEditConfig(OpenAIImageEditConfig):
|
||||
"""Azure AI Foundry MAI image editing (e.g. MAI-Image-2.5)."""
|
||||
|
||||
DEFAULT_SIZE = "1024x1024"
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
return ["prompt", "image", "model", "n", "size"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
image_edit_optional_params: ImageEditOptionalRequestParams,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> Dict:
|
||||
optional_params: Dict[str, Any] = {}
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
|
||||
for key, value in dict(image_edit_optional_params).items():
|
||||
if value is None or key in optional_params:
|
||||
continue
|
||||
|
||||
if key in supported_params:
|
||||
if key == "size" and value:
|
||||
size_param = cast(str, value)
|
||||
self._validate_size_param(size_param)
|
||||
optional_params[key] = size_param
|
||||
else:
|
||||
optional_params[key] = value
|
||||
elif not drop_params:
|
||||
raise ValueError(
|
||||
f"Parameter {key} is not supported for model {model}. "
|
||||
f"Supported parameters are {supported_params}. "
|
||||
f"Set drop_params=True to drop unsupported parameters."
|
||||
)
|
||||
|
||||
if "size" not in optional_params:
|
||||
optional_params["size"] = self.DEFAULT_SIZE
|
||||
|
||||
return optional_params
|
||||
|
||||
def _validate_size_param(self, size: str) -> None:
|
||||
known_sizes = {
|
||||
"1024x1024",
|
||||
"1792x1024",
|
||||
"1024x1792",
|
||||
"512x512",
|
||||
"256x256",
|
||||
}
|
||||
|
||||
if size in known_sizes:
|
||||
return
|
||||
|
||||
if "x" in size:
|
||||
try:
|
||||
tuple(map(int, size.lower().split("x", 1)))
|
||||
return
|
||||
except ValueError:
|
||||
raise ValueError(
|
||||
f"Invalid size format: '{size}'. Expected format 'WIDTHxHEIGHT' (e.g., '1024x1024')."
|
||||
)
|
||||
|
||||
raise ValueError(
|
||||
f"Unsupported size value: '{size}'. "
|
||||
f"Use a known size (e.g., '1024x1024') or a custom 'WIDTHxHEIGHT' string."
|
||||
)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
api_key: Optional[str] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
api_key = AzureFoundryModelInfo.get_api_key(api_key)
|
||||
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
f"Azure AI API key is required for model {model}. "
|
||||
"Set AZURE_AI_API_KEY environment variable or pass api_key parameter."
|
||||
)
|
||||
|
||||
headers.update({"api-key": api_key})
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
model: str,
|
||||
api_base: Optional[str],
|
||||
litellm_params: dict,
|
||||
) -> str:
|
||||
api_base = AzureFoundryModelInfo.get_api_base(api_base)
|
||||
|
||||
if api_base is None:
|
||||
raise ValueError(
|
||||
"Azure AI API base is required. Set AZURE_AI_API_BASE environment variable or pass api_base parameter."
|
||||
)
|
||||
|
||||
api_version = (
|
||||
litellm_params.get("api_version")
|
||||
or get_secret_str("AZURE_AI_API_VERSION")
|
||||
or "preview"
|
||||
)
|
||||
|
||||
return AzureFoundryMAIImageGenerationConfig.get_mai_image_edit_url(
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
)
|
||||
|
||||
def transform_image_edit_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: Optional[str],
|
||||
image: Optional[FileTypes],
|
||||
image_edit_optional_request_params: Dict,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[Dict, RequestFiles]:
|
||||
request_params = {
|
||||
"model": model,
|
||||
**image_edit_optional_request_params,
|
||||
}
|
||||
if prompt is not None:
|
||||
request_params["prompt"] = prompt
|
||||
|
||||
data_without_files = {
|
||||
key: value
|
||||
for key, value in request_params.items()
|
||||
if key not in ["image", "mask"]
|
||||
}
|
||||
files_list: List[Tuple[str, Any]] = []
|
||||
|
||||
if image is not None:
|
||||
image_list = [image] if not isinstance(image, list) else image
|
||||
for _image in image_list:
|
||||
if _image is not None:
|
||||
self._add_image_to_files(
|
||||
files_list=files_list,
|
||||
image=_image,
|
||||
field_name="image",
|
||||
)
|
||||
break
|
||||
|
||||
return data_without_files, files_list
|
||||
|
||||
def transform_image_edit_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
) -> ImageResponse:
|
||||
try:
|
||||
response = raw_response.json()
|
||||
except Exception:
|
||||
raise OpenAIError(
|
||||
message=raw_response.text, status_code=raw_response.status_code
|
||||
)
|
||||
|
||||
if "usage" in response:
|
||||
response["usage"] = (
|
||||
AzureFoundryMAIImageGenerationConfig.normalize_mai_image_usage(
|
||||
response.get("usage")
|
||||
)
|
||||
)
|
||||
|
||||
logging_obj.post_call(
|
||||
input="",
|
||||
api_key="",
|
||||
additional_args={"complete_input_dict": {}},
|
||||
original_response=response,
|
||||
)
|
||||
|
||||
return convert_to_model_response_object(
|
||||
response_object=response,
|
||||
model_response_object=ImageResponse(),
|
||||
response_type="image_generation",
|
||||
)
|
||||
|
|
@ -7,12 +7,14 @@ from .dall_e_2_transformation import AzureFoundryDallE2ImageGenerationConfig
|
|||
from .dall_e_3_transformation import AzureFoundryDallE3ImageGenerationConfig
|
||||
from .flux_transformation import AzureFoundryFluxImageGenerationConfig
|
||||
from .gpt_transformation import AzureFoundryGPTImageGenerationConfig
|
||||
from .mai_transformation import AzureFoundryMAIImageGenerationConfig
|
||||
|
||||
__all__ = [
|
||||
"AzureFoundryFluxImageGenerationConfig",
|
||||
"AzureFoundryGPTImageGenerationConfig",
|
||||
"AzureFoundryDallE2ImageGenerationConfig",
|
||||
"AzureFoundryDallE3ImageGenerationConfig",
|
||||
"AzureFoundryMAIImageGenerationConfig",
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -24,6 +26,8 @@ def get_azure_ai_image_generation_config(model: str) -> BaseImageGenerationConfi
|
|||
return AzureFoundryDallE2ImageGenerationConfig()
|
||||
elif "dalle3" in model:
|
||||
return AzureFoundryDallE3ImageGenerationConfig()
|
||||
elif AzureFoundryMAIImageGenerationConfig.is_mai_model(model):
|
||||
return AzureFoundryMAIImageGenerationConfig()
|
||||
elif "flux" in model:
|
||||
return AzureFoundryFluxImageGenerationConfig()
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
from typing import Any
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
||||
calculate_image_response_cost_from_usage,
|
||||
)
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
|
||||
|
|
@ -9,19 +12,28 @@ def cost_calculator(
|
|||
image_response: Any,
|
||||
) -> float:
|
||||
"""
|
||||
Recraft image generation cost calculator
|
||||
Azure AI image generation cost calculator
|
||||
"""
|
||||
_model_info = litellm.get_model_info(
|
||||
model=model,
|
||||
custom_llm_provider=litellm.LlmProviders.AZURE_AI.value,
|
||||
)
|
||||
output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0
|
||||
num_images: int = 0
|
||||
|
||||
if isinstance(image_response, ImageResponse):
|
||||
token_based_cost = calculate_image_response_cost_from_usage(
|
||||
model=model,
|
||||
image_response=image_response,
|
||||
custom_llm_provider=litellm.LlmProviders.AZURE_AI.value,
|
||||
)
|
||||
if token_based_cost is not None:
|
||||
return token_based_cost
|
||||
|
||||
output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0
|
||||
num_images: int = 0
|
||||
if image_response.data:
|
||||
num_images = len(image_response.data)
|
||||
return output_cost_per_image * num_images
|
||||
else:
|
||||
raise ValueError(
|
||||
f"image_response must be of type ImageResponse got type={type(image_response)}"
|
||||
)
|
||||
|
||||
raise ValueError(
|
||||
f"image_response must be of type ImageResponse got type={type(image_response)}"
|
||||
)
|
||||
|
|
|
|||
236
litellm/llms/azure_ai/image_generation/mai_transformation.py
Normal file
236
litellm/llms/azure_ai/image_generation/mai_transformation.py
Normal file
|
|
@ -0,0 +1,236 @@
|
|||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
from litellm.llms.openai.common_utils import OpenAIError
|
||||
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
|
||||
from litellm.types.utils import ImageResponse
|
||||
from litellm.utils import convert_to_model_response_object
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
|
||||
"""Azure AI Foundry MAI image generation (e.g. MAI-Image-2.5)."""
|
||||
|
||||
DEFAULT_WIDTH = 1024
|
||||
DEFAULT_HEIGHT = 1024
|
||||
|
||||
@staticmethod
|
||||
def get_mai_image_generation_url(
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str],
|
||||
) -> str:
|
||||
if api_base is None:
|
||||
raise ValueError("api_base is required for Azure AI MAI image generation")
|
||||
|
||||
api_version = api_version or "preview"
|
||||
path, separator, query = api_base.partition("?")
|
||||
path = path.rstrip("/")
|
||||
|
||||
if "/mai/" in path:
|
||||
prefix, _, _ = path.partition("/images/")
|
||||
path = f"{prefix}/images/generations"
|
||||
else:
|
||||
path = f"{path}/mai/v1/images/generations"
|
||||
|
||||
if separator:
|
||||
return f"{path}?{query}"
|
||||
return f"{path}?api-version={api_version}"
|
||||
|
||||
@staticmethod
|
||||
def get_mai_image_edit_url(
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str],
|
||||
) -> str:
|
||||
if api_base is None:
|
||||
raise ValueError("api_base is required for Azure AI MAI image editing")
|
||||
|
||||
api_version = api_version or "preview"
|
||||
path, separator, query = api_base.partition("?")
|
||||
path = path.rstrip("/")
|
||||
|
||||
if "/mai/" in path:
|
||||
prefix, _, _ = path.partition("/images/")
|
||||
path = f"{prefix}/images/edits"
|
||||
else:
|
||||
path = f"{path}/mai/v1/images/edits"
|
||||
|
||||
if separator:
|
||||
return f"{path}?{query}"
|
||||
return f"{path}?api-version={api_version}"
|
||||
|
||||
@staticmethod
|
||||
def is_mai_model(model: str) -> bool:
|
||||
model_normalized = model.lower().replace("-", "").replace("_", "")
|
||||
return "maiimage" in model_normalized
|
||||
|
||||
@staticmethod
|
||||
def normalize_mai_image_usage(usage: Optional[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
"""Map Azure MAI usage fields to OpenAI ImageUsage schema."""
|
||||
if usage is None:
|
||||
return {
|
||||
"input_tokens": 0,
|
||||
"input_tokens_details": {"image_tokens": 0, "text_tokens": 0},
|
||||
"output_tokens": 0,
|
||||
"total_tokens": 0,
|
||||
}
|
||||
|
||||
normalized_usage = dict(usage)
|
||||
input_tokens_details = normalized_usage.get("input_tokens_details")
|
||||
if not isinstance(input_tokens_details, dict):
|
||||
input_tokens_details = {}
|
||||
|
||||
text_tokens = normalized_usage.get("num_input_text_tokens")
|
||||
if text_tokens is None:
|
||||
text_tokens = input_tokens_details.get("text_tokens")
|
||||
if text_tokens is None:
|
||||
text_tokens = normalized_usage.get("input_tokens", 0) or 0
|
||||
|
||||
image_tokens = normalized_usage.get("num_input_image_tokens")
|
||||
if image_tokens is None:
|
||||
image_tokens = input_tokens_details.get("image_tokens")
|
||||
if image_tokens is None:
|
||||
image_tokens = 0
|
||||
|
||||
output_tokens = normalized_usage.get("output_tokens")
|
||||
if output_tokens is None:
|
||||
output_tokens = normalized_usage.get("num_output_tokens")
|
||||
if output_tokens is None:
|
||||
output_tokens = normalized_usage.get("output_image_tokens")
|
||||
if output_tokens is None:
|
||||
output_tokens = 0
|
||||
|
||||
input_tokens = normalized_usage.get("input_tokens")
|
||||
if input_tokens is None:
|
||||
input_tokens = text_tokens + image_tokens
|
||||
|
||||
total_tokens = normalized_usage.get("total_tokens")
|
||||
if total_tokens is None:
|
||||
total_tokens = input_tokens + output_tokens
|
||||
|
||||
normalized_usage.update(
|
||||
{
|
||||
"input_tokens": input_tokens,
|
||||
"input_tokens_details": {
|
||||
"image_tokens": image_tokens,
|
||||
"text_tokens": text_tokens,
|
||||
},
|
||||
"output_tokens": output_tokens,
|
||||
"total_tokens": total_tokens,
|
||||
}
|
||||
)
|
||||
return normalized_usage
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIImageGenerationOptionalParams]:
|
||||
return ["n", "size"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
|
||||
for k, v in non_default_params.items():
|
||||
if k in optional_params:
|
||||
continue
|
||||
|
||||
if k in supported_params:
|
||||
if k == "size" and v:
|
||||
self._map_size_param(v, optional_params)
|
||||
else:
|
||||
optional_params[k] = v
|
||||
elif k in ("width", "height"):
|
||||
optional_params[k] = v
|
||||
elif not drop_params:
|
||||
raise ValueError(
|
||||
f"Parameter {k} is not supported for model {model}. "
|
||||
f"Supported parameters are {supported_params} and width/height. "
|
||||
f"Set drop_params=True to drop unsupported parameters."
|
||||
)
|
||||
|
||||
if "width" not in optional_params:
|
||||
optional_params["width"] = self.DEFAULT_WIDTH
|
||||
if "height" not in optional_params:
|
||||
optional_params["height"] = self.DEFAULT_HEIGHT
|
||||
|
||||
optional_params.pop("size", None)
|
||||
return optional_params
|
||||
|
||||
def _map_size_param(self, size: str, optional_params: dict) -> None:
|
||||
size_mapping = {
|
||||
"1024x1024": (1024, 1024),
|
||||
"1792x1024": (1792, 1024),
|
||||
"1024x1792": (1024, 1792),
|
||||
"512x512": (512, 512),
|
||||
"256x256": (256, 256),
|
||||
}
|
||||
|
||||
if size in size_mapping:
|
||||
width, height = size_mapping[size]
|
||||
optional_params["width"] = width
|
||||
optional_params["height"] = height
|
||||
elif "x" in size:
|
||||
try:
|
||||
width, height = map(int, size.lower().split("x"))
|
||||
optional_params["width"] = width
|
||||
optional_params["height"] = height
|
||||
except ValueError:
|
||||
raise ValueError(
|
||||
f"Invalid size format: '{size}'. Expected format 'WIDTHxHEIGHT' (e.g., '1024x1024')."
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported size value: '{size}'. "
|
||||
f"Use a known size (e.g., '1024x1024') or a custom 'WIDTHxHEIGHT' string."
|
||||
)
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ImageResponse:
|
||||
try:
|
||||
response = raw_response.json()
|
||||
except Exception:
|
||||
raise OpenAIError(
|
||||
message=raw_response.text, status_code=raw_response.status_code
|
||||
)
|
||||
|
||||
if "usage" in response:
|
||||
response["usage"] = self.normalize_mai_image_usage(response.get("usage"))
|
||||
|
||||
logging_obj.post_call(
|
||||
input=request_data.get("prompt", ""),
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
original_response=response,
|
||||
)
|
||||
|
||||
image_response: ImageResponse = convert_to_model_response_object(
|
||||
response_object=response,
|
||||
model_response_object=model_response,
|
||||
response_type="image_generation",
|
||||
)
|
||||
|
||||
width = optional_params.get("width", self.DEFAULT_WIDTH)
|
||||
height = optional_params.get("height", self.DEFAULT_HEIGHT)
|
||||
image_response.size = f"{width}x{height}" # type: ignore[assignment]
|
||||
return image_response
|
||||
|
|
@ -442,6 +442,14 @@ class BaseConfig(ABC):
|
|||
"""Hook for providers to post-process streaming responses. Default: pass-through."""
|
||||
return stream
|
||||
|
||||
def apply_assembled_streaming_response_metadata(
|
||||
self,
|
||||
response: "ModelResponse",
|
||||
chunks: List[Any],
|
||||
) -> None:
|
||||
"""Hook for providers to merge chunk metadata into assembled streaming responses."""
|
||||
return None
|
||||
|
||||
def calculate_additional_costs(
|
||||
self, model: str, prompt_tokens: int, completion_tokens: int
|
||||
) -> Optional[dict]:
|
||||
|
|
|
|||
|
|
@ -62,6 +62,26 @@ class BaseResponsesAPIConfig(ABC):
|
|||
"""
|
||||
return False
|
||||
|
||||
def sign_request(
|
||||
self,
|
||||
headers: dict,
|
||||
optional_params: dict,
|
||||
request_data: dict,
|
||||
api_base: str,
|
||||
api_key: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
stream: Optional[bool] = None,
|
||||
fake_stream: Optional[bool] = None,
|
||||
) -> Tuple[dict, Optional[bytes]]:
|
||||
"""Sign the request after the body is finalized.
|
||||
|
||||
Default is a no-op (returns headers unchanged, no signed body). Providers
|
||||
whose endpoint requires request signing (e.g. Bedrock Mantle SigV4)
|
||||
override this and return the signed body bytes so the handler sends those
|
||||
exact bytes.
|
||||
"""
|
||||
return headers, None
|
||||
|
||||
@abstractmethod
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -1649,12 +1649,14 @@ class AmazonConverseConfig(BaseConfig):
|
|||
bedrock_tool_config["toolChoice"] = tool_choice_values
|
||||
|
||||
data: CommonRequestObject = {
|
||||
"additionalModelRequestFields": additional_request_params,
|
||||
"system": system_content_blocks,
|
||||
"inferenceConfig": self._transform_inference_params(
|
||||
inference_params=inference_params
|
||||
),
|
||||
}
|
||||
if additional_request_params:
|
||||
data["additionalModelRequestFields"] = additional_request_params
|
||||
if system_content_blocks:
|
||||
data["system"] = system_content_blocks
|
||||
|
||||
# Handle all config blocks
|
||||
for config_name, config_class in self.get_config_blocks().items():
|
||||
|
|
|
|||
|
|
@ -60,6 +60,9 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "bedrock"
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
return AnthropicConfig.get_supported_openai_params(self, model)
|
||||
|
||||
|
|
|
|||
|
|
@ -17,6 +17,9 @@ class BedrockClaudePlatformConfig(BedrockClaudePlatformMixin, AnthropicConfig):
|
|||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "bedrock"
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -4,14 +4,26 @@ Amazon Bedrock Mantle - Responses API backend.
|
|||
gpt-5.5 / gpt-5.4 on Mantle are exposed ONLY on the `/openai/v1/responses`
|
||||
path (not the standard `/v1/responses`). Payloads and SSE follow the OpenAI
|
||||
Responses spec, so this config inherits OpenAIResponsesAPIConfig and overrides
|
||||
only the endpoint URL and Bearer authentication.
|
||||
only the endpoint URL and authentication.
|
||||
|
||||
Auth: AWS Bedrock API key as Bearer token (BEDROCK_MANTLE_API_KEY or the
|
||||
standard AWS_BEARER_TOKEN_BEDROCK), NOT SigV4.
|
||||
Auth: Bearer token (BEDROCK_MANTLE_API_KEY or the standard
|
||||
AWS_BEARER_TOKEN_BEDROCK, or litellm_params.api_key) when present; otherwise
|
||||
AWS SigV4 (service name "bedrock") using the standard credential chain (IAM
|
||||
role / access key / profile / web identity), signed via the shared
|
||||
BaseAWSLLM._sign_request after the request body is finalized.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
import re
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from botocore.exceptions import (
|
||||
CredentialRetrievalError,
|
||||
NoCredentialsError,
|
||||
PartialCredentialsError,
|
||||
ProfileNotFound,
|
||||
)
|
||||
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -29,22 +41,44 @@ _BASE_SUFFIXES_TO_STRIP = (
|
|||
"/v1",
|
||||
)
|
||||
|
||||
# Standard Mantle host: https://bedrock-mantle.<region>.api.aws (group 1 = region).
|
||||
_MANTLE_HOST_RE = re.compile(
|
||||
r"^https?://bedrock-mantle\.([^/.]+)\.api\.aws", re.IGNORECASE
|
||||
)
|
||||
|
||||
|
||||
class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
def __init__(self, aws_signer: Optional[BaseAWSLLM] = None):
|
||||
super().__init__()
|
||||
self._aws_signer = aws_signer or BaseAWSLLM()
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> LlmProviders:
|
||||
return LlmProviders.BEDROCK_MANTLE
|
||||
|
||||
@staticmethod
|
||||
def _resolve_region(params: dict) -> str:
|
||||
region = params.get("aws_region_name")
|
||||
if region:
|
||||
return region
|
||||
base = params.get("api_base") or get_secret_str("BEDROCK_MANTLE_API_BASE")
|
||||
if base:
|
||||
match = _MANTLE_HOST_RE.match(base.rstrip("/"))
|
||||
if match:
|
||||
return match.group(1)
|
||||
return (
|
||||
get_secret_str("BEDROCK_MANTLE_REGION")
|
||||
or get_secret_str("AWS_REGION_NAME")
|
||||
or get_secret_str("AWS_REGION")
|
||||
or BEDROCK_MANTLE_DEFAULT_REGION
|
||||
)
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
litellm_params: dict,
|
||||
) -> str:
|
||||
region = (
|
||||
get_secret_str("BEDROCK_MANTLE_REGION")
|
||||
or get_secret_str("AWS_REGION")
|
||||
or BEDROCK_MANTLE_DEFAULT_REGION
|
||||
)
|
||||
region = self._resolve_region({**litellm_params, "api_base": api_base})
|
||||
base = (
|
||||
api_base
|
||||
or get_secret_str("BEDROCK_MANTLE_API_BASE")
|
||||
|
|
@ -55,6 +89,11 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
if base.endswith(suffix):
|
||||
base = base[: -len(suffix)]
|
||||
break
|
||||
# For the standard Mantle host (including the default-region base that
|
||||
# responses/main.py auto-injects into litellm_params.api_base), pin to the
|
||||
# single resolved region so aws_region_name wins; preserve custom proxy hosts.
|
||||
if _MANTLE_HOST_RE.match(base):
|
||||
base = f"https://bedrock-mantle.{region}.api.aws"
|
||||
return f"{base}/openai/v1/responses"
|
||||
|
||||
def validate_environment(
|
||||
|
|
@ -66,12 +105,8 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
or get_secret_str("BEDROCK_MANTLE_API_KEY")
|
||||
or get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
|
||||
)
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
"Bedrock Mantle API key is required. Set BEDROCK_MANTLE_API_KEY "
|
||||
"(or AWS_BEARER_TOKEN_BEDROCK) or pass api_key."
|
||||
)
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
return headers
|
||||
|
||||
def supports_native_file_search(self) -> bool:
|
||||
|
|
@ -79,3 +114,58 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
return False
|
||||
|
||||
def sign_request(
|
||||
self,
|
||||
headers: dict,
|
||||
optional_params: dict,
|
||||
request_data: dict,
|
||||
api_base: str,
|
||||
api_key: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
stream: Optional[bool] = None,
|
||||
fake_stream: Optional[bool] = None,
|
||||
) -> Tuple[dict, Optional[bytes]]:
|
||||
bearer = (
|
||||
api_key
|
||||
or get_secret_str("BEDROCK_MANTLE_API_KEY")
|
||||
or get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
|
||||
)
|
||||
if not bearer:
|
||||
# SigV4 path. Pin the credential-scope region to the region of the actual
|
||||
# signing URL (api_base, already region-resolved by get_complete_url) so the
|
||||
# SigV4 scope and the URL host can never disagree. Resolve from api_base first,
|
||||
# then fall back to the regular precedence. Also drop any caller Authorization
|
||||
# so _sign_request's restore-original-Authorization step cannot override the
|
||||
# SigV4 header.
|
||||
optional_params = {
|
||||
**optional_params,
|
||||
"aws_region_name": self._resolve_region(
|
||||
{**optional_params, "api_base": api_base}
|
||||
),
|
||||
}
|
||||
headers = {k: v for k, v in headers.items() if k.lower() != "authorization"}
|
||||
try:
|
||||
return self._aws_signer._sign_request(
|
||||
service_name="bedrock",
|
||||
headers=headers,
|
||||
optional_params=optional_params,
|
||||
request_data=request_data,
|
||||
api_base=api_base,
|
||||
api_key=bearer,
|
||||
model=model,
|
||||
stream=stream,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
except (
|
||||
NoCredentialsError,
|
||||
PartialCredentialsError,
|
||||
ProfileNotFound,
|
||||
CredentialRetrievalError,
|
||||
) as e:
|
||||
raise ValueError(
|
||||
"Bedrock Mantle auth failed: no Bearer token and no usable AWS "
|
||||
"credentials. Set BEDROCK_MANTLE_API_KEY (or AWS_BEARER_TOKEN_BEDROCK) "
|
||||
"or pass api_key for Bearer auth, or provide AWS credentials "
|
||||
"(IAM role / access key / profile / web identity) for SigV4."
|
||||
) from e
|
||||
|
|
|
|||
|
|
@ -2318,6 +2318,31 @@ class BaseLLMHTTPHandler:
|
|||
# but never included in the outbound provider payload.
|
||||
request_context["litellm_params"] = dict(litellm_params)
|
||||
|
||||
is_stream_request = bool(stream)
|
||||
if is_stream_request and fake_stream is True:
|
||||
stream, data = self._prepare_fake_stream_request(
|
||||
stream=stream,
|
||||
data=data,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
|
||||
# Sign after the body is final (post-transform/normalize/extra_body and post
|
||||
# fake-stream prep) so signed bytes match what we send. No-op for providers
|
||||
# that inherit the default sign_request.
|
||||
headers, signed_body = responses_api_provider_config.sign_request(
|
||||
headers=headers,
|
||||
optional_params=dict(litellm_params),
|
||||
request_data=data,
|
||||
api_base=api_base,
|
||||
api_key=litellm_params.api_key,
|
||||
model=model,
|
||||
stream=stream,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
body_kwargs: Dict[str, Any] = (
|
||||
{"data": signed_body} if signed_body is not None else {"json": data}
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
@ -2330,22 +2355,14 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
try:
|
||||
if stream:
|
||||
# For streaming, use stream=True in the request
|
||||
if fake_stream is True:
|
||||
stream, data = self._prepare_fake_stream_request(
|
||||
stream=stream,
|
||||
data=data,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
|
||||
if is_stream_request:
|
||||
response = sync_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout
|
||||
or float(response_api_optional_request_params.get("timeout", 0)),
|
||||
stream=stream,
|
||||
**body_kwargs,
|
||||
)
|
||||
if fake_stream is True:
|
||||
return MockResponsesAPIStreamingIterator(
|
||||
|
|
@ -2370,13 +2387,12 @@ class BaseLLMHTTPHandler:
|
|||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
else:
|
||||
# For non-streaming requests
|
||||
response = sync_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout
|
||||
or float(response_api_optional_request_params.get("timeout", 0)),
|
||||
**body_kwargs,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
|
|
@ -2464,6 +2480,28 @@ class BaseLLMHTTPHandler:
|
|||
# but never included in the outbound provider payload.
|
||||
request_context["litellm_params"] = dict(litellm_params)
|
||||
|
||||
is_stream_request = bool(stream)
|
||||
if is_stream_request and fake_stream is True:
|
||||
stream, data = self._prepare_fake_stream_request(
|
||||
stream=stream,
|
||||
data=data,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
|
||||
headers, signed_body = responses_api_provider_config.sign_request(
|
||||
headers=headers,
|
||||
optional_params=dict(litellm_params),
|
||||
request_data=data,
|
||||
api_base=api_base,
|
||||
api_key=litellm_params.api_key,
|
||||
model=model,
|
||||
stream=stream,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
body_kwargs: Dict[str, Any] = (
|
||||
{"data": signed_body} if signed_body is not None else {"json": data}
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
@ -2476,22 +2514,14 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
try:
|
||||
if stream:
|
||||
# For streaming, we need to use stream=True in the request
|
||||
if fake_stream is True:
|
||||
stream, data = self._prepare_fake_stream_request(
|
||||
stream=stream,
|
||||
data=data,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
|
||||
if is_stream_request:
|
||||
response = await async_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout
|
||||
or float(response_api_optional_request_params.get("timeout", 0)),
|
||||
stream=stream,
|
||||
**body_kwargs,
|
||||
)
|
||||
|
||||
if fake_stream is True:
|
||||
|
|
@ -2518,13 +2548,12 @@ class BaseLLMHTTPHandler:
|
|||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
else:
|
||||
# For non-streaming, proceed as before
|
||||
response = await async_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout
|
||||
or float(response_api_optional_request_params.get("timeout", 0)),
|
||||
**body_kwargs,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -4005,6 +4034,18 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
data = BaseResponsesAPIConfig.normalize_responses_api_request_dict(data)
|
||||
|
||||
headers, signed_body = responses_api_provider_config.sign_request(
|
||||
headers=headers,
|
||||
optional_params=dict(litellm_params),
|
||||
request_data=data,
|
||||
api_base=url,
|
||||
api_key=litellm_params.api_key,
|
||||
model=model,
|
||||
)
|
||||
body_kwargs: Dict[str, Any] = (
|
||||
{"data": signed_body} if signed_body is not None else {"json": data}
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
@ -4018,7 +4059,7 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
try:
|
||||
response = sync_httpx_client.post(
|
||||
url=url, headers=headers, json=data, timeout=timeout
|
||||
url=url, headers=headers, timeout=timeout, **body_kwargs
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -4088,6 +4129,18 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
data = BaseResponsesAPIConfig.normalize_responses_api_request_dict(data)
|
||||
|
||||
headers, signed_body = responses_api_provider_config.sign_request(
|
||||
headers=headers,
|
||||
optional_params=dict(litellm_params),
|
||||
request_data=data,
|
||||
api_base=url,
|
||||
api_key=litellm_params.api_key,
|
||||
model=model,
|
||||
)
|
||||
body_kwargs: Dict[str, Any] = (
|
||||
{"data": signed_body} if signed_body is not None else {"json": data}
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
@ -4101,7 +4154,7 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
try:
|
||||
response = await async_httpx_client.post(
|
||||
url=url, headers=headers, json=data, timeout=timeout
|
||||
url=url, headers=headers, timeout=timeout, **body_kwargs
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -26,6 +26,9 @@ class DeepSeekAnthropicMessagesConfig(AnthropicMessagesConfig):
|
|||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "deepseek"
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def get_api_key(api_key: Optional[str] = None) -> Optional[str]:
|
||||
return api_key or get_secret_str("DEEPSEEK_API_KEY") or litellm.api_key
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
GitHub Copilot Responses API Configuration.
|
||||
|
||||
This module provides the configuration for GitHub Copilot's Responses API,
|
||||
which is required for models like gpt-5.1-codex that only support the /responses endpoint.
|
||||
which is required for models like gpt-5.3-codex that only support the /responses endpoint.
|
||||
|
||||
Implementation based on analysis of the copilot-api project by caozhiyuan:
|
||||
https://github.com/caozhiyuan/copilot-api
|
||||
|
|
@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional, Union
|
|||
|
||||
import os
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
from litellm.exceptions import AuthenticationError
|
||||
|
|
@ -22,6 +23,7 @@ from litellm.types.llms.openai import (
|
|||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import _cached_get_model_info_helper
|
||||
|
||||
from ..authenticator import Authenticator
|
||||
from ..common_utils import (
|
||||
|
|
@ -38,6 +40,47 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
def github_copilot_supports_responses_api(model: str) -> bool:
|
||||
"""
|
||||
Gate native /v1/responses dispatch per github_copilot model.
|
||||
|
||||
Resolution (first match wins): mode "responses" -> True; mode "chat" ->
|
||||
False (opt-out wins for dual-endpoint models); "/v1/responses" in
|
||||
supported_endpoints -> True; else False. Unknown model -> False (the bridge
|
||||
always works since every Copilot model supports /chat/completions).
|
||||
|
||||
Reads merged model info (per-deployment model_info applied via the router's
|
||||
register_model, which also clears the cache used here).
|
||||
"""
|
||||
try:
|
||||
info = _cached_get_model_info_helper(
|
||||
model=model, custom_llm_provider="github_copilot"
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"github_copilot_supports_responses_api: get_model_info failed "
|
||||
"for %s: %s",
|
||||
model,
|
||||
e,
|
||||
)
|
||||
return False
|
||||
|
||||
mode = info.get("mode")
|
||||
if mode == "responses":
|
||||
return True
|
||||
if mode == "chat":
|
||||
return False
|
||||
|
||||
# supported_endpoints is dropped by ModelInfoBase; read it from the raw
|
||||
# model_cost entry via the resolved key.
|
||||
key = info.get("key")
|
||||
raw_info = litellm.model_cost.get(key) if isinstance(key, str) else None
|
||||
endpoints = (
|
||||
raw_info.get("supported_endpoints") if isinstance(raw_info, dict) else None
|
||||
)
|
||||
return isinstance(endpoints, list) and "/v1/responses" in endpoints
|
||||
|
||||
|
||||
class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
"""
|
||||
Configuration for GitHub Copilot's Responses API.
|
||||
|
|
|
|||
|
|
@ -28,6 +28,9 @@ class MinimaxMessagesConfig(AnthropicMessagesConfig):
|
|||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "minimax"
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def get_api_key(api_key: Optional[str] = None) -> Optional[str]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -187,6 +187,7 @@ def create_responses_config_class(provider: SimpleProviderConfig):
|
|||
from litellm.llms.openai_like.responses.transformation import (
|
||||
OpenAILikeResponsesConfig,
|
||||
)
|
||||
from litellm.types.llms.openai import ResponseInputParam
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
class JSONProviderResponsesConfig(OpenAILikeResponsesConfig):
|
||||
|
|
@ -223,5 +224,23 @@ def create_responses_config_class(provider: SimpleProviderConfig):
|
|||
api_base = api_base.rstrip("/")
|
||||
return f"{api_base}/responses"
|
||||
|
||||
def transform_responses_api_request(
|
||||
self,
|
||||
model: str,
|
||||
input: Union[str, ResponseInputParam],
|
||||
response_api_optional_request_params: dict,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
if provider.special_handling.get("force_store_false"):
|
||||
response_api_optional_request_params["store"] = False
|
||||
return super().transform_responses_api_request(
|
||||
model=model,
|
||||
input=input,
|
||||
response_api_optional_request_params=response_api_optional_request_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
_responses_config_cache[provider.slug] = JSONProviderResponsesConfig
|
||||
return JSONProviderResponsesConfig
|
||||
|
|
|
|||
|
|
@ -132,5 +132,14 @@
|
|||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
}
|
||||
},
|
||||
"parasail": {
|
||||
"base_url": "https://api.parasail.io/v1",
|
||||
"api_key_env": "PARASAIL_API_KEY",
|
||||
"api_base_env": "PARASAIL_API_BASE",
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"special_handling": {
|
||||
"force_store_false": true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -12,7 +12,11 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_defs
|
|||
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.llms.vertex_ai import PartType, Schema
|
||||
from litellm.types.llms.vertex_ai import (
|
||||
VERTEX_AI_PROVIDER_METADATA_FIELDS,
|
||||
PartType,
|
||||
Schema,
|
||||
)
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
from litellm.utils import supports_response_schema, supports_system_messages
|
||||
|
||||
|
|
@ -27,6 +31,47 @@ class VertexAIError(BaseLLMException):
|
|||
super().__init__(message=message, status_code=status_code, headers=headers)
|
||||
|
||||
|
||||
def redact_vertex_ai_metadata_from_logged_object(obj: Any) -> None:
|
||||
if isinstance(obj, dict):
|
||||
for field in VERTEX_AI_PROVIDER_METADATA_FIELDS:
|
||||
if field in obj:
|
||||
obj[field] = []
|
||||
hidden_params = obj.get("_hidden_params")
|
||||
if isinstance(hidden_params, dict):
|
||||
for field in VERTEX_AI_PROVIDER_METADATA_FIELDS:
|
||||
hidden_params.pop(field, None)
|
||||
return
|
||||
|
||||
for field in VERTEX_AI_PROVIDER_METADATA_FIELDS:
|
||||
if hasattr(obj, field):
|
||||
setattr(obj, field, [])
|
||||
hidden_params = getattr(obj, "_hidden_params", None)
|
||||
if isinstance(hidden_params, dict):
|
||||
for field in VERTEX_AI_PROVIDER_METADATA_FIELDS:
|
||||
hidden_params.pop(field, None)
|
||||
|
||||
|
||||
def redact_vertex_ai_metadata_from_litellm_params(model_call_details: dict) -> None:
|
||||
"""
|
||||
success_handler() merges response._hidden_params into
|
||||
litellm_params.metadata['hidden_params'] before redaction runs, so the Vertex
|
||||
metadata must be scrubbed from that copy too.
|
||||
"""
|
||||
litellm_params = model_call_details.get("litellm_params")
|
||||
if not isinstance(litellm_params, dict):
|
||||
return
|
||||
|
||||
for metadata_key in ("metadata", "litellm_metadata"):
|
||||
metadata = litellm_params.get(metadata_key)
|
||||
if not isinstance(metadata, dict):
|
||||
continue
|
||||
hidden_params = metadata.get("hidden_params")
|
||||
if not isinstance(hidden_params, dict):
|
||||
continue
|
||||
for field in VERTEX_AI_PROVIDER_METADATA_FIELDS:
|
||||
hidden_params.pop(field, None)
|
||||
|
||||
|
||||
def vertex_request_labels_from_litellm_params(
|
||||
litellm_params: Optional[dict],
|
||||
) -> Optional[Dict[str, str]]:
|
||||
|
|
|
|||
|
|
@ -63,6 +63,7 @@ from litellm.types.llms.openai import (
|
|||
OpenAIChatCompletionFinishReason,
|
||||
)
|
||||
from litellm.types.llms.vertex_ai import (
|
||||
VERTEX_AI_PROVIDER_METADATA_FIELDS,
|
||||
VERTEX_CREDENTIALS_TYPES,
|
||||
Candidates,
|
||||
ContentType,
|
||||
|
|
@ -1111,6 +1112,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
{
|
||||
"voice": "alloy",
|
||||
"format": "mp3",
|
||||
"language_code": "en-US",
|
||||
}
|
||||
|
||||
Expected output:
|
||||
|
|
@ -1119,7 +1121,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
prebuiltVoiceConfig: {
|
||||
voiceName: "alloy",
|
||||
}
|
||||
}
|
||||
},
|
||||
languageCode: "en-US",
|
||||
}
|
||||
"""
|
||||
from litellm.types.llms.vertex_ai import (
|
||||
|
|
@ -1145,6 +1148,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
voice_config: VoiceConfig = {"prebuiltVoiceConfig": prebuilt_voice_config}
|
||||
speech_config["voiceConfig"] = voice_config
|
||||
|
||||
if "language_code" in value:
|
||||
speech_config["languageCode"] = value["language_code"]
|
||||
|
||||
return cast(dict, speech_config)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -2253,6 +2259,71 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
citation_metadata,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_stream_chunk_attr(chunk: Any, field_name: str) -> Any:
|
||||
if isinstance(chunk, dict):
|
||||
value = chunk.get(field_name)
|
||||
if value is not None:
|
||||
return value
|
||||
model_extra = chunk.get("model_extra")
|
||||
if isinstance(model_extra, dict):
|
||||
value = model_extra.get(field_name)
|
||||
if value is not None:
|
||||
return value
|
||||
hidden_params = chunk.get("_hidden_params")
|
||||
if isinstance(hidden_params, dict):
|
||||
return hidden_params.get(field_name)
|
||||
return None
|
||||
return getattr(chunk, field_name, None)
|
||||
|
||||
@staticmethod
|
||||
def _set_stream_metadata_on_response(
|
||||
model_response: Any,
|
||||
grounding_metadata: List[dict],
|
||||
url_context_metadata: List[dict],
|
||||
safety_ratings: List[dict],
|
||||
citation_metadata: List[dict],
|
||||
) -> None:
|
||||
setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore
|
||||
if grounding_metadata:
|
||||
model_response._hidden_params["vertex_ai_grounding_metadata"] = (
|
||||
grounding_metadata
|
||||
)
|
||||
setattr(model_response, "vertex_ai_url_context_metadata", url_context_metadata) # type: ignore
|
||||
if url_context_metadata:
|
||||
model_response._hidden_params["vertex_ai_url_context_metadata"] = (
|
||||
url_context_metadata
|
||||
)
|
||||
setattr(model_response, "vertex_ai_safety_ratings", safety_ratings) # type: ignore
|
||||
setattr(model_response, "vertex_ai_safety_results", safety_ratings) # type: ignore
|
||||
if safety_ratings:
|
||||
model_response._hidden_params["vertex_ai_safety_ratings"] = safety_ratings
|
||||
model_response._hidden_params["vertex_ai_safety_results"] = safety_ratings
|
||||
setattr(model_response, "vertex_ai_citation_metadata", citation_metadata) # type: ignore
|
||||
if citation_metadata:
|
||||
model_response._hidden_params["vertex_ai_citation_metadata"] = (
|
||||
citation_metadata
|
||||
)
|
||||
|
||||
def apply_assembled_streaming_response_metadata(
|
||||
self,
|
||||
response: ModelResponse,
|
||||
chunks: List[Any],
|
||||
) -> None:
|
||||
for field_name in VERTEX_AI_PROVIDER_METADATA_FIELDS:
|
||||
merged: List[Any] = []
|
||||
for chunk in chunks:
|
||||
value = VertexGeminiConfig._get_stream_chunk_attr(chunk, field_name)
|
||||
if not value:
|
||||
continue
|
||||
if isinstance(value, list):
|
||||
merged.extend(value)
|
||||
else:
|
||||
merged.append(value)
|
||||
if merged:
|
||||
setattr(response, field_name, merged)
|
||||
response._hidden_params[field_name] = merged
|
||||
|
||||
@staticmethod
|
||||
def _convert_grounding_metadata_to_annotations(
|
||||
grounding_metadata: List[dict],
|
||||
|
|
@ -3385,10 +3456,13 @@ class ModelResponseIterator:
|
|||
if choice.finish_reason == "stop":
|
||||
choice.finish_reason = "tool_calls"
|
||||
|
||||
setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore
|
||||
setattr(model_response, "vertex_ai_url_context_metadata", url_context_metadata) # type: ignore
|
||||
setattr(model_response, "vertex_ai_safety_ratings", safety_ratings) # type: ignore
|
||||
setattr(model_response, "vertex_ai_citation_metadata", citation_metadata) # type: ignore
|
||||
VertexGeminiConfig._set_stream_metadata_on_response(
|
||||
model_response,
|
||||
grounding_metadata,
|
||||
url_context_metadata,
|
||||
safety_ratings,
|
||||
citation_metadata,
|
||||
)
|
||||
|
||||
return (
|
||||
grounding_metadata,
|
||||
|
|
|
|||
|
|
@ -17,6 +17,9 @@ from ..output_params_utils import sanitize_vertex_anthropic_output_params
|
|||
|
||||
|
||||
class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, VertexBase):
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
def validate_anthropic_messages_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -52,6 +52,9 @@ class VertexAIAnthropicConfig(AnthropicConfig):
|
|||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "vertex_ai"
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
def _add_context_management_beta_headers(
|
||||
self, beta_set: set, context_management: dict
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -7761,6 +7761,9 @@ def stream_chunk_builder( # noqa: PLR0915
|
|||
"cost",
|
||||
logging_obj._response_cost_calculator(result=response),
|
||||
)
|
||||
processor.apply_provider_assembled_streaming_metadata(
|
||||
response, chunks, logging_obj
|
||||
)
|
||||
return response
|
||||
|
||||
tool_call_chunks = [
|
||||
|
|
@ -7940,6 +7943,9 @@ def stream_chunk_builder( # noqa: PLR0915
|
|||
usage, "cost", logging_obj._response_cost_calculator(result=response)
|
||||
)
|
||||
|
||||
processor.apply_provider_assembled_streaming_metadata(
|
||||
response, chunks, logging_obj
|
||||
)
|
||||
return response
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
|
|
|
|||
|
|
@ -6889,6 +6889,43 @@
|
|||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"azure_ai/MAI-Image-2.5": {
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.05,
|
||||
"output_cost_per_image_token": 4.7e-05,
|
||||
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/new-mai-models-in-microsoft-foundry-across-text-image-voice-and-speech/4524632",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
]
|
||||
},
|
||||
"azure_ai/MAI-Image-2.5-Flash": {
|
||||
"input_cost_per_image_token": 1.75e-06,
|
||||
"input_cost_per_token": 1.75e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0338,
|
||||
"output_cost_per_image_token": 3.3e-05,
|
||||
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/new-mai-models-in-microsoft-foundry-across-text-image-voice-and-speech/4524632",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
]
|
||||
},
|
||||
"azure_ai/MAI-Image-2e": {
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.02,
|
||||
"output_cost_per_image_token": 1.95e-05,
|
||||
"source": "https://aka.ms/mai-image-2e-foundryblog",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"azure_ai/Llama-3.2-11B-Vision-Instruct": {
|
||||
"input_cost_per_token": 3.7e-07,
|
||||
"litellm_provider": "azure_ai",
|
||||
|
|
@ -14286,10 +14323,10 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 4.4e-06,
|
||||
"source": "https://fireworks.ai/models/fireworks/glm-5p1",
|
||||
"supports_function_calling": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": false,
|
||||
"supports_tool_choice": false
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/gpt-oss-120b": {
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
|
|
@ -14567,10 +14604,10 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 4.4e-06,
|
||||
"source": "https://fireworks.ai/models/fireworks/glm-5p1",
|
||||
"supports_function_calling": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": false,
|
||||
"supports_tool_choice": false
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"fireworks_ai/kimi-k2p5": {
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
|
|
|
|||
|
|
@ -512,12 +512,13 @@ async def exchange_token_with_server(
|
|||
result = {
|
||||
"access_token": access_token,
|
||||
"token_type": token_response.get("token_type", "Bearer"),
|
||||
"expires_in": token_response.get("expires_in", 3600),
|
||||
}
|
||||
|
||||
if "refresh_token" in token_response and token_response["refresh_token"]:
|
||||
if token_response.get("expires_in") is not None:
|
||||
result["expires_in"] = token_response["expires_in"]
|
||||
if token_response.get("refresh_token"):
|
||||
result["refresh_token"] = token_response["refresh_token"]
|
||||
if "scope" in token_response and token_response["scope"]:
|
||||
if token_response.get("scope"):
|
||||
result["scope"] = token_response["scope"]
|
||||
|
||||
# RFC 6749 §5.1: token responses must not be cached.
|
||||
|
|
|
|||
|
|
@ -386,8 +386,15 @@ if MCP_AVAILABLE:
|
|||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
apply_tool_filters: bool = True,
|
||||
):
|
||||
"""Helper function to get tools for a single server."""
|
||||
"""Helper function to get tools for a single server.
|
||||
|
||||
When ``apply_tool_filters`` is False the raw server catalog is returned
|
||||
without the allowed_tools/disallowed_tools gate or the per-key tool
|
||||
permissions. This is the admin-only configuration view; every runtime
|
||||
path keeps the default True so callable tools stay filtered.
|
||||
"""
|
||||
tools = await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
|
|
@ -397,6 +404,9 @@ if MCP_AVAILABLE:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
if not apply_tool_filters:
|
||||
return _create_tool_response_objects(tools, server.mcp_info)
|
||||
|
||||
# Always apply allowed_tools/disallowed_tools so the blacklist is
|
||||
# enforced even when no allowlist is set (matches the SSE/HTTP path).
|
||||
tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
|
@ -463,6 +473,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header: Optional[str],
|
||||
raw_headers_from_request: dict,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
apply_tool_filters: bool = True,
|
||||
) -> dict:
|
||||
"""Handle tool listing for a single server_id request."""
|
||||
# Resolve a server name to its UUID if needed
|
||||
|
|
@ -527,6 +538,7 @@ if MCP_AVAILABLE:
|
|||
raw_headers_from_request,
|
||||
user_api_key_dict,
|
||||
extra_headers=user_oauth_extra_headers,
|
||||
apply_tool_filters=apply_tool_filters,
|
||||
)
|
||||
except MCPUpstreamAuthError:
|
||||
# Surface the upstream 401/403 to the caller so it can emit the
|
||||
|
|
@ -552,6 +564,14 @@ if MCP_AVAILABLE:
|
|||
server_id: Optional[str] = Query(
|
||||
None, description="The server id to list tools for"
|
||||
),
|
||||
include_disabled_tools: bool = Query(
|
||||
False,
|
||||
description=(
|
||||
"Admin only. Return the full server tool catalog without the "
|
||||
"allowed_tools filter or per-key tool permissions, so the MCP "
|
||||
"settings UI can configure the allowlist. Ignored for non-admins."
|
||||
),
|
||||
),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
) -> dict:
|
||||
"""
|
||||
|
|
@ -579,6 +599,13 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
|
||||
try:
|
||||
# The full catalog (allowlist filter skipped) is admin-only so the
|
||||
# REST endpoint can't be used to enumerate deliberately-disabled tools.
|
||||
apply_tool_filters = not (
|
||||
include_disabled_tools
|
||||
and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
# Extract auth headers from request
|
||||
headers = request.headers
|
||||
raw_headers_from_request = dict(headers)
|
||||
|
|
@ -620,6 +647,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header=mcp_auth_header,
|
||||
raw_headers_from_request=raw_headers_from_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
apply_tool_filters=apply_tool_filters,
|
||||
)
|
||||
else:
|
||||
if not allowed_server_ids:
|
||||
|
|
@ -677,6 +705,7 @@ if MCP_AVAILABLE:
|
|||
raw_headers_from_request,
|
||||
user_api_key_dict,
|
||||
extra_headers=user_oauth_extra_headers,
|
||||
apply_tool_filters=apply_tool_filters,
|
||||
)
|
||||
list_tools_result.extend(tools_result)
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -2312,6 +2312,24 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="List of MCP server fields that must be filled in for a submission to pass standards checks (e.g. ['description', 'source_url', 'alias']).",
|
||||
)
|
||||
disable_budget_reservation: Optional[bool] = Field(
|
||||
None,
|
||||
description=(
|
||||
"If True, disables the optimistic per-request budget reservation "
|
||||
"introduced in v1.84.0. "
|
||||
"WARNING: This weakens hard budget enforcement. Without the reservation, "
|
||||
"a burst of concurrent requests from a single key can each pass the "
|
||||
"read-time spend check before any of them is charged, allowing a "
|
||||
"configured budget to be exceeded under high concurrency. "
|
||||
"Budgets are still evaluated on every request at read time, so "
|
||||
"an already-exhausted budget is still rejected. "
|
||||
"Enable only if your deployment is experiencing phantom "
|
||||
"BudgetExceededError responses caused by leaked reservations "
|
||||
"(see GitHub issue #27639). "
|
||||
"A proxy-level WARNING is logged on every request while this flag "
|
||||
"is active as a reminder that hard enforcement is relaxed."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class ConfigYAML(LiteLLMPydanticObjectBase):
|
||||
|
|
@ -4177,6 +4195,16 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
default=None,
|
||||
description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.",
|
||||
)
|
||||
team_claim_fallback: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"If True, when a configured team_id_jwt_field / team_ids_jwt_field "
|
||||
"claim is present but does not resolve to any known team, defer to "
|
||||
"the single-team DB fallback (caller's only team membership) "
|
||||
"instead of raising. Default False preserves strict claim-based "
|
||||
"authorization."
|
||||
),
|
||||
)
|
||||
issuers: Optional[List[JWTIssuerConfig]] = Field(
|
||||
default=None,
|
||||
description="Optional issuer-bound JWT validation rules. When a token's `iss` matches a configured issuer, validation uses that issuer's JWKS, audience, and claim mappings. Tokens with an unlisted `iss` fall back to the global JWT_AUDIENCE/JWT_ISSUER validation path — this is additive routing, not an allow-list.",
|
||||
|
|
|
|||
|
|
@ -1299,15 +1299,29 @@ class JWTAuthManager:
|
|||
|
||||
# First try to get team by team_id
|
||||
if individual_team_id:
|
||||
team_object = await get_team_object(
|
||||
team_id=individual_team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert,
|
||||
)
|
||||
return individual_team_id, team_object
|
||||
try:
|
||||
team_object = await get_team_object(
|
||||
team_id=individual_team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert,
|
||||
)
|
||||
return individual_team_id, team_object
|
||||
except HTTPException as e:
|
||||
if (
|
||||
e.status_code != 404
|
||||
or not jwt_handler.litellm_jwtauth.team_claim_fallback
|
||||
):
|
||||
raise
|
||||
# Claim doesn't map to a known team — defer to fallback.
|
||||
verbose_proxy_logger.debug(
|
||||
"JWT team_id claim '%s' did not resolve to a team: %s",
|
||||
individual_team_id,
|
||||
e.detail,
|
||||
)
|
||||
return None, None
|
||||
|
||||
# If no team_id found, try to resolve via team_alias_jwt_field
|
||||
team_alias = jwt_handler.get_team_alias(
|
||||
|
|
@ -1431,6 +1445,7 @@ class JWTAuthManager:
|
|||
)
|
||||
return None, None
|
||||
|
||||
any_claim_team_resolved = False
|
||||
for team_id in team_ids:
|
||||
try:
|
||||
team_object = await get_team_object(
|
||||
|
|
@ -1441,6 +1456,9 @@ class JWTAuthManager:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
if team_object is not None:
|
||||
any_claim_team_resolved = True
|
||||
|
||||
if team_object and team_object.models is not None:
|
||||
team_models = team_object.models
|
||||
if isinstance(team_models, list) and (
|
||||
|
|
@ -1478,12 +1496,17 @@ class JWTAuthManager:
|
|||
if denied_auth_enforced_pass_through_route:
|
||||
JWTAuthManager._raise_team_passthrough_route_denial(route=route)
|
||||
|
||||
if requested_model:
|
||||
if requested_model and (
|
||||
any_claim_team_resolved
|
||||
or not jwt_handler.litellm_jwtauth.team_claim_fallback
|
||||
):
|
||||
# Claim resolved but no model access, or fallback disabled — deny.
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"No team has access to the requested model: {requested_model}. Checked teams={team_ids}. Check `/models` to see all available models.",
|
||||
)
|
||||
|
||||
# No claim team resolved and fallback enabled — defer to fallback.
|
||||
return None, None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -2425,6 +2425,7 @@ async def _run_centralized_common_checks( # noqa: PLR0915
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
skip_budget_checks=skip_budget_checks,
|
||||
general_settings=general_settings,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -2445,12 +2446,23 @@ async def _reserve_budget_after_common_checks(
|
|||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
skip_budget_checks: bool,
|
||||
general_settings: dict,
|
||||
end_user_id: Optional[str] = None,
|
||||
end_user_object: Optional[LiteLLM_EndUserTable] = None,
|
||||
) -> None:
|
||||
user_api_key_auth_obj.budget_reservation = None
|
||||
if skip_budget_checks:
|
||||
return
|
||||
if general_settings.get("disable_budget_reservation") is True:
|
||||
verbose_proxy_logger.warning(
|
||||
"disable_budget_reservation is enabled: skipping optimistic budget "
|
||||
"reservation. Budget enforcement is read-time only — concurrent "
|
||||
"requests can each pass the spend check before their cost is recorded, "
|
||||
"so a configured budget may be briefly exceeded under high concurrency. "
|
||||
"Set disable_budget_reservation to False or remove it to restore "
|
||||
"hard per-request budget enforcement."
|
||||
)
|
||||
return
|
||||
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
reserve_budget_for_request,
|
||||
|
|
|
|||
|
|
@ -105,6 +105,16 @@ def _extract_text_from_content(content: object) -> str:
|
|||
return ""
|
||||
|
||||
|
||||
def _merge_metadata_bags(request_data: Mapping[str, Any]) -> Optional[dict[str, Any]]:
|
||||
merged: dict[str, Any] = {}
|
||||
present = False
|
||||
for bag in (request_data.get("metadata"), request_data.get("litellm_metadata")):
|
||||
if isinstance(bag, Mapping):
|
||||
present = True
|
||||
merged.update(bag)
|
||||
return merged if present else None
|
||||
|
||||
|
||||
class CrowdStrikeAIDRHandler(CustomGuardrail):
|
||||
"""
|
||||
CrowdStrike AIDR AI Guardrail handler to interact with the CrowdStrike AIDR
|
||||
|
|
@ -312,11 +322,27 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
event_type = "output"
|
||||
hook_name = "apply_guardrail (response)"
|
||||
|
||||
ai_guard_payload = {
|
||||
ai_guard_payload: dict[str, Any] = {
|
||||
"guard_input": guard_input.model_dump(mode="json"),
|
||||
"event_type": event_type,
|
||||
}
|
||||
|
||||
model = inputs.get("model")
|
||||
if model:
|
||||
ai_guard_payload["model"] = model
|
||||
|
||||
metadata = _merge_metadata_bags(request_data)
|
||||
if metadata is not None:
|
||||
user_id = metadata.get("user_api_key_user_id")
|
||||
if user_id:
|
||||
ai_guard_payload["user_id"] = user_id
|
||||
|
||||
extra_info: dict[str, str] = {}
|
||||
user_email = metadata.get("user_api_key_user_email")
|
||||
if user_email:
|
||||
extra_info["user_name"] = user_email
|
||||
ai_guard_payload["extra_info"] = extra_info
|
||||
|
||||
ai_guard_response = await self._call_crowdstrike_aidr_guard(
|
||||
ai_guard_payload, hook_name
|
||||
)
|
||||
|
|
|
|||
|
|
@ -129,6 +129,7 @@ services = Union[
|
|||
"datadog_llm_observability",
|
||||
"generic_api",
|
||||
"arize",
|
||||
"galileo",
|
||||
"sqs",
|
||||
],
|
||||
str,
|
||||
|
|
@ -206,6 +207,7 @@ async def health_services_endpoint( # noqa: PLR0915
|
|||
"datadog_llm_observability",
|
||||
"generic_api",
|
||||
"arize",
|
||||
"galileo",
|
||||
"sqs",
|
||||
]:
|
||||
raise HTTPException(
|
||||
|
|
@ -295,6 +297,19 @@ async def health_services_endpoint( # noqa: PLR0915
|
|||
else "Arize is healthy"
|
||||
),
|
||||
}
|
||||
elif service == "galileo":
|
||||
from litellm.integrations.galileo import GalileoObserve
|
||||
|
||||
galileo_logger = GalileoObserve()
|
||||
response = await galileo_logger.async_health_check()
|
||||
return {
|
||||
"status": response["status"],
|
||||
"message": (
|
||||
response["error_message"]
|
||||
if response["status"] == "unhealthy"
|
||||
else "Galileo is healthy"
|
||||
),
|
||||
}
|
||||
elif service == "langfuse":
|
||||
from litellm.integrations.langfuse.langfuse import LangFuseLogger
|
||||
|
||||
|
|
|
|||
|
|
@ -561,7 +561,7 @@ async def _setup_new_team_model_assignment(
|
|||
|
||||
|
||||
async def _get_team_deployments(
|
||||
team_id: str, prisma_client: PrismaClient
|
||||
team_id: str, prisma_client: PrismaClient, table: Optional[Any] = None
|
||||
) -> List[LiteLLM_ProxyModelTable]:
|
||||
"""
|
||||
Fetch all deployments for a given team_id from the database.
|
||||
|
|
@ -572,9 +572,13 @@ async def _get_team_deployments(
|
|||
Note: prisma-client-py 0.11.0 does not support JSON path filtering, so we filter
|
||||
by the model_name prefix (team models use "model_name_{team_id}_*") and confirm
|
||||
team_id in model_info with Python-side filtering.
|
||||
|
||||
Pass ``table`` (a transaction's proxy-model table) to run the read inside an
|
||||
existing transaction.
|
||||
"""
|
||||
prefix = f"model_name_{team_id}_"
|
||||
response = await ModelRepository(prisma_client).table.find_many(
|
||||
table = table or ModelRepository(prisma_client).table
|
||||
response = await table.find_many(
|
||||
where={
|
||||
"model_name": {"startswith": prefix},
|
||||
}
|
||||
|
|
@ -596,6 +600,42 @@ async def _get_team_deployments(
|
|||
return result
|
||||
|
||||
|
||||
async def delete_team_models(
|
||||
team_ids: List[str],
|
||||
prisma_client: PrismaClient,
|
||||
llm_router: Optional[Any],
|
||||
) -> List[str]:
|
||||
"""
|
||||
Delete every BYOK model owned by the given teams, from the DB and the router.
|
||||
|
||||
The DB rows are removed inside a single transaction, so deletion is atomic
|
||||
across all team_ids. Each team's rows are deleted by the exact model_ids read
|
||||
in the same transaction, which keeps the deleted set identical to the set
|
||||
handed to the router. The router is synced only after the transaction commits,
|
||||
so a rollback can never leave a deployment live in the router without its row.
|
||||
|
||||
Returns the model_ids that were deleted.
|
||||
"""
|
||||
deleted_model_ids: List[str] = []
|
||||
async with prisma_client.db.tx() as tx:
|
||||
for team_id in team_ids:
|
||||
rows = await _get_team_deployments(
|
||||
team_id, prisma_client, table=tx.litellm_proxymodeltable
|
||||
)
|
||||
model_ids = [row.model_id for row in rows]
|
||||
if model_ids:
|
||||
await tx.litellm_proxymodeltable.delete_many(
|
||||
where={"model_id": {"in": model_ids}}
|
||||
)
|
||||
deleted_model_ids.extend(model_ids)
|
||||
|
||||
if llm_router is not None:
|
||||
for model_id in deleted_model_ids:
|
||||
llm_router.delete_deployment(id=model_id)
|
||||
|
||||
return deleted_model_ids
|
||||
|
||||
|
||||
async def _get_team_public_model_names(
|
||||
team_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
|
|
@ -860,6 +900,7 @@ class ModelManagementAuthChecks:
|
|||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient,
|
||||
premium_user: bool,
|
||||
allow_missing_team: bool = False,
|
||||
) -> Literal[True]:
|
||||
## Check team model auth
|
||||
if (
|
||||
|
|
@ -870,6 +911,18 @@ class ModelManagementAuthChecks:
|
|||
where={"team_id": model_params.model_info.team_id}
|
||||
)
|
||||
if team_obj_row is None:
|
||||
# The team was deleted. Callers that opt in (e.g. model deletion) may
|
||||
# act on the orphaned model, but only as a proxy admin -- without the
|
||||
# team there is no team-admin membership left to verify.
|
||||
if allow_missing_team:
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
return True
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Only a proxy admin can delete a model whose team has been deleted."
|
||||
},
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
|
|
@ -955,6 +1008,7 @@ async def delete_model(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
premium_user=premium_user,
|
||||
allow_missing_team=True,
|
||||
)
|
||||
|
||||
# update DB
|
||||
|
|
|
|||
|
|
@ -779,6 +779,7 @@ async def _check_user_team_limits(
|
|||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: Any,
|
||||
existing_team_max_budget: Optional[float] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Check user team limits for standalone teams (not org-scoped).
|
||||
|
|
@ -789,28 +790,44 @@ async def _check_user_team_limits(
|
|||
|
||||
Should only be called for standalone teams (when organization_id is None).
|
||||
For org-scoped teams, use _check_org_team_limits() instead.
|
||||
|
||||
`existing_team_max_budget` is the team's current `max_budget` on the
|
||||
/team/update path. When the incoming `max_budget` is unchanged or lower
|
||||
than the team's current budget, the personal-budget comparison is skipped
|
||||
so a team admin can edit other fields (e.g. tpm_limit, team name) without
|
||||
being blocked by a budget the team already has. The UI sends the full team
|
||||
object on every update, so the unchanged `max_budget` would otherwise fail.
|
||||
"""
|
||||
# Validate team budget against user's max_budget
|
||||
if data.max_budget is not None and user_api_key_dict.user_id is not None:
|
||||
user_obj = await get_user_object(
|
||||
user_id=user_api_key_dict.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
# On /team/update, allow unchanged or lower budgets without checking
|
||||
# the caller's personal max_budget. Only increases above the team's
|
||||
# current budget are validated against the user's personal limit.
|
||||
budget_unchanged_or_lower = (
|
||||
existing_team_max_budget is not None
|
||||
and data.max_budget <= existing_team_max_budget
|
||||
)
|
||||
|
||||
if (
|
||||
user_obj is not None
|
||||
and user_obj.max_budget is not None
|
||||
and data.max_budget > user_obj.max_budget
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"max budget higher than user max. User max budget={user_obj.max_budget}. User role={user_api_key_dict.user_role}"
|
||||
},
|
||||
if not budget_unchanged_or_lower:
|
||||
user_obj = await get_user_object(
|
||||
user_id=user_api_key_dict.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
)
|
||||
|
||||
if (
|
||||
user_obj is not None
|
||||
and user_obj.max_budget is not None
|
||||
and data.max_budget > user_obj.max_budget
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"max budget higher than user max. User max budget={user_obj.max_budget}. User role={user_api_key_dict.user_role}"
|
||||
},
|
||||
)
|
||||
|
||||
# Validate team models against user's allowed models
|
||||
if data.models is not None and len(user_api_key_dict.models) > 0:
|
||||
for m in data.models:
|
||||
|
|
@ -1824,6 +1841,7 @@ async def update_team( # noqa: PLR0915
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
existing_team_max_budget=existing_team_row.max_budget,
|
||||
)
|
||||
|
||||
updated_kv = data.json(exclude_unset=True)
|
||||
|
|
@ -3271,6 +3289,20 @@ async def delete_team(
|
|||
|
||||
await prisma_client.delete_data(team_id_list=data.team_ids, table_name="key")
|
||||
|
||||
## DELETE ASSOCIATED BYOK MODELS
|
||||
# Runs before the team rows are deleted so a mid-flight failure never leaves
|
||||
# the team gone with its models orphaned.
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
delete_team_models,
|
||||
)
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
await delete_team_models(
|
||||
team_ids=data.team_ids,
|
||||
prisma_client=prisma_client,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
# ## DELETE TEAM MEMBERSHIPS
|
||||
for team_row in team_rows:
|
||||
### get all team members
|
||||
|
|
|
|||
|
|
@ -1220,7 +1220,7 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]:
|
|||
generic_role_mappings_group_claim = os.getenv(
|
||||
"GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", None
|
||||
)
|
||||
generic_role_mappoings_default_role = os.getenv(
|
||||
generic_role_mappings_default_role = os.getenv(
|
||||
"GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", None
|
||||
)
|
||||
if generic_role_mappings is not None:
|
||||
|
|
@ -1239,7 +1239,7 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]:
|
|||
role_mappings_data = {
|
||||
"provider": "generic",
|
||||
"group_claim": generic_role_mappings_group_claim,
|
||||
"default_role": generic_role_mappoings_default_role,
|
||||
"default_role": generic_role_mappings_default_role,
|
||||
"roles": generic_user_role_mappings_data,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -32,6 +32,71 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
|||
from litellm.types.utils import ImageResponse, LlmProviders, PassthroughCallTypes
|
||||
from litellm.utils import ModelResponse, TextCompletionResponse
|
||||
|
||||
# Hostnames that route to OpenAI-compatible APIs.
|
||||
#
|
||||
# `api.openai.com` is OpenAI proper. The two Azure domains below are *shared by
|
||||
# every Azure Cognitive Service* (Speech, Vision, Language, ...), not just Azure
|
||||
# OpenAI: `openai.azure.com` is the classic Azure OpenAI domain, while
|
||||
# `cognitiveservices.azure.com` is used by newer "Azure AI Foundry" /
|
||||
# Cognitive Services-hosted Azure OpenAI deployments. Because the hostname alone
|
||||
# cannot tell Azure OpenAI apart from the other Cognitive Services on those
|
||||
# domains, requests there must additionally carry an OpenAI-style path segment.
|
||||
_OPENAI_HOSTNAMES = ("api.openai.com",)
|
||||
_AZURE_OPENAI_HOSTNAMES = ("openai.azure.com", "cognitiveservices.azure.com")
|
||||
# Path markers that identify an Azure request as Azure OpenAI rather than Speech
|
||||
# / Vision / Language / ... `/openai/` is the native Azure OpenAI path prefix;
|
||||
# `/v1/` is the OpenAI-v1 surface used by LiteLLM's pass-through routing. Other
|
||||
# Cognitive Services use service-named prefixes and versions like `/v3.1/`,
|
||||
# `/v1.0/`, so they do not collide with these markers.
|
||||
_AZURE_OPENAI_PATH_MARKERS = ("/openai/", "/v1/")
|
||||
|
||||
|
||||
def _hostname_matches(hostname: str, suffixes: tuple) -> bool:
|
||||
"""True if hostname equals one of `suffixes` or is a subdomain of it.
|
||||
|
||||
Uses suffix matching (not a bare substring test) so look-alikes such as
|
||||
`cognitiveservices.azure.com.attacker.example` are not accepted.
|
||||
"""
|
||||
return any(
|
||||
hostname == suffix or hostname.endswith("." + suffix) for suffix in suffixes
|
||||
)
|
||||
|
||||
|
||||
def _is_openai_compatible_host(hostname: Optional[str]) -> bool:
|
||||
"""True if the hostname is OpenAI proper or one of the Azure OpenAI domains.
|
||||
|
||||
Hostname-only check, kept for the route-level helpers that additionally
|
||||
require a specific OpenAI path (e.g. `/v1/chat/completions`). When only the
|
||||
hostname would otherwise gate dispatch, use `_is_openai_compatible_url` so
|
||||
non-OpenAI Azure Cognitive Services on the shared domains are excluded.
|
||||
"""
|
||||
if not hostname:
|
||||
return False
|
||||
return _hostname_matches(hostname, _OPENAI_HOSTNAMES) or _hostname_matches(
|
||||
hostname, _AZURE_OPENAI_HOSTNAMES
|
||||
)
|
||||
|
||||
|
||||
def _is_openai_compatible_url(url_route: Optional[str]) -> bool:
|
||||
"""True if the URL targets an OpenAI-compatible API surface.
|
||||
|
||||
For the shared Azure Cognitive Services domains we additionally require an
|
||||
OpenAI-style path segment (`/openai/` or `/v1/`) so non-OpenAI Azure services
|
||||
(Speech, Vision, Language, ...) on the same domain are not misclassified as
|
||||
OpenAI routes.
|
||||
"""
|
||||
if not url_route:
|
||||
return False
|
||||
parsed_url = urlparse(url_route)
|
||||
hostname = parsed_url.hostname
|
||||
if not hostname:
|
||||
return False
|
||||
if _hostname_matches(hostname, _OPENAI_HOSTNAMES):
|
||||
return True
|
||||
if _hostname_matches(hostname, _AZURE_OPENAI_HOSTNAMES):
|
||||
return any(marker in parsed_url.path for marker in _AZURE_OPENAI_PATH_MARKERS)
|
||||
return False
|
||||
|
||||
|
||||
class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
|
||||
"""
|
||||
|
|
@ -52,12 +117,8 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
|
|||
if not url_route:
|
||||
return False
|
||||
parsed_url = urlparse(url_route)
|
||||
return bool(
|
||||
parsed_url.hostname
|
||||
and (
|
||||
"api.openai.com" in parsed_url.hostname
|
||||
or "openai.azure.com" in parsed_url.hostname
|
||||
)
|
||||
return (
|
||||
_is_openai_compatible_host(parsed_url.hostname)
|
||||
and "/v1/chat/completions" in parsed_url.path
|
||||
)
|
||||
|
||||
|
|
@ -67,12 +128,8 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
|
|||
if not url_route:
|
||||
return False
|
||||
parsed_url = urlparse(url_route)
|
||||
return bool(
|
||||
parsed_url.hostname
|
||||
and (
|
||||
"api.openai.com" in parsed_url.hostname
|
||||
or "openai.azure.com" in parsed_url.hostname
|
||||
)
|
||||
return (
|
||||
_is_openai_compatible_host(parsed_url.hostname)
|
||||
and "/v1/images/generations" in parsed_url.path
|
||||
)
|
||||
|
||||
|
|
@ -82,12 +139,8 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
|
|||
if not url_route:
|
||||
return False
|
||||
parsed_url = urlparse(url_route)
|
||||
return bool(
|
||||
parsed_url.hostname
|
||||
and (
|
||||
"api.openai.com" in parsed_url.hostname
|
||||
or "openai.azure.com" in parsed_url.hostname
|
||||
)
|
||||
return (
|
||||
_is_openai_compatible_host(parsed_url.hostname)
|
||||
and "/v1/images/edits" in parsed_url.path
|
||||
)
|
||||
|
||||
|
|
@ -97,13 +150,8 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
|
|||
if not url_route:
|
||||
return False
|
||||
parsed_url = urlparse(url_route)
|
||||
return bool(
|
||||
parsed_url.hostname
|
||||
and (
|
||||
"api.openai.com" in parsed_url.hostname
|
||||
or "openai.azure.com" in parsed_url.hostname
|
||||
)
|
||||
and ("/v1/responses" in parsed_url.path or "/responses" in parsed_url.path)
|
||||
return _is_openai_compatible_host(parsed_url.hostname) and (
|
||||
"/v1/responses" in parsed_url.path or "/responses" in parsed_url.path
|
||||
)
|
||||
|
||||
def _get_user_from_metadata(
|
||||
|
|
|
|||
|
|
@ -434,15 +434,20 @@ class PassThroughEndpointLogging:
|
|||
return False
|
||||
|
||||
def is_openai_route(self, url_route: str):
|
||||
"""Check if the URL route is an OpenAI API route."""
|
||||
"""Check if the URL route is an OpenAI API route.
|
||||
|
||||
Uses the URL-aware helper so that non-OpenAI Azure Cognitive Services
|
||||
(Speech, Vision, Language, ...) sharing the `*.cognitiveservices.azure.com`
|
||||
/ `*.openai.azure.com` domains are not misclassified as OpenAI routes.
|
||||
"""
|
||||
if not url_route:
|
||||
return False
|
||||
parsed_url = urlparse(url_route)
|
||||
return parsed_url.hostname and (
|
||||
"api.openai.com" in parsed_url.hostname
|
||||
or "openai.azure.com" in parsed_url.hostname
|
||||
from .llm_provider_handlers.openai_passthrough_logging_handler import (
|
||||
_is_openai_compatible_url,
|
||||
)
|
||||
|
||||
return _is_openai_compatible_url(url_route)
|
||||
|
||||
def is_gemini_route(
|
||||
self, url_route: str, custom_llm_provider: Optional[str] = None
|
||||
):
|
||||
|
|
|
|||
|
|
@ -11742,10 +11742,8 @@ async def _find_model_by_id(
|
|||
|
||||
@router.get(
|
||||
"/v2/model/info",
|
||||
description="v2 - returns models available to the user based on their API key permissions. Shows model info from config.yaml (except api key and api base). Filter to just user-added models with ?user_models_only=true",
|
||||
tags=["model management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
include_in_schema=False,
|
||||
)
|
||||
async def model_info_v2(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
|
|
@ -11781,7 +11779,49 @@ async def model_info_v2(
|
|||
),
|
||||
):
|
||||
"""
|
||||
BETA ENDPOINT. Might change unexpectedly. Use `/v1/model/info` for now.
|
||||
Paginated model metadata for proxy deployments (pricing, provider, team access).
|
||||
|
||||
Returns configured router deployments with enriched `model_info` (costs, provider,
|
||||
context window, etc.). Sensitive fields such as API keys and api_base are omitted.
|
||||
|
||||
Query parameters:
|
||||
model: Filter to a single public `model_name`.
|
||||
user_models_only: When true, only return models created by the calling user.
|
||||
include_team_models: When true, populate `access_via_team_ids` and `direct_access`
|
||||
on each model and filter to deployments the caller can use.
|
||||
page / size: Pagination controls (defaults: page=1, size=50).
|
||||
search: Case-insensitive partial match on model name or team public name.
|
||||
modelId: Return a single deployment by LiteLLM model id.
|
||||
teamId: Filter to models with direct access or team membership for this team id.
|
||||
sortBy / sortOrder: Sort by model_name, created_at, updated_at, costs, or status.
|
||||
|
||||
Example request:
|
||||
```
|
||||
curl -X GET 'http://localhost:4000/v2/model/info?include_team_models=true&page=1&size=50' \\
|
||||
--header 'Authorization: Bearer sk-1234'
|
||||
```
|
||||
|
||||
Example response:
|
||||
```json
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "openai/gpt-4.1"},
|
||||
"model_info": {
|
||||
"id": "abc123",
|
||||
"litellm_provider": "openai",
|
||||
"access_via_team_ids": ["team-1"],
|
||||
"direct_access": true
|
||||
}
|
||||
}
|
||||
],
|
||||
"total_count": 1,
|
||||
"current_page": 1,
|
||||
"total_pages": 1,
|
||||
"size": 50
|
||||
}
|
||||
```
|
||||
"""
|
||||
global llm_model_list, general_settings, user_config_file_path, proxy_config, llm_router
|
||||
|
||||
|
|
|
|||
|
|
@ -148,7 +148,9 @@ class LiteLLMCompletionResponsesConfig:
|
|||
# which is equivalent to "required" in OpenAI format
|
||||
return "required"
|
||||
elif tool_choice_type == "function":
|
||||
# function type without name - fall back to required
|
||||
function_name = tool_choice.get("name")
|
||||
if function_name:
|
||||
return {"type": "function", "function": {"name": function_name}}
|
||||
return "required"
|
||||
|
||||
# Return as-is for unknown formats
|
||||
|
|
|
|||
|
|
@ -232,6 +232,7 @@ class VoiceConfig(TypedDict):
|
|||
|
||||
class SpeechConfig(TypedDict, total=False):
|
||||
voiceConfig: VoiceConfig
|
||||
languageCode: str
|
||||
|
||||
|
||||
class GenerationConfig(TypedDict, total=False):
|
||||
|
|
@ -757,3 +758,12 @@ class VertexPartnerProvider(str, Enum):
|
|||
llama = "llama"
|
||||
ai21 = "ai21"
|
||||
claude = "claude"
|
||||
|
||||
|
||||
VERTEX_AI_PROVIDER_METADATA_FIELDS = (
|
||||
"vertex_ai_grounding_metadata",
|
||||
"vertex_ai_url_context_metadata",
|
||||
"vertex_ai_safety_ratings",
|
||||
"vertex_ai_safety_results",
|
||||
"vertex_ai_citation_metadata",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1540,6 +1540,11 @@ class ServerToolUse(BaseModel):
|
|||
web_search_requests: Optional[int] = None
|
||||
tool_search_requests: Optional[int] = None
|
||||
|
||||
def __getitem__(self, key: str) -> Optional[int]:
|
||||
if key not in self.__class__.model_fields:
|
||||
raise KeyError(key)
|
||||
return getattr(self, key)
|
||||
|
||||
|
||||
class Usage(SafeAttributeModel, CompletionUsage):
|
||||
_cache_creation_input_tokens: int = PrivateAttr(
|
||||
|
|
@ -1570,7 +1575,7 @@ class Usage(SafeAttributeModel, CompletionUsage):
|
|||
completion_tokens_details: Optional[
|
||||
Union[CompletionTokensDetailsWrapper, dict]
|
||||
] = None,
|
||||
server_tool_use: Optional[ServerToolUse] = None,
|
||||
server_tool_use: Optional[Union[ServerToolUse, dict]] = None,
|
||||
cost: Optional[float] = None,
|
||||
**params,
|
||||
):
|
||||
|
|
@ -1671,6 +1676,9 @@ class Usage(SafeAttributeModel, CompletionUsage):
|
|||
prompt_tokens_details=_prompt_tokens_details or None,
|
||||
)
|
||||
|
||||
if isinstance(server_tool_use, dict):
|
||||
server_tool_use = ServerToolUse(**server_tool_use)
|
||||
|
||||
if server_tool_use is not None:
|
||||
self.server_tool_use = server_tool_use
|
||||
else: # maintain openai compatibility in usage object if possible
|
||||
|
|
@ -3392,6 +3400,7 @@ class LlmProviders(str, Enum):
|
|||
POE = "poe"
|
||||
CHUTES = "chutes"
|
||||
NEOSANTARA = "neosantara"
|
||||
PARASAIL = "parasail"
|
||||
XIAOMI_MIMO = "xiaomi_mimo"
|
||||
TENSORMESH = "tensormesh"
|
||||
LITELLM_AGENT = "litellm_agent"
|
||||
|
|
|
|||
|
|
@ -8895,7 +8895,13 @@ class ProviderConfigManager:
|
|||
elif litellm.LlmProviders.XAI == provider:
|
||||
return litellm.XAIResponsesAPIConfig()
|
||||
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
|
||||
return litellm.GithubCopilotResponsesAPIConfig()
|
||||
from litellm.llms.github_copilot.responses.transformation import (
|
||||
github_copilot_supports_responses_api,
|
||||
)
|
||||
|
||||
if model is None or github_copilot_supports_responses_api(model=model):
|
||||
return litellm.GithubCopilotResponsesAPIConfig()
|
||||
return None
|
||||
elif litellm.LlmProviders.CHATGPT == provider:
|
||||
return litellm.ChatGPTResponsesAPIConfig()
|
||||
elif litellm.LlmProviders.LITELLM_PROXY == provider:
|
||||
|
|
|
|||
|
|
@ -6898,6 +6898,43 @@
|
|||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"azure_ai/MAI-Image-2.5": {
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.05,
|
||||
"output_cost_per_image_token": 4.7e-05,
|
||||
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/new-mai-models-in-microsoft-foundry-across-text-image-voice-and-speech/4524632",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
]
|
||||
},
|
||||
"azure_ai/MAI-Image-2.5-Flash": {
|
||||
"input_cost_per_image_token": 1.75e-06,
|
||||
"input_cost_per_token": 1.75e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0338,
|
||||
"output_cost_per_image_token": 3.3e-05,
|
||||
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/new-mai-models-in-microsoft-foundry-across-text-image-voice-and-speech/4524632",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
]
|
||||
},
|
||||
"azure_ai/MAI-Image-2e": {
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.02,
|
||||
"output_cost_per_image_token": 1.95e-05,
|
||||
"source": "https://aka.ms/mai-image-2e-foundryblog",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"azure_ai/Llama-3.2-11B-Vision-Instruct": {
|
||||
"input_cost_per_token": 3.7e-07,
|
||||
"litellm_provider": "azure_ai",
|
||||
|
|
@ -14295,10 +14332,10 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 4.4e-06,
|
||||
"source": "https://fireworks.ai/models/fireworks/glm-5p1",
|
||||
"supports_function_calling": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": false,
|
||||
"supports_tool_choice": false
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/gpt-oss-120b": {
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
|
|
@ -14576,10 +14613,10 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 4.4e-06,
|
||||
"source": "https://fireworks.ai/models/fireworks/glm-5p1",
|
||||
"supports_function_calling": false,
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": false,
|
||||
"supports_tool_choice": false
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"fireworks_ai/kimi-k2p5": {
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
|
|
|
|||
|
|
@ -1834,6 +1834,23 @@
|
|||
"search": true
|
||||
}
|
||||
},
|
||||
"parasail": {
|
||||
"display_name": "Parasail (`parasail`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/parasail",
|
||||
"endpoints": {
|
||||
"chat_completions": true,
|
||||
"messages": false,
|
||||
"responses": true,
|
||||
"embeddings": false,
|
||||
"image_generations": false,
|
||||
"audio_transcriptions": false,
|
||||
"audio_speech": false,
|
||||
"moderations": false,
|
||||
"batches": false,
|
||||
"rerank": false,
|
||||
"a2a": false
|
||||
}
|
||||
},
|
||||
"perplexity": {
|
||||
"display_name": "Perplexity AI (`perplexity`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/perplexity",
|
||||
|
|
|
|||
|
|
@ -53,7 +53,7 @@ proxy = [
|
|||
"orjson>=3.11.6,<4.0",
|
||||
"apscheduler>=3.11.2,<4.0",
|
||||
"fastapi-sso>=0.19.0,<1.0",
|
||||
"PyJWT>=2.12.0,<3.0",
|
||||
"PyJWT>=2.13.0,<3.0",
|
||||
"python-multipart>=0.0.27,<1.0",
|
||||
"cryptography>=46.0.7,<47.0",
|
||||
"pynacl>=1.6.2,<2.0",
|
||||
|
|
|
|||
|
|
@ -2712,6 +2712,10 @@ def test_bedrock_top_k_param(model, expected_params):
|
|||
data = json.loads(mock_post.call_args.kwargs["data"])
|
||||
if "mistral" in model:
|
||||
assert data["top_k"] == 2
|
||||
elif expected_params == {}:
|
||||
# Models that don't support top_k produce no additionalModelRequestFields;
|
||||
# the empty block is now omitted entirely rather than sent as `{}`.
|
||||
assert "additionalModelRequestFields" not in data
|
||||
else:
|
||||
assert data["additionalModelRequestFields"] == expected_params
|
||||
|
||||
|
|
@ -3059,8 +3063,6 @@ async def test_bedrock_max_completion_tokens(model: str):
|
|||
|
||||
assert request_body == {
|
||||
"messages": [{"role": "user", "content": [{"text": "Hello!"}]}],
|
||||
"additionalModelRequestFields": {},
|
||||
"system": [],
|
||||
"inferenceConfig": {"maxTokens": 10},
|
||||
}
|
||||
|
||||
|
|
|
|||
48
tests/test_litellm/caching/test_caching.py
Normal file
48
tests/test_litellm/caching/test_caching.py
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
import logging
|
||||
import re
|
||||
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.caching import LiteLLMCacheType
|
||||
|
||||
|
||||
def test_cache_key_debug_log_does_not_include_prompt_material(caplog):
|
||||
cache = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
prompt_marker = "secret prompt material "
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
||||
cache_key = cache.get_cache_key(
|
||||
model="gpt-4.1-mini",
|
||||
messages=[
|
||||
{"role": "system", "content": prompt_marker * 100},
|
||||
{"role": "user", "content": "hello"},
|
||||
],
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "lookup",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"query": {"type": "string"}},
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
response_format={
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "lookup_response",
|
||||
"schema": {"type": "object"},
|
||||
},
|
||||
},
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert re.fullmatch(r"[0-9a-f]{64}", cache_key)
|
||||
|
||||
created_cache_key_logs = [
|
||||
record.getMessage() for record in caplog.records if "Created cache key:" in record.getMessage()
|
||||
]
|
||||
assert created_cache_key_logs
|
||||
assert all(prompt_marker not in message for message in created_cache_key_logs)
|
||||
assert any(cache_key in message for message in created_cache_key_logs)
|
||||
|
|
@ -523,3 +523,468 @@ async def test_redis_semantic_cache_async_set_cache_stores_cache_key_filter(
|
|||
filters={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"},
|
||||
ttl=60,
|
||||
)
|
||||
|
||||
|
||||
def test_redis_semantic_cache_set_cache_uses_responses_string_input():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
redis_semantic_cache.llmcache = MagicMock()
|
||||
redis_semantic_cache._get_cache_filters = MagicMock(
|
||||
return_value={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"}
|
||||
)
|
||||
redis_semantic_cache._get_ttl = MagicMock(return_value=None)
|
||||
|
||||
redis_semantic_cache.set_cache(
|
||||
key="test_key",
|
||||
value={"content": "Paris"},
|
||||
input="What is the capital of France?",
|
||||
)
|
||||
|
||||
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"},
|
||||
)
|
||||
|
||||
|
||||
def test_redis_semantic_cache_get_cache_uses_responses_string_input():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
redis_semantic_cache.similarity_threshold = 0.8
|
||||
redis_semantic_cache.llmcache = MagicMock()
|
||||
redis_semantic_cache.llmcache.check = MagicMock(
|
||||
return_value=[
|
||||
{
|
||||
"prompt": "What is the capital of France?",
|
||||
"response": '{"content": "Paris"}',
|
||||
"vector_distance": 0.1,
|
||||
RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key",
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
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",
|
||||
input="What is the capital of France?",
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
assert result == {"content": "Paris"}
|
||||
assert metadata["semantic-similarity"] == pytest.approx(0.9)
|
||||
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_set_cache_flattens_structured_responses_input():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
redis_semantic_cache.llmcache = MagicMock()
|
||||
redis_semantic_cache._get_cache_filters = MagicMock(
|
||||
return_value={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"}
|
||||
)
|
||||
redis_semantic_cache._get_ttl = MagicMock(return_value=None)
|
||||
|
||||
redis_semantic_cache.set_cache(
|
||||
key="test_key",
|
||||
value={"content": "Paris"},
|
||||
input=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "input_text", "text": "What is the capital of France?"},
|
||||
{"type": "input_text", "text": "Answer briefly."},
|
||||
{
|
||||
"type": "input_image",
|
||||
"image_url": "https://example.com/paris.png",
|
||||
},
|
||||
],
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
redis_semantic_cache.llmcache.store.assert_called_once_with(
|
||||
"What is the capital of France?\nAnswer briefly.",
|
||||
"{'content': 'Paris'}",
|
||||
filters={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"},
|
||||
)
|
||||
|
||||
|
||||
def test_redis_semantic_cache_prompt_extraction_prefers_messages():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
prompt = RedisSemanticCache._get_prompt_from_kwargs(
|
||||
messages=[{"content": "message prompt"}],
|
||||
input="responses prompt",
|
||||
)
|
||||
|
||||
assert prompt == "message prompt"
|
||||
|
||||
|
||||
def test_redis_semantic_cache_prompt_extraction_handles_model_objects():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
class ModelDumpInput:
|
||||
def model_dump(self):
|
||||
return {"content": [{"text": "model dump prompt"}]}
|
||||
|
||||
class DictInput:
|
||||
def dict(self):
|
||||
return {"content": [{"output_text": "dict prompt"}]}
|
||||
|
||||
prompt = RedisSemanticCache._get_prompt_from_kwargs(
|
||||
input=[
|
||||
ModelDumpInput(),
|
||||
DictInput(),
|
||||
{"content": [{"input_text": "inline prompt"}]},
|
||||
{"content": [{"type": "input_image", "image_url": "https://example.com"}]},
|
||||
]
|
||||
)
|
||||
|
||||
assert prompt == "model dump prompt\ndict prompt\ninline prompt"
|
||||
|
||||
|
||||
def test_redis_semantic_cache_prompt_extraction_returns_none_without_text():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
assert RedisSemanticCache._get_prompt_from_kwargs() is None
|
||||
assert RedisSemanticCache._get_prompt_from_kwargs(input=None) is None
|
||||
assert RedisSemanticCache._get_prompt_from_kwargs(input=" ") is None
|
||||
assert (
|
||||
RedisSemanticCache._get_prompt_from_kwargs(
|
||||
input=[{"type": "input_image", "image_url": "https://example.com"}]
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_redis_semantic_cache_prompt_extraction_skips_blank_dict_text_keys():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
prompt = RedisSemanticCache._get_prompt_from_kwargs(
|
||||
input={"text": " ", "input_text": "fallback prompt"}
|
||||
)
|
||||
|
||||
assert prompt == "fallback prompt"
|
||||
|
||||
|
||||
def test_redis_semantic_cache_prompt_extraction_skips_blank_object_text_keys():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
class ResponseInput:
|
||||
text = " "
|
||||
input_text = "fallback prompt"
|
||||
|
||||
prompt = RedisSemanticCache._get_prompt_from_kwargs(input=ResponseInput())
|
||||
|
||||
assert prompt == "fallback prompt"
|
||||
|
||||
|
||||
def test_redis_semantic_cache_prompt_extraction_handles_object_content():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
class ResponseInput:
|
||||
content = [{"text": "object content prompt"}]
|
||||
|
||||
prompt = RedisSemanticCache._get_prompt_from_kwargs(input=ResponseInput())
|
||||
|
||||
assert prompt == "object content prompt"
|
||||
|
||||
|
||||
def test_redis_semantic_cache_set_cache_skips_blank_responses_input():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
redis_semantic_cache.llmcache = MagicMock()
|
||||
|
||||
redis_semantic_cache.set_cache(
|
||||
key="test_key",
|
||||
value={"content": "Paris"},
|
||||
input=" ",
|
||||
)
|
||||
|
||||
redis_semantic_cache.llmcache.store.assert_not_called()
|
||||
|
||||
|
||||
def test_redis_semantic_cache_get_cache_sets_similarity_on_blank_responses_input():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
redis_semantic_cache.llmcache = MagicMock()
|
||||
metadata = {}
|
||||
|
||||
result = redis_semantic_cache.get_cache(
|
||||
key="test_key",
|
||||
input=" ",
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
assert result is None
|
||||
assert metadata["semantic-similarity"] == 0.0
|
||||
redis_semantic_cache.llmcache.check.assert_not_called()
|
||||
|
||||
|
||||
def test_redis_semantic_cache_get_cache_sets_similarity_when_no_results():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
redis_semantic_cache.llmcache = MagicMock()
|
||||
redis_semantic_cache.llmcache.check = MagicMock(return_value=[])
|
||||
|
||||
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",
|
||||
input="What is the capital of France?",
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
assert result is None
|
||||
assert metadata["semantic-similarity"] == 0.0
|
||||
redis_semantic_cache.llmcache.check.assert_called_once_with(
|
||||
prompt="What is the capital of France?",
|
||||
filter_expression="cache-key-filter",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_semantic_cache_async_paths_use_responses_string_input():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
redis_semantic_cache.similarity_threshold = 0.8
|
||||
redis_semantic_cache.llmcache = MagicMock()
|
||||
redis_semantic_cache.llmcache.astore = AsyncMock()
|
||||
redis_semantic_cache.llmcache.acheck = AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"prompt": "What is the capital of France?",
|
||||
"response": '{"content": "Paris"}',
|
||||
"vector_distance": 0.1,
|
||||
RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key",
|
||||
}
|
||||
]
|
||||
)
|
||||
redis_semantic_cache._get_cache_filters = MagicMock(
|
||||
return_value={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"}
|
||||
)
|
||||
redis_semantic_cache._get_ttl = MagicMock(return_value=None)
|
||||
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"},
|
||||
input="What is the capital of France?",
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
redis_semantic_cache,
|
||||
"_get_cache_key_filter_expression",
|
||||
return_value="cache-key-filter",
|
||||
):
|
||||
metadata = {}
|
||||
result = await redis_semantic_cache.async_get_cache(
|
||||
key="test_key",
|
||||
input="What is the capital of France?",
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
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"},
|
||||
)
|
||||
assert result == {"content": "Paris"}
|
||||
assert metadata["semantic-similarity"] == pytest.approx(0.9)
|
||||
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_paths_set_similarity_on_misses():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
redis_semantic_cache.llmcache = MagicMock()
|
||||
redis_semantic_cache.llmcache.astore = AsyncMock()
|
||||
redis_semantic_cache.llmcache.acheck = AsyncMock(return_value=[])
|
||||
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"},
|
||||
input=" ",
|
||||
)
|
||||
|
||||
redis_semantic_cache.llmcache.astore.assert_not_called()
|
||||
redis_semantic_cache._get_async_embedding.assert_not_called()
|
||||
|
||||
blank_metadata = {}
|
||||
blank_result = await redis_semantic_cache.async_get_cache(
|
||||
key="test_key",
|
||||
input=" ",
|
||||
metadata=blank_metadata,
|
||||
)
|
||||
|
||||
assert blank_result is None
|
||||
assert blank_metadata["semantic-similarity"] == 0.0
|
||||
redis_semantic_cache.llmcache.acheck.assert_not_called()
|
||||
redis_semantic_cache._get_async_embedding.assert_not_called()
|
||||
|
||||
with patch.object(
|
||||
redis_semantic_cache,
|
||||
"_get_cache_key_filter_expression",
|
||||
return_value="cache-key-filter",
|
||||
):
|
||||
miss_metadata = {}
|
||||
miss_result = await redis_semantic_cache.async_get_cache(
|
||||
key="test_key",
|
||||
input="What is the capital of France?",
|
||||
metadata=miss_metadata,
|
||||
)
|
||||
|
||||
assert miss_result is None
|
||||
assert miss_metadata["semantic-similarity"] == 0.0
|
||||
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",
|
||||
)
|
||||
|
||||
|
||||
def test_cache_get_cache_passes_responses_input_to_backend_cache():
|
||||
from litellm.caching.caching import Cache
|
||||
|
||||
cache = Cache.__new__(Cache)
|
||||
cache.cache = MagicMock()
|
||||
cache.cache.get_cache = MagicMock(return_value=None)
|
||||
cache.should_use_cache = MagicMock(return_value=True)
|
||||
cache.get_cache_key = MagicMock(return_value="test_key")
|
||||
|
||||
metadata = {}
|
||||
cache.get_cache(
|
||||
input="What is the capital of France?",
|
||||
metadata=metadata,
|
||||
cache={},
|
||||
)
|
||||
|
||||
cache.cache.get_cache.assert_called_once_with(
|
||||
"test_key",
|
||||
input="What is the capital of France?",
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
|
||||
def test_cache_get_cache_filters_sensitive_kwargs_from_backend_cache():
|
||||
from litellm.caching.caching import Cache
|
||||
|
||||
cache = Cache.__new__(Cache)
|
||||
cache.cache = MagicMock()
|
||||
cache.should_use_cache = MagicMock(return_value=True)
|
||||
cache.get_cache_key = MagicMock(return_value="test_key")
|
||||
cache._get_cache_logic = MagicMock(return_value={"content": "Paris"})
|
||||
|
||||
def _cache_hit(_cache_key, **cache_kwargs):
|
||||
cache_kwargs["metadata"]["semantic-similarity"] = 0.7
|
||||
return {"content": "Paris"}
|
||||
|
||||
cache.cache.get_cache = MagicMock(side_effect=_cache_hit)
|
||||
|
||||
metadata = {"user_api_key": "sk-secret", "trace_id": "trace-id"}
|
||||
result = cache.get_cache(
|
||||
input="What is the capital of France?",
|
||||
metadata=metadata,
|
||||
cache={"s-maxage": 10},
|
||||
api_key="sk-secret",
|
||||
headers={"authorization": "Bearer sk-secret"},
|
||||
)
|
||||
|
||||
assert result == {"content": "Paris"}
|
||||
assert metadata == {
|
||||
"user_api_key": "sk-secret",
|
||||
"trace_id": "trace-id",
|
||||
"semantic-similarity": 0.7,
|
||||
}
|
||||
|
||||
forwarded_kwargs = cache.cache.get_cache.call_args.kwargs
|
||||
assert forwarded_kwargs == {
|
||||
"input": "What is the capital of France?",
|
||||
"metadata": {"semantic-similarity": 0.7},
|
||||
}
|
||||
assert forwarded_kwargs["metadata"] is not metadata
|
||||
cache._get_cache_logic.assert_called_once_with(
|
||||
cached_result={"content": "Paris"},
|
||||
max_age=10,
|
||||
)
|
||||
|
||||
|
||||
def test_cache_get_cache_filters_sensitive_kwargs_without_metadata():
|
||||
from litellm.caching.caching import Cache
|
||||
|
||||
cache = Cache.__new__(Cache)
|
||||
cache.cache = MagicMock()
|
||||
cache.cache.get_cache = MagicMock(return_value={"content": "Paris"})
|
||||
cache.should_use_cache = MagicMock(return_value=True)
|
||||
cache.get_cache_key = MagicMock(return_value="test_key")
|
||||
cache._get_cache_logic = MagicMock(return_value={"content": "Paris"})
|
||||
|
||||
result = cache.get_cache(
|
||||
input="What is the capital of France?",
|
||||
cache={"s-maxage": 10},
|
||||
api_key="sk-secret",
|
||||
headers={"authorization": "Bearer sk-secret"},
|
||||
)
|
||||
|
||||
assert result == {"content": "Paris"}
|
||||
cache.cache.get_cache.assert_called_once_with(
|
||||
"test_key",
|
||||
input="What is the capital of France?",
|
||||
)
|
||||
|
||||
|
||||
def test_cache_get_cache_passes_responses_input_to_dynamic_cache():
|
||||
from litellm.caching.caching import Cache
|
||||
|
||||
cache = Cache.__new__(Cache)
|
||||
cache.should_use_cache = MagicMock(return_value=True)
|
||||
cache.get_cache_key = MagicMock(return_value="test_key")
|
||||
cache._get_cache_logic = MagicMock(return_value={"content": "Paris"})
|
||||
dynamic_cache_object = MagicMock()
|
||||
dynamic_cache_object.get_cache = MagicMock(return_value={"content": "Paris"})
|
||||
|
||||
metadata = {}
|
||||
result = cache.get_cache(
|
||||
dynamic_cache_object=dynamic_cache_object,
|
||||
input="What is the capital of France?",
|
||||
metadata=metadata,
|
||||
cache={},
|
||||
)
|
||||
|
||||
assert result == {"content": "Paris"}
|
||||
dynamic_cache_object.get_cache.assert_called_once_with(
|
||||
"test_key",
|
||||
input="What is the capital of France?",
|
||||
metadata=metadata,
|
||||
)
|
||||
cache._get_cache_logic.assert_called_once_with(
|
||||
cached_result={"content": "Paris"},
|
||||
max_age=float("inf"),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -72,3 +72,18 @@ async def test_should_reject_invalid_limit(monkeypatch: pytest.MonkeyPatch):
|
|||
await db.get_usage_data(limit="invalid")
|
||||
|
||||
assert query_mock.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_join_organization_table(monkeypatch: pytest.MonkeyPatch):
|
||||
db, query_mock = _setup_db(monkeypatch, [])
|
||||
|
||||
await db.get_usage_data()
|
||||
|
||||
query_text, *_ = query_mock.await_args.args
|
||||
assert (
|
||||
"COALESCE(vt.organization_id, tt.organization_id) as organization_id"
|
||||
in query_text
|
||||
)
|
||||
assert "ot.organization_alias as organization_alias" in query_text
|
||||
assert 'LEFT JOIN "LiteLLM_OrganizationTable" ot' in query_text
|
||||
|
|
|
|||
|
|
@ -0,0 +1,180 @@
|
|||
"""Tests for FocusGCSDestination."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.focus.destinations.base import FocusTimeWindow
|
||||
|
||||
|
||||
def _make_window(frequency: str = "hourly") -> FocusTimeWindow:
|
||||
return FocusTimeWindow(
|
||||
start_time=datetime(2026, 1, 1, 10, 0, 0, tzinfo=timezone.utc),
|
||||
end_time=datetime(2026, 1, 1, 11, 0, 0, tzinfo=timezone.utc),
|
||||
frequency=frequency,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deliver_posts_to_gcs_upload_endpoint():
|
||||
"""deliver() must POST raw bytes to the GCS upload endpoint."""
|
||||
from litellm.integrations.focus.destinations.gcs_destination import (
|
||||
FocusGCSDestination,
|
||||
)
|
||||
|
||||
dest = FocusGCSDestination(
|
||||
prefix="focus_exports",
|
||||
config={"bucket_name": "my-bucket", "service_account_json": None},
|
||||
)
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.post = AsyncMock(return_value=mock_response)
|
||||
dest.async_httpx_client = mock_client
|
||||
|
||||
with patch.object(
|
||||
dest,
|
||||
"construct_request_headers",
|
||||
new=AsyncMock(return_value={"Authorization": "Bearer tok-123"}),
|
||||
):
|
||||
await dest.deliver(
|
||||
content=b"col1,col2\nval1,val2\n",
|
||||
time_window=_make_window(),
|
||||
filename="usage_20260101T100000Z_20260101T110000Z.csv",
|
||||
)
|
||||
|
||||
mock_client.post.assert_called_once()
|
||||
call_kwargs = mock_client.post.call_args
|
||||
url = call_kwargs.kwargs.get("url") or call_kwargs.args[0]
|
||||
assert "my-bucket" in url
|
||||
assert "uploadType=media" in url
|
||||
headers = call_kwargs.kwargs["headers"]
|
||||
assert headers["Authorization"] == "Bearer tok-123"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deliver_raises_on_gcs_error():
|
||||
"""deliver() must raise RuntimeError when GCS returns non-200."""
|
||||
from litellm.integrations.focus.destinations.gcs_destination import (
|
||||
FocusGCSDestination,
|
||||
)
|
||||
|
||||
dest = FocusGCSDestination(
|
||||
prefix="focus_exports",
|
||||
config={"bucket_name": "my-bucket"},
|
||||
)
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 403
|
||||
mock_response.text = "Permission denied"
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.post = AsyncMock(return_value=mock_response)
|
||||
dest.async_httpx_client = mock_client
|
||||
|
||||
with patch.object(
|
||||
dest,
|
||||
"construct_request_headers",
|
||||
new=AsyncMock(return_value={"Authorization": "Bearer tok-bad"}),
|
||||
):
|
||||
with pytest.raises(RuntimeError, match="GCS upload failed"):
|
||||
await dest.deliver(
|
||||
content=b"data",
|
||||
time_window=_make_window(),
|
||||
filename="usage.csv",
|
||||
)
|
||||
|
||||
|
||||
def test_build_object_key_hourly():
|
||||
"""Hourly key must include date= and hour= components."""
|
||||
from litellm.integrations.focus.destinations.gcs_destination import (
|
||||
FocusGCSDestination,
|
||||
)
|
||||
|
||||
dest = FocusGCSDestination(prefix="focus_exports", config={"bucket_name": "b"})
|
||||
key = dest._build_object_key(
|
||||
time_window=_make_window("hourly"), filename="usage.parquet"
|
||||
)
|
||||
|
||||
assert key == "focus_exports/date=2026-01-01/hour=10/usage.parquet"
|
||||
|
||||
|
||||
def test_build_object_key_daily():
|
||||
"""Daily key must include date= but not hour=."""
|
||||
from litellm.integrations.focus.destinations.gcs_destination import (
|
||||
FocusGCSDestination,
|
||||
)
|
||||
|
||||
dest = FocusGCSDestination(prefix="focus_exports", config={"bucket_name": "b"})
|
||||
window = FocusTimeWindow(
|
||||
start_time=datetime(2026, 1, 1, 0, 0, 0, tzinfo=timezone.utc),
|
||||
end_time=datetime(2026, 1, 2, 0, 0, 0, tzinfo=timezone.utc),
|
||||
frequency="daily",
|
||||
)
|
||||
key = dest._build_object_key(time_window=window, filename="usage.parquet")
|
||||
|
||||
assert key == "focus_exports/date=2026-01-01/usage.parquet"
|
||||
|
||||
|
||||
def test_missing_bucket_name_raises():
|
||||
"""Constructing without bucket_name must raise ValueError."""
|
||||
from litellm.integrations.focus.destinations.gcs_destination import (
|
||||
FocusGCSDestination,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="bucket_name"):
|
||||
FocusGCSDestination(prefix="focus_exports", config={})
|
||||
|
||||
|
||||
def test_global_gcs_service_account_not_overwritten_when_absent(monkeypatch):
|
||||
"""service_account_json absent from config must not overwrite GCS_PATH_SERVICE_ACCOUNT.
|
||||
|
||||
GCSBucketBase sets self.path_service_account_json from GCS_PATH_SERVICE_ACCOUNT.
|
||||
If config has no service_account_json key, we must leave the parent value intact
|
||||
so deployments using the global credential don't silently fall back to ADC.
|
||||
"""
|
||||
monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/global/sa.json")
|
||||
|
||||
from litellm.integrations.focus.destinations.gcs_destination import (
|
||||
FocusGCSDestination,
|
||||
)
|
||||
|
||||
dest = FocusGCSDestination(prefix="focus_exports", config={"bucket_name": "b"})
|
||||
|
||||
assert dest.path_service_account_json == "/global/sa.json"
|
||||
|
||||
|
||||
def test_explicit_service_account_overrides_global(monkeypatch):
|
||||
"""Explicit service_account_json in config must take precedence over GCS_PATH_SERVICE_ACCOUNT."""
|
||||
monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/global/sa.json")
|
||||
|
||||
from litellm.integrations.focus.destinations.gcs_destination import (
|
||||
FocusGCSDestination,
|
||||
)
|
||||
|
||||
dest = FocusGCSDestination(
|
||||
prefix="focus_exports",
|
||||
config={"bucket_name": "b", "service_account_json": "/focus/sa.json"},
|
||||
)
|
||||
|
||||
assert dest.path_service_account_json == "/focus/sa.json"
|
||||
|
||||
|
||||
def test_factory_creates_gcs_destination(monkeypatch):
|
||||
"""FocusDestinationFactory.create(provider='gcs') must return FocusGCSDestination."""
|
||||
monkeypatch.setenv("FOCUS_GCS_BUCKET_NAME", "env-bucket")
|
||||
|
||||
from litellm.integrations.focus.destinations.factory import FocusDestinationFactory
|
||||
from litellm.integrations.focus.destinations.gcs_destination import (
|
||||
FocusGCSDestination,
|
||||
)
|
||||
|
||||
dest = FocusDestinationFactory.create(provider="gcs", prefix="focus_exports")
|
||||
|
||||
assert isinstance(dest, FocusGCSDestination)
|
||||
assert dest.BUCKET_NAME == "env-bucket"
|
||||
62
tests/test_litellm/integrations/focus/test_transformer.py
Normal file
62
tests/test_litellm/integrations/focus/test_transformer.py
Normal file
|
|
@ -0,0 +1,62 @@
|
|||
"""Tests for FocusTransformer organization metadata in Tags."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import date
|
||||
|
||||
import polars as pl
|
||||
|
||||
from litellm.integrations.focus.transformer import FocusTransformer
|
||||
|
||||
|
||||
def test_should_include_organization_fields_in_tags():
|
||||
frame = pl.DataFrame(
|
||||
{
|
||||
"date": [date(2024, 1, 2)],
|
||||
"spend": [1.25],
|
||||
"api_requests": [1],
|
||||
"api_key": ["hashed-key"],
|
||||
"api_key_alias": ["prod-key"],
|
||||
"model": ["gpt-4o"],
|
||||
"model_group": ["gpt-4o"],
|
||||
"custom_llm_provider": ["openai"],
|
||||
"team_id": ["team-1"],
|
||||
"team_alias": ["Platform"],
|
||||
"organization_id": ["org-123"],
|
||||
"organization_alias": ["Acme Corp"],
|
||||
"user_id": ["user-1"],
|
||||
"user_email": ["user@example.com"],
|
||||
}
|
||||
)
|
||||
|
||||
normalized = FocusTransformer().transform(frame)
|
||||
|
||||
tags = json.loads(normalized["Tags"][0])
|
||||
assert tags["organization_id"] == "org-123"
|
||||
assert tags["organization_alias"] == "Acme Corp"
|
||||
assert tags["team_id"] == "team-1"
|
||||
|
||||
|
||||
def test_should_omit_missing_organization_fields_from_tags():
|
||||
frame = pl.DataFrame(
|
||||
{
|
||||
"date": [date(2024, 1, 2)],
|
||||
"spend": [0.5],
|
||||
"api_requests": [1],
|
||||
"api_key": ["hashed-key"],
|
||||
"api_key_alias": ["prod-key"],
|
||||
"model": ["gpt-4o-mini"],
|
||||
"model_group": ["gpt-4o-mini"],
|
||||
"custom_llm_provider": ["openai"],
|
||||
"team_id": ["team-1"],
|
||||
"team_alias": ["Platform"],
|
||||
}
|
||||
)
|
||||
|
||||
normalized = FocusTransformer().transform(frame)
|
||||
|
||||
tags = json.loads(normalized["Tags"][0])
|
||||
assert "organization_id" not in tags
|
||||
assert "organization_alias" not in tags
|
||||
assert tags["team_id"] == "team-1"
|
||||
|
|
@ -357,12 +357,18 @@ def test_galileo_record_to_v2_span_with_tags_and_offset():
|
|||
|
||||
def test_galileo_get_output_str_variants(galileo_v2_env):
|
||||
logger = GalileoObserve()
|
||||
assert logger.get_output_str_from_response(None, {}) is None
|
||||
assert logger.get_output_str_from_response(None, {}) == ""
|
||||
assert (
|
||||
logger.get_output_str_from_response(
|
||||
EmbeddingResponse(), {"call_type": "embedding"}
|
||||
)
|
||||
is None
|
||||
== "embedding-output"
|
||||
)
|
||||
assert (
|
||||
logger.get_output_str_from_response(
|
||||
EmbeddingResponse(), {"call_type": "aembedding"}
|
||||
)
|
||||
== "embedding-output"
|
||||
)
|
||||
|
||||
text_resp = TextCompletionResponse()
|
||||
|
|
@ -414,7 +420,7 @@ def test_galileo_get_output_str_variants(galileo_v2_env):
|
|||
{"call_type": "acompletion", "messages": [{"role": "user", "content": "hi"}]},
|
||||
)
|
||||
|
||||
assert logger.get_output_str_from_response("not-a-supported-type", {}) is None
|
||||
assert logger.get_output_str_from_response("not-a-supported-type", {}) == ""
|
||||
|
||||
|
||||
def test_galileo_get_input_output_error_status_message(galileo_v2_env):
|
||||
|
|
@ -445,6 +451,48 @@ def test_galileo_get_output_str_rerank_response(galileo_v2_env):
|
|||
assert '"relevance_score": 0.98' in output
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_galileo_async_log_success_embedding(galileo_v2_env):
|
||||
import datetime
|
||||
|
||||
logger = GalileoObserve()
|
||||
embedding_response = EmbeddingResponse(
|
||||
data=[{"object": "embedding", "embedding": [0.1, 0.2, 0.3], "index": 0}]
|
||||
)
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.is_success = True
|
||||
mock_response.status_code = 201
|
||||
|
||||
with patch.object(logger.async_httpx_handler, "post", return_value=mock_response):
|
||||
await logger.async_log_success_event(
|
||||
kwargs={
|
||||
"call_type": "aembedding",
|
||||
"model": "text-embedding-3-small",
|
||||
"input": "hello world",
|
||||
"standard_logging_object": {
|
||||
"call_type": "aembedding",
|
||||
"model": "text-embedding-3-small",
|
||||
"prompt_tokens": 2,
|
||||
"completion_tokens": 0,
|
||||
"total_tokens": 2,
|
||||
"response_cost": 0.0,
|
||||
"startTime": datetime.datetime(
|
||||
2026, 5, 25, 12, 0, 0, tzinfo=datetime.timezone.utc
|
||||
).timestamp(),
|
||||
"endTime": datetime.datetime(
|
||||
2026, 5, 25, 12, 0, 1, tzinfo=datetime.timezone.utc
|
||||
).timestamp(),
|
||||
},
|
||||
},
|
||||
response_obj=embedding_response,
|
||||
start_time=datetime.datetime(2026, 5, 25, 12, 0, 0),
|
||||
end_time=datetime.datetime(2026, 5, 25, 12, 0, 1),
|
||||
)
|
||||
|
||||
assert logger.in_memory_records == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_galileo_async_log_success_rerank(galileo_v2_env):
|
||||
import datetime
|
||||
|
|
@ -524,6 +572,158 @@ def test_galileo_get_ingest_request_legacy(monkeypatch):
|
|||
assert payload["traces"][0]["input"] == "hi"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_galileo_async_health_check_success(galileo_v2_env):
|
||||
logger = GalileoObserve()
|
||||
current_user_resp = MagicMock()
|
||||
current_user_resp.status_code = 200
|
||||
|
||||
with patch.object(
|
||||
logger.async_httpx_handler, "get", new_callable=AsyncMock
|
||||
) as mock_get:
|
||||
mock_get.return_value = current_user_resp
|
||||
result = await logger.async_health_check()
|
||||
|
||||
assert result["status"] == "healthy"
|
||||
mock_get.assert_awaited_once_with(
|
||||
url="https://api.galileo.ai/current_user",
|
||||
headers={
|
||||
"accept": "application/json",
|
||||
"Content-Type": "application/json",
|
||||
"Galileo-API-Key": "test-api-key",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_galileo_async_health_check_api_error(galileo_v2_env):
|
||||
logger = GalileoObserve()
|
||||
current_user_resp = MagicMock()
|
||||
current_user_resp.status_code = 401
|
||||
|
||||
with patch.object(
|
||||
logger.async_httpx_handler, "get", new_callable=AsyncMock
|
||||
) as mock_get:
|
||||
mock_get.return_value = current_user_resp
|
||||
result = await logger.async_health_check()
|
||||
|
||||
assert result["status"] == "unhealthy"
|
||||
assert "HTTP 401" in result["error_message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_galileo_async_health_check_missing_project_id(monkeypatch):
|
||||
monkeypatch.setenv("GALILEO_API_KEY", "test-api-key")
|
||||
monkeypatch.setenv("GALILEO_BASE_URL", "https://api.galileo.ai")
|
||||
monkeypatch.delenv("GALILEO_PROJECT_ID", raising=False)
|
||||
logger = GalileoObserve()
|
||||
|
||||
result = await logger.async_health_check()
|
||||
|
||||
assert result["status"] == "unhealthy"
|
||||
assert "GALILEO_PROJECT_ID" in result["error_message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_galileo_async_health_check_missing_base_url(monkeypatch):
|
||||
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
|
||||
monkeypatch.delenv("GALILEO_BASE_URL", raising=False)
|
||||
monkeypatch.setenv("GALILEO_PROJECT_ID", "p")
|
||||
monkeypatch.setenv("GALILEO_USERNAME", "u")
|
||||
monkeypatch.setenv("GALILEO_PASSWORD", "pw")
|
||||
logger = GalileoObserve()
|
||||
|
||||
result = await logger.async_health_check()
|
||||
|
||||
assert result["status"] == "unhealthy"
|
||||
assert "GALILEO_BASE_URL" in result["error_message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_galileo_async_health_check_missing_credentials(monkeypatch):
|
||||
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
|
||||
monkeypatch.delenv("GALILEO_USERNAME", raising=False)
|
||||
monkeypatch.delenv("GALILEO_PASSWORD", raising=False)
|
||||
monkeypatch.setenv("GALILEO_PROJECT_ID", "p")
|
||||
monkeypatch.setenv("GALILEO_BASE_URL", "https://galileo.example")
|
||||
logger = GalileoObserve()
|
||||
|
||||
result = await logger.async_health_check()
|
||||
|
||||
assert result["status"] == "unhealthy"
|
||||
assert "GALILEO_USERNAME" in result["error_message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_galileo_async_health_check_auth_failed(monkeypatch):
|
||||
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
|
||||
monkeypatch.setenv("GALILEO_PROJECT_ID", "p")
|
||||
monkeypatch.setenv("GALILEO_BASE_URL", "https://galileo.example")
|
||||
monkeypatch.setenv("GALILEO_USERNAME", "u")
|
||||
monkeypatch.setenv("GALILEO_PASSWORD", "pw")
|
||||
logger = GalileoObserve()
|
||||
|
||||
with patch.object(
|
||||
logger.async_httpx_handler, "post", new_callable=AsyncMock
|
||||
) as mock_post:
|
||||
mock_post.side_effect = Exception("login failed")
|
||||
result = await logger.async_health_check()
|
||||
|
||||
assert result["status"] == "unhealthy"
|
||||
assert result["error_message"] == "Galileo authentication failed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_galileo_async_health_check_request_exception(galileo_v2_env):
|
||||
logger = GalileoObserve()
|
||||
|
||||
with patch.object(
|
||||
logger.async_httpx_handler, "get", new_callable=AsyncMock
|
||||
) as mock_get:
|
||||
mock_get.side_effect = Exception("connection refused")
|
||||
result = await logger.async_health_check()
|
||||
|
||||
assert result["status"] == "unhealthy"
|
||||
assert "connection refused" in result["error_message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_galileo_async_log_success_empty_model_response(galileo_v2_env):
|
||||
import datetime
|
||||
|
||||
logger = GalileoObserve()
|
||||
logger.batch_size = 2
|
||||
empty_response = ModelResponse(choices=[])
|
||||
|
||||
await logger.async_log_success_event(
|
||||
kwargs={
|
||||
"call_type": "acompletion",
|
||||
"model": "gpt-5.2",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"standard_logging_object": {
|
||||
"call_type": "acompletion",
|
||||
"model": "gpt-5.2",
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": 0,
|
||||
"total_tokens": 1,
|
||||
"response_cost": 0.0,
|
||||
"startTime": datetime.datetime(
|
||||
2026, 5, 25, 12, 0, 0, tzinfo=datetime.timezone.utc
|
||||
).timestamp(),
|
||||
"endTime": datetime.datetime(
|
||||
2026, 5, 25, 12, 0, 1, tzinfo=datetime.timezone.utc
|
||||
).timestamp(),
|
||||
},
|
||||
},
|
||||
response_obj=empty_response,
|
||||
start_time=datetime.datetime(2026, 5, 25, 12, 0, 0),
|
||||
end_time=datetime.datetime(2026, 5, 25, 12, 0, 1),
|
||||
)
|
||||
|
||||
assert len(logger.in_memory_records) == 1
|
||||
assert logger.in_memory_records[0]["output_text"] == ""
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_galileo_ensure_headers_v2_missing_key(monkeypatch):
|
||||
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
|
||||
|
|
|
|||
|
|
@ -1,17 +1,14 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
|
||||
StandardBuiltInToolCostTracking,
|
||||
)
|
||||
from litellm.types.llms.openai import FileSearchTool, WebSearchOptions
|
||||
from litellm.types.utils import ModelInfo, ModelResponse, StandardBuiltInToolsParams
|
||||
from litellm.types.utils import ModelResponse, StandardBuiltInToolsParams
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../..")
|
||||
|
|
@ -139,6 +136,22 @@ def test_get_cost_for_anthropic_web_search():
|
|||
assert cost > 0.0
|
||||
|
||||
|
||||
def test_get_cost_for_anthropic_web_search_with_server_tool_use_dict():
|
||||
"""
|
||||
Anthropic-compatible passthrough responses can construct Usage from a raw
|
||||
usage payload. Ensure dict server_tool_use values are normalized before
|
||||
built-in tool cost tracking reads server_tool_use.web_search_requests.
|
||||
"""
|
||||
from litellm.types.utils import ServerToolUse, Usage
|
||||
|
||||
usage = Usage(server_tool_use={"web_search_requests": 1})
|
||||
|
||||
assert isinstance(usage.server_tool_use, ServerToolUse)
|
||||
assert StandardBuiltInToolCostTracking.response_object_includes_web_search_call(
|
||||
response_object=None, usage=usage
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model", ["gemini/gemini-2.0-flash-001", "gemini-2.0-flash-001"]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2165,6 +2165,41 @@ def test_get_assembled_streaming_response_returns_result_for_streaming():
|
|||
assert assembled is result
|
||||
|
||||
|
||||
def test_streaming_success_handler_includes_vertex_ai_metadata_in_standard_logging():
|
||||
"""Assembled streaming responses should include Vertex AI metadata in logging payload."""
|
||||
import datetime
|
||||
|
||||
from litellm.types.utils import Choices, Message
|
||||
|
||||
logging_obj = _make_logging_obj(stream=True)
|
||||
grounding_metadata = [{"webSearchQueries": ["weather in SF"]}]
|
||||
url_context_metadata = [{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}]
|
||||
result = ModelResponse(
|
||||
id="resp-1",
|
||||
choices=[
|
||||
Choices(
|
||||
index=0,
|
||||
message=Message(role="assistant", content="hello"),
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
model="gemini-2.5-flash",
|
||||
)
|
||||
setattr(result, "vertex_ai_grounding_metadata", grounding_metadata)
|
||||
setattr(result, "vertex_ai_url_context_metadata", url_context_metadata)
|
||||
result._hidden_params["vertex_ai_grounding_metadata"] = grounding_metadata
|
||||
result._hidden_params["vertex_ai_url_context_metadata"] = url_context_metadata
|
||||
|
||||
start = datetime.datetime.now()
|
||||
end = datetime.datetime.now()
|
||||
logging_obj.success_handler(result=result, start_time=start, end_time=end)
|
||||
|
||||
payload = logging_obj.model_call_details.get("standard_logging_object")
|
||||
assert payload is not None
|
||||
assert payload["response"]["vertex_ai_grounding_metadata"] == grounding_metadata
|
||||
assert payload["response"]["vertex_ai_url_context_metadata"] == url_context_metadata
|
||||
|
||||
|
||||
def test_get_assembled_streaming_response_returns_none_for_non_streaming_text_completion():
|
||||
"""Non-streaming TextCompletionResponse should also return None."""
|
||||
import datetime
|
||||
|
|
|
|||
|
|
@ -349,3 +349,96 @@ class TestPerformRedaction:
|
|||
|
||||
assert redacted.output[0].content[0].text == "redacted-by-litellm"
|
||||
assert response.output[0].content[0].text == "sensitive output"
|
||||
|
||||
def test_redacts_vertex_provider_metadata_in_standard_logging_response(self):
|
||||
details = {
|
||||
"standard_logging_object": {
|
||||
"messages": [{"role": "user", "content": "sensitive prompt"}],
|
||||
"response": {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"content": "sensitive answer",
|
||||
"role": "assistant",
|
||||
}
|
||||
}
|
||||
],
|
||||
"vertex_ai_grounding_metadata": [
|
||||
{"webSearchQueries": ["sensitive search term"]}
|
||||
],
|
||||
"vertex_ai_url_context_metadata": [
|
||||
{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}
|
||||
],
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
perform_redaction(details, None)
|
||||
|
||||
response = details["standard_logging_object"]["response"]
|
||||
assert response["choices"][0]["message"]["content"] == "redacted-by-litellm"
|
||||
assert response["vertex_ai_grounding_metadata"] == []
|
||||
assert response["vertex_ai_url_context_metadata"] == []
|
||||
|
||||
def test_redacts_vertex_provider_metadata_on_streaming_model_response(self):
|
||||
response = litellm.ModelResponse(
|
||||
id="resp-1",
|
||||
choices=[
|
||||
litellm.Choices(
|
||||
message=litellm.Message(
|
||||
content="sensitive answer",
|
||||
role="assistant",
|
||||
)
|
||||
)
|
||||
],
|
||||
model="gemini-2.5-flash",
|
||||
)
|
||||
setattr(
|
||||
response,
|
||||
"vertex_ai_grounding_metadata",
|
||||
[{"webSearchQueries": ["sensitive search term"]}],
|
||||
)
|
||||
response._hidden_params["vertex_ai_grounding_metadata"] = [
|
||||
{"webSearchQueries": ["sensitive search term"]}
|
||||
]
|
||||
|
||||
details = {
|
||||
"stream": True,
|
||||
"complete_streaming_response": response,
|
||||
}
|
||||
|
||||
perform_redaction(details, response)
|
||||
|
||||
assert response.choices[0].message.content == "redacted-by-litellm"
|
||||
assert getattr(response, "vertex_ai_grounding_metadata") == []
|
||||
assert "vertex_ai_grounding_metadata" not in response._hidden_params
|
||||
|
||||
def test_redacts_vertex_provider_metadata_from_metadata_hidden_params(self):
|
||||
"""Streaming success_handler copies _hidden_params into metadata before redaction."""
|
||||
details = {
|
||||
"stream": True,
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"hidden_params": {
|
||||
"response_cost": 0.01,
|
||||
"vertex_ai_grounding_metadata": [
|
||||
{"webSearchQueries": ["sensitive search term"]}
|
||||
],
|
||||
"vertex_ai_url_context_metadata": [
|
||||
{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}
|
||||
],
|
||||
"vertex_ai_safety_ratings": [{"category": "HARM"}],
|
||||
"vertex_ai_citation_metadata": [{"citations": ["source"]}],
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
perform_redaction(details, None)
|
||||
|
||||
hidden_params = details["litellm_params"]["metadata"]["hidden_params"]
|
||||
assert hidden_params["response_cost"] == 0.01
|
||||
assert "vertex_ai_grounding_metadata" not in hidden_params
|
||||
assert "vertex_ai_url_context_metadata" not in hidden_params
|
||||
assert "vertex_ai_safety_ratings" not in hidden_params
|
||||
assert "vertex_ai_citation_metadata" not in hidden_params
|
||||
|
|
|
|||
|
|
@ -613,3 +613,153 @@ def test_stream_chunk_builder_dict_snapshot_preserves_hidden_provider_fields():
|
|||
assert (
|
||||
response._hidden_params["provider_specific_fields"]["traffic_type"] == "default"
|
||||
)
|
||||
|
||||
|
||||
def test_stream_chunk_builder_propagates_vertex_ai_metadata_from_chunks():
|
||||
"""Vertex AI metadata on streaming chunks must appear on assembled response."""
|
||||
grounding_metadata = [{"webSearchQueries": ["weather in SF"]}]
|
||||
url_context_metadata = [{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}]
|
||||
|
||||
chunk1 = ModelResponseStream(
|
||||
id="chatcmpl-vertex-1",
|
||||
created=1,
|
||||
model="gemini-2.5-flash",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(content="The weather", role="assistant"),
|
||||
)
|
||||
],
|
||||
)
|
||||
setattr(chunk1, "vertex_ai_grounding_metadata", grounding_metadata)
|
||||
chunk1._hidden_params["vertex_ai_grounding_metadata"] = grounding_metadata
|
||||
|
||||
chunk2 = ModelResponseStream(
|
||||
id="chatcmpl-vertex-1",
|
||||
created=1,
|
||||
model="gemini-2.5-flash",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(content=" is sunny.", role="assistant"),
|
||||
)
|
||||
],
|
||||
)
|
||||
setattr(chunk2, "vertex_ai_url_context_metadata", url_context_metadata)
|
||||
chunk2._hidden_params["vertex_ai_url_context_metadata"] = url_context_metadata
|
||||
|
||||
response = stream_chunk_builder(chunks=[chunk1, chunk2])
|
||||
assert response is not None
|
||||
assert getattr(response, "vertex_ai_grounding_metadata") == grounding_metadata
|
||||
assert getattr(response, "vertex_ai_url_context_metadata") == url_context_metadata
|
||||
assert response._hidden_params["vertex_ai_grounding_metadata"] == grounding_metadata
|
||||
assert (
|
||||
response._hidden_params["vertex_ai_url_context_metadata"]
|
||||
== url_context_metadata
|
||||
)
|
||||
|
||||
dumped = response.model_dump()
|
||||
assert dumped["vertex_ai_grounding_metadata"] == grounding_metadata
|
||||
assert dumped["vertex_ai_url_context_metadata"] == url_context_metadata
|
||||
|
||||
|
||||
def test_stream_chunk_builder_uses_assembled_model_for_provider_metadata():
|
||||
grounding_metadata = [{"webSearchQueries": ["weather in SF"]}]
|
||||
|
||||
chunk1 = ModelResponseStream(
|
||||
id="chatcmpl-vertex-router",
|
||||
created=1,
|
||||
model="gpt-4o",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(content="The weather", role="assistant"),
|
||||
)
|
||||
],
|
||||
)
|
||||
chunk2 = ModelResponseStream(
|
||||
id="chatcmpl-vertex-router",
|
||||
created=1,
|
||||
model="gemini-2.5-flash",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(content=" is sunny.", role=None),
|
||||
)
|
||||
],
|
||||
)
|
||||
setattr(chunk2, "vertex_ai_grounding_metadata", grounding_metadata)
|
||||
chunk2._hidden_params["vertex_ai_grounding_metadata"] = grounding_metadata
|
||||
|
||||
response = stream_chunk_builder(chunks=[chunk1, chunk2])
|
||||
assert response is not None
|
||||
assert response.model == "gemini-2.5-flash"
|
||||
assert getattr(response, "vertex_ai_grounding_metadata") == grounding_metadata
|
||||
|
||||
|
||||
def test_stream_chunk_builder_propagates_vertex_ai_safety_results():
|
||||
"""Assembled response must expose safety data under the non-streaming field name."""
|
||||
safety_ratings = [
|
||||
[{"category": "HARM_CATEGORY_HATE_SPEECH", "probability": "NEGLIGIBLE"}]
|
||||
]
|
||||
|
||||
chunk = ModelResponseStream(
|
||||
id="chatcmpl-vertex-safety",
|
||||
created=1,
|
||||
model="gemini-2.5-flash",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(content="hello", role="assistant"),
|
||||
)
|
||||
],
|
||||
)
|
||||
setattr(chunk, "vertex_ai_safety_ratings", safety_ratings)
|
||||
setattr(chunk, "vertex_ai_safety_results", safety_ratings)
|
||||
chunk._hidden_params["vertex_ai_safety_ratings"] = safety_ratings
|
||||
chunk._hidden_params["vertex_ai_safety_results"] = safety_ratings
|
||||
|
||||
response = stream_chunk_builder(chunks=[chunk])
|
||||
assert response is not None
|
||||
assert getattr(response, "vertex_ai_safety_results") == safety_ratings
|
||||
assert response._hidden_params["vertex_ai_safety_results"] == safety_ratings
|
||||
assert response.model_dump()["vertex_ai_safety_results"] == safety_ratings
|
||||
|
||||
|
||||
def test_stream_chunk_builder_propagates_vertex_ai_metadata_from_dict_chunks():
|
||||
"""Dict snapshot chunks (model_dump) should also propagate Vertex AI metadata."""
|
||||
chunk_dict = ModelResponseStream(
|
||||
id="chatcmpl-vertex-2",
|
||||
created=1,
|
||||
model="gemini-2.5-flash",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(content="hello", role="assistant"),
|
||||
)
|
||||
],
|
||||
).model_dump()
|
||||
chunk_dict["_hidden_params"] = {
|
||||
"vertex_ai_grounding_metadata": [{"webSearchQueries": ["test query"]}]
|
||||
}
|
||||
|
||||
response = stream_chunk_builder(chunks=[chunk_dict])
|
||||
assert response is not None
|
||||
assert getattr(response, "vertex_ai_grounding_metadata") == [
|
||||
{"webSearchQueries": ["test query"]}
|
||||
]
|
||||
assert response.model_dump()["vertex_ai_grounding_metadata"] == [
|
||||
{"webSearchQueries": ["test query"]}
|
||||
]
|
||||
|
|
|
|||
|
|
@ -5092,6 +5092,175 @@ def test_map_tool_helper_collision_prefers_definitions_over_components_schemas()
|
|||
assert transformed["input_schema"]["properties"]["from_components"] == expected
|
||||
|
||||
|
||||
BILLING_HEADER_BLOCK = {
|
||||
"type": "text",
|
||||
"text": "x-anthropic-billing-header: cc_version=1.0.abc; cc_entrypoint=cli; cch=00000;",
|
||||
}
|
||||
|
||||
|
||||
def _system_with_billing_header(real_text: str) -> list:
|
||||
return [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [BILLING_HEADER_BLOCK, {"type": "text", "text": real_text}],
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_translate_system_message_keeps_billing_header_for_first_party_anthropic():
|
||||
config = AnthropicConfig()
|
||||
assert config.should_strip_billing_metadata() is False
|
||||
|
||||
result = config.translate_system_message(
|
||||
messages=_system_with_billing_header(
|
||||
"You are Claude Code, Anthropic's official CLI for Claude."
|
||||
)
|
||||
)
|
||||
|
||||
texts = [block["text"] for block in result]
|
||||
assert any(t.startswith("x-anthropic-billing-header:") for t in texts)
|
||||
assert "You are Claude Code, Anthropic's official CLI for Claude." in texts
|
||||
|
||||
|
||||
def test_translate_system_message_strips_billing_header_for_bedrock():
|
||||
from litellm.llms.bedrock.claude_platform.transformation import (
|
||||
BedrockClaudePlatformConfig,
|
||||
)
|
||||
|
||||
config = BedrockClaudePlatformConfig()
|
||||
assert config.should_strip_billing_metadata() is True
|
||||
|
||||
result = config.translate_system_message(
|
||||
messages=_system_with_billing_header("real system prompt")
|
||||
)
|
||||
|
||||
texts = [block["text"] for block in result]
|
||||
assert all(not t.startswith("x-anthropic-billing-header:") for t in texts)
|
||||
assert "real system prompt" in texts
|
||||
|
||||
|
||||
def test_anthropic_messages_request_keeps_billing_header_for_first_party():
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
config = AnthropicMessagesConfig()
|
||||
assert config.should_strip_billing_metadata() is False
|
||||
|
||||
optional_params = {
|
||||
"max_tokens": 16,
|
||||
"system": [
|
||||
BILLING_HEADER_BLOCK,
|
||||
{"type": "text", "text": "real system prompt"},
|
||||
],
|
||||
}
|
||||
result = config.transform_anthropic_messages_request(
|
||||
model="claude-3-5-sonnet-latest",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
anthropic_messages_optional_request_params=optional_params,
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
texts = [block["text"] for block in result["system"]]
|
||||
assert any(t.startswith("x-anthropic-billing-header:") for t in texts)
|
||||
|
||||
|
||||
def test_anthropic_messages_request_strips_billing_header_for_minimax():
|
||||
from litellm.llms.minimax.messages.transformation import MinimaxMessagesConfig
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
config = MinimaxMessagesConfig()
|
||||
assert config.should_strip_billing_metadata() is True
|
||||
|
||||
optional_params = {
|
||||
"max_tokens": 16,
|
||||
"system": [
|
||||
BILLING_HEADER_BLOCK,
|
||||
{"type": "text", "text": "real system prompt"},
|
||||
],
|
||||
}
|
||||
result = config.transform_anthropic_messages_request(
|
||||
model="MiniMax-M2",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
anthropic_messages_optional_request_params=optional_params,
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
texts = [block["text"] for block in result.get("system", [])]
|
||||
assert all(not t.startswith("x-anthropic-billing-header:") for t in texts)
|
||||
|
||||
|
||||
def test_translate_system_message_strips_billing_header_for_bedrock_invoke():
|
||||
from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
|
||||
AmazonAnthropicClaudeConfig,
|
||||
)
|
||||
|
||||
config = AmazonAnthropicClaudeConfig()
|
||||
assert config.should_strip_billing_metadata() is True
|
||||
|
||||
result = config.translate_system_message(
|
||||
messages=_system_with_billing_header("real system prompt")
|
||||
)
|
||||
|
||||
texts = [block["text"] for block in result]
|
||||
assert all(not t.startswith("x-anthropic-billing-header:") for t in texts)
|
||||
assert "real system prompt" in texts
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"module_path, class_name, expected_strip",
|
||||
[
|
||||
("litellm.llms.anthropic.chat.transformation", "AnthropicConfig", False),
|
||||
(
|
||||
"litellm.llms.anthropic.experimental_pass_through.messages.transformation",
|
||||
"AnthropicMessagesConfig",
|
||||
False,
|
||||
),
|
||||
(
|
||||
"litellm.llms.bedrock.claude_platform.transformation",
|
||||
"BedrockClaudePlatformConfig",
|
||||
True,
|
||||
),
|
||||
(
|
||||
"litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation",
|
||||
"AmazonAnthropicClaudeConfig",
|
||||
True,
|
||||
),
|
||||
(
|
||||
"litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.transformation",
|
||||
"VertexAIAnthropicConfig",
|
||||
True,
|
||||
),
|
||||
(
|
||||
"litellm.llms.azure_ai.anthropic.transformation",
|
||||
"AzureAnthropicConfig",
|
||||
True,
|
||||
),
|
||||
("litellm.llms.minimax.messages.transformation", "MinimaxMessagesConfig", True),
|
||||
(
|
||||
"litellm.llms.azure_ai.anthropic.messages_transformation",
|
||||
"AzureAnthropicMessagesConfig",
|
||||
True,
|
||||
),
|
||||
(
|
||||
"litellm.llms.deepseek.messages.transformation",
|
||||
"DeepSeekAnthropicMessagesConfig",
|
||||
True,
|
||||
),
|
||||
(
|
||||
"litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.experimental_pass_through.transformation",
|
||||
"VertexAIPartnerModelsAnthropicMessagesConfig",
|
||||
True,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_should_strip_billing_metadata_by_provider(
|
||||
module_path, class_name, expected_strip
|
||||
):
|
||||
import importlib
|
||||
|
||||
config_cls = getattr(importlib.import_module(module_path), class_name)
|
||||
assert config_cls().should_strip_billing_metadata() is expected_strip
|
||||
def test_namespace_tool_flat_nested_tools_are_extracted():
|
||||
"""Codex sends nested tools in flat format {type, name, description, parameters} with no 'function' wrapper.
|
||||
These must be normalized and mapped without raising KeyError: 'function'."""
|
||||
|
|
|
|||
|
|
@ -57,6 +57,22 @@ def test_azure_providers_image_generation_json_body_keeps_model():
|
|||
assert out == data
|
||||
|
||||
|
||||
def test_azure_image_generation_mai_base_model_uses_mai_url():
|
||||
azure_chat = AzureChatCompletion()
|
||||
url = azure_chat.create_azure_base_url(
|
||||
azure_client_params={
|
||||
"azure_endpoint": "https://my-resource.services.ai.azure.com",
|
||||
"api_version": "preview",
|
||||
},
|
||||
model="image-deployment-alias",
|
||||
base_model="MAI-Image-2.5",
|
||||
)
|
||||
assert (
|
||||
url
|
||||
== "https://my-resource.services.ai.azure.com/mai/v1/images/generations?api-version=preview"
|
||||
)
|
||||
|
||||
|
||||
def test_azure_image_generation_flattens_extra_body():
|
||||
"""
|
||||
Test that Azure image generation correctly flattens extra_body parameters.
|
||||
|
|
|
|||
|
|
@ -0,0 +1,171 @@
|
|||
import io
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../../.."))
|
||||
|
||||
from litellm.llms.azure_ai.image_edit import (
|
||||
AzureFoundryMAIImageEditConfig,
|
||||
get_azure_ai_image_edit_config,
|
||||
)
|
||||
from litellm.llms.azure_ai.image_generation.mai_transformation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
)
|
||||
|
||||
|
||||
class TestAzureMAIImageEdit:
|
||||
def test_get_mai_image_edit_url(self):
|
||||
url = AzureFoundryMAIImageGenerationConfig.get_mai_image_edit_url(
|
||||
api_base="https://my-resource.services.ai.azure.com",
|
||||
api_version="preview",
|
||||
)
|
||||
assert (
|
||||
url
|
||||
== "https://my-resource.services.ai.azure.com/mai/v1/images/edits?api-version=preview"
|
||||
)
|
||||
|
||||
def test_get_mai_image_edit_url_rewrites_generation_url(self):
|
||||
url = AzureFoundryMAIImageGenerationConfig.get_mai_image_edit_url(
|
||||
api_base=(
|
||||
"https://my-resource.services.ai.azure.com/mai/v1/images/generations"
|
||||
"?api-version=preview"
|
||||
),
|
||||
api_version="preview",
|
||||
)
|
||||
assert (
|
||||
url
|
||||
== "https://my-resource.services.ai.azure.com/mai/v1/images/edits?api-version=preview"
|
||||
)
|
||||
|
||||
def test_get_mai_image_edit_url_appends_edits_to_mai_root(self):
|
||||
url = AzureFoundryMAIImageGenerationConfig.get_mai_image_edit_url(
|
||||
api_base="https://my-resource.services.ai.azure.com/mai/v1",
|
||||
api_version="preview",
|
||||
)
|
||||
assert (
|
||||
url
|
||||
== "https://my-resource.services.ai.azure.com/mai/v1/images/edits?api-version=preview"
|
||||
)
|
||||
|
||||
def test_get_azure_ai_image_edit_config_returns_mai(self):
|
||||
config = get_azure_ai_image_edit_config("MAI-Image-2.5")
|
||||
assert isinstance(config, AzureFoundryMAIImageEditConfig)
|
||||
|
||||
def test_validate_environment_uses_api_key_header(self):
|
||||
config = AzureFoundryMAIImageEditConfig()
|
||||
headers: dict = {}
|
||||
config.validate_environment(headers, "MAI-Image-2.5", api_key="test-key")
|
||||
assert headers["api-key"] == "test-key"
|
||||
assert "Api-Key" not in headers
|
||||
|
||||
def test_get_complete_url(self):
|
||||
config = AzureFoundryMAIImageEditConfig()
|
||||
url = config.get_complete_url(
|
||||
model="MAI-Image-2.5",
|
||||
api_base="https://my-resource.services.ai.azure.com",
|
||||
litellm_params={"api_version": "preview"},
|
||||
)
|
||||
assert "/mai/v1/images/edits" in url
|
||||
assert "api-version=preview" in url
|
||||
|
||||
def test_map_openai_params_keeps_size(self):
|
||||
config = AzureFoundryMAIImageEditConfig()
|
||||
optional_params = config.map_openai_params(
|
||||
image_edit_optional_params={"size": "1792x1024", "n": 1},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
assert optional_params["size"] == "1792x1024"
|
||||
assert optional_params["n"] == 1
|
||||
assert "width" not in optional_params
|
||||
assert "height" not in optional_params
|
||||
|
||||
def test_map_openai_params_defaults_size(self):
|
||||
config = AzureFoundryMAIImageEditConfig()
|
||||
optional_params = config.map_openai_params(
|
||||
image_edit_optional_params={},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
assert optional_params["size"] == "1024x1024"
|
||||
|
||||
def test_map_openai_params_unsupported_size_raises(self):
|
||||
config = AzureFoundryMAIImageEditConfig()
|
||||
with pytest.raises(ValueError, match="Unsupported size value: 'auto'"):
|
||||
config.map_openai_params(
|
||||
image_edit_optional_params={"size": "auto"},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
def test_map_openai_params_invalid_size_format_raises(self):
|
||||
config = AzureFoundryMAIImageEditConfig()
|
||||
with pytest.raises(ValueError, match="Invalid size format: '1024xabc'"):
|
||||
config.map_openai_params(
|
||||
image_edit_optional_params={"size": "1024xabc"},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
def test_transform_image_edit_request_uses_image_field(self):
|
||||
config = AzureFoundryMAIImageEditConfig()
|
||||
image_bytes = io.BytesIO(b"fake-image-bytes")
|
||||
|
||||
data, files = config.transform_image_edit_request(
|
||||
model="MAI-Image-2.5",
|
||||
prompt="Turn this into a studio product shot",
|
||||
image=image_bytes,
|
||||
image_edit_optional_request_params={"size": "1024x1024", "n": 1},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert data["model"] == "MAI-Image-2.5"
|
||||
assert data["prompt"] == "Turn this into a studio product shot"
|
||||
assert data["size"] == "1024x1024"
|
||||
assert data["n"] == 1
|
||||
assert len(files) == 1
|
||||
assert files[0][0] == "image"
|
||||
assert files[0][0] != "image[]"
|
||||
|
||||
def test_normalize_mai_image_usage_maps_edit_response_fields(self):
|
||||
usage = AzureFoundryMAIImageGenerationConfig.normalize_mai_image_usage(
|
||||
{
|
||||
"num_output_tokens": 1024,
|
||||
"output_image_tokens": 1024,
|
||||
}
|
||||
)
|
||||
assert usage["output_tokens"] == 1024
|
||||
assert usage["input_tokens"] == 0
|
||||
assert usage["total_tokens"] == 1024
|
||||
assert usage["input_tokens_details"]["text_tokens"] == 0
|
||||
assert usage["input_tokens_details"]["image_tokens"] == 0
|
||||
|
||||
def test_transform_image_edit_response_parses_mai_usage(self):
|
||||
config = AzureFoundryMAIImageEditConfig()
|
||||
raw_response = MagicMock(spec=httpx.Response)
|
||||
raw_response.status_code = 200
|
||||
raw_response.text = ""
|
||||
raw_response.json.return_value = {
|
||||
"created": 1780897477,
|
||||
"data": [{"b64_json": "abc123"}],
|
||||
"usage": {
|
||||
"num_output_tokens": 1024,
|
||||
"output_image_tokens": 1024,
|
||||
},
|
||||
}
|
||||
|
||||
logging_obj = MagicMock()
|
||||
image_response = config.transform_image_edit_response(
|
||||
model="MAI-Image-2.5",
|
||||
raw_response=raw_response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
assert image_response.data[0].b64_json == "abc123"
|
||||
assert image_response.usage.output_tokens == 1024
|
||||
assert image_response.usage.total_tokens == 1024
|
||||
|
|
@ -0,0 +1,380 @@
|
|||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../../.."))
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure.azure import AzureChatCompletion
|
||||
from litellm.llms.azure.image_generation import get_azure_image_generation_config
|
||||
from litellm.llms.azure.image_generation.http_utils import (
|
||||
azure_deployment_image_generation_json_body,
|
||||
)
|
||||
from litellm.llms.azure_ai.image_generation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
get_azure_ai_image_generation_config,
|
||||
)
|
||||
from litellm.llms.azure_ai.image_generation.cost_calculator import (
|
||||
cost_calculator as azure_ai_image_cost_calculator,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
ImageObject,
|
||||
ImageResponse,
|
||||
ImageUsage,
|
||||
ImageUsageInputTokensDetails,
|
||||
)
|
||||
from litellm.utils import get_optional_params_image_gen
|
||||
|
||||
|
||||
class TestAzureMAIImageGeneration:
|
||||
def test_is_mai_model(self):
|
||||
assert AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-Image-2.5")
|
||||
assert AzureFoundryMAIImageGenerationConfig.is_mai_model(
|
||||
"azure_ai/MAI-Image-2.5"
|
||||
)
|
||||
assert AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-Image-2.5-Flash")
|
||||
assert AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-Image-2e")
|
||||
assert not AzureFoundryMAIImageGenerationConfig.is_mai_model("flux.2-pro")
|
||||
assert not AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-DS-R1")
|
||||
|
||||
def test_mai_flash_and_2e_model_pricing_in_cost_map(self):
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
flash_info = litellm.get_model_info(
|
||||
model="azure_ai/MAI-Image-2.5-Flash",
|
||||
custom_llm_provider="azure_ai",
|
||||
)
|
||||
assert flash_info["input_cost_per_token"] == 1.75e-06
|
||||
assert flash_info["input_cost_per_image_token"] == 1.75e-06
|
||||
assert flash_info["output_cost_per_image_token"] == 3.3e-05
|
||||
|
||||
image_2e_info = litellm.get_model_info(
|
||||
model="azure_ai/MAI-Image-2e",
|
||||
custom_llm_provider="azure_ai",
|
||||
)
|
||||
assert image_2e_info["input_cost_per_token"] == 5e-06
|
||||
assert image_2e_info["output_cost_per_image_token"] == 1.95e-05
|
||||
|
||||
def test_get_mai_image_generation_url(self):
|
||||
url = AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url(
|
||||
api_base="https://my-resource.services.ai.azure.com",
|
||||
api_version="preview",
|
||||
)
|
||||
assert (
|
||||
url
|
||||
== "https://my-resource.services.ai.azure.com/mai/v1/images/generations?api-version=preview"
|
||||
)
|
||||
|
||||
def test_get_mai_image_generation_url_preserves_full_path(self):
|
||||
api = (
|
||||
"https://my-resource.services.ai.azure.com/mai/v1/images/generations"
|
||||
"?api-version=preview"
|
||||
)
|
||||
url = AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url(
|
||||
api_base=api,
|
||||
api_version="preview",
|
||||
)
|
||||
assert url == api
|
||||
|
||||
def test_get_mai_image_generation_url_appends_generations_to_mai_root(self):
|
||||
url = AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url(
|
||||
api_base="https://my-resource.services.ai.azure.com/mai/v1",
|
||||
api_version="preview",
|
||||
)
|
||||
assert (
|
||||
url
|
||||
== "https://my-resource.services.ai.azure.com/mai/v1/images/generations?api-version=preview"
|
||||
)
|
||||
|
||||
def test_get_azure_ai_image_generation_config_returns_mai(self):
|
||||
config = get_azure_ai_image_generation_config("MAI-Image-2.5")
|
||||
assert isinstance(config, AzureFoundryMAIImageGenerationConfig)
|
||||
|
||||
def test_azure_image_generation_config_returns_mai(self):
|
||||
config = get_azure_image_generation_config("MAI-Image-2.5")
|
||||
assert isinstance(config, AzureFoundryMAIImageGenerationConfig)
|
||||
|
||||
def test_map_openai_params_size_to_width_height(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
optional_params = config.map_openai_params(
|
||||
non_default_params={"size": "1024x1024", "n": 1},
|
||||
optional_params={},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
assert optional_params["width"] == 1024
|
||||
assert optional_params["height"] == 1024
|
||||
assert optional_params["n"] == 1
|
||||
assert "size" not in optional_params
|
||||
|
||||
def test_map_openai_params_defaults(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
optional_params = config.map_openai_params(
|
||||
non_default_params={},
|
||||
optional_params={},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
assert optional_params["width"] == 1024
|
||||
assert optional_params["height"] == 1024
|
||||
|
||||
def test_get_optional_params_image_gen_mai(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
optional_params = get_optional_params_image_gen(
|
||||
model="MAI-Image-2.5",
|
||||
size="1792x1024",
|
||||
n=1,
|
||||
custom_llm_provider="azure_ai",
|
||||
provider_config=config,
|
||||
drop_params=True,
|
||||
)
|
||||
assert optional_params["width"] == 1792
|
||||
assert optional_params["height"] == 1024
|
||||
assert "size" not in optional_params
|
||||
|
||||
def test_azure_create_azure_base_url_mai(self):
|
||||
azure_chat = AzureChatCompletion()
|
||||
url = azure_chat.create_azure_base_url(
|
||||
azure_client_params={
|
||||
"azure_endpoint": "https://my-resource.services.ai.azure.com",
|
||||
"api_version": "preview",
|
||||
},
|
||||
model="MAI-Image-2.5",
|
||||
)
|
||||
assert "/mai/v1/images/generations" in url
|
||||
assert "api-version=preview" in url
|
||||
|
||||
def test_mai_json_body_keeps_model(self):
|
||||
api = (
|
||||
"https://my-resource.services.ai.azure.com/mai/v1/images/generations"
|
||||
"?api-version=preview"
|
||||
)
|
||||
data = {
|
||||
"model": "MAI-Image-2.5",
|
||||
"prompt": "A photograph of a red fox",
|
||||
"width": 1024,
|
||||
"height": 1024,
|
||||
"n": 1,
|
||||
}
|
||||
out = azure_deployment_image_generation_json_body(api, data)
|
||||
assert out == data
|
||||
|
||||
def test_map_openai_params_custom_size(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
optional_params = config.map_openai_params(
|
||||
non_default_params={"size": "768x768"},
|
||||
optional_params={},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
assert optional_params["width"] == 768
|
||||
assert optional_params["height"] == 768
|
||||
|
||||
def test_map_openai_params_width_only_gets_height_default(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
optional_params = config.map_openai_params(
|
||||
non_default_params={"width": 1792},
|
||||
optional_params={},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
assert optional_params["width"] == 1792
|
||||
assert optional_params["height"] == config.DEFAULT_HEIGHT
|
||||
|
||||
def test_map_openai_params_height_only_gets_width_default(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
optional_params = config.map_openai_params(
|
||||
non_default_params={"height": 1792},
|
||||
optional_params={},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
assert optional_params["width"] == config.DEFAULT_WIDTH
|
||||
assert optional_params["height"] == 1792
|
||||
|
||||
def test_map_openai_params_unsupported_size_raises(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
with pytest.raises(ValueError, match="Unsupported size value: 'auto'"):
|
||||
config.map_openai_params(
|
||||
non_default_params={"size": "auto"},
|
||||
optional_params={},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
def test_map_openai_params_invalid_custom_size_raises(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
with pytest.raises(ValueError, match="Invalid size format: '1024xabc'"):
|
||||
config.map_openai_params(
|
||||
non_default_params={"size": "1024xabc"},
|
||||
optional_params={},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
def test_map_openai_params_unsupported_param_raises(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
with pytest.raises(ValueError, match="Parameter quality is not supported"):
|
||||
config.map_openai_params(
|
||||
non_default_params={"quality": "hd"},
|
||||
optional_params={},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
def test_transform_image_generation_response_normalizes_mai_usage(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
raw_response = MagicMock(spec=httpx.Response)
|
||||
raw_response.json.return_value = {
|
||||
"created": 1780897477,
|
||||
"data": [{"b64_json": "abc123"}],
|
||||
"usage": {
|
||||
"num_output_tokens": 1024,
|
||||
"num_input_text_tokens": 22,
|
||||
"output_image_tokens": 1024,
|
||||
},
|
||||
}
|
||||
|
||||
logging_obj = MagicMock()
|
||||
image_response = config.transform_image_generation_response(
|
||||
model="MAI-Image-2.5",
|
||||
raw_response=raw_response,
|
||||
model_response=ImageResponse(),
|
||||
logging_obj=logging_obj,
|
||||
request_data={"prompt": "A red fox"},
|
||||
optional_params={"width": 1024, "height": 1024},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert image_response.data[0].b64_json == "abc123"
|
||||
assert image_response.usage.output_tokens == 1024
|
||||
assert image_response.usage.input_tokens == 22
|
||||
assert image_response.usage.total_tokens == 1046
|
||||
|
||||
def test_transform_image_generation_response_non_json_raises_openai_error(self):
|
||||
from litellm.llms.openai.common_utils import OpenAIError
|
||||
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
raw_response = MagicMock(spec=httpx.Response)
|
||||
raw_response.json.side_effect = ValueError("not json")
|
||||
raw_response.text = "upstream gateway error"
|
||||
raw_response.status_code = 502
|
||||
|
||||
with pytest.raises(OpenAIError) as exc_info:
|
||||
config.transform_image_generation_response(
|
||||
model="MAI-Image-2.5",
|
||||
raw_response=raw_response,
|
||||
model_response=ImageResponse(),
|
||||
logging_obj=MagicMock(),
|
||||
request_data={"prompt": "A red fox"},
|
||||
optional_params={"width": 1024, "height": 1024},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 502
|
||||
assert exc_info.value.message == "upstream gateway error"
|
||||
|
||||
def test_normalize_mai_usage_preserves_zero_output_tokens(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
normalized = config.normalize_mai_image_usage(
|
||||
{
|
||||
"num_output_tokens": 0,
|
||||
"output_image_tokens": 1024,
|
||||
"num_input_text_tokens": 22,
|
||||
}
|
||||
)
|
||||
assert normalized["output_tokens"] == 0
|
||||
assert normalized["input_tokens"] == 22
|
||||
assert normalized["total_tokens"] == 22
|
||||
|
||||
def test_azure_sync_image_generation_uses_mai_response_transform(self):
|
||||
raw_response = MagicMock(spec=httpx.Response)
|
||||
raw_response.json.return_value = {
|
||||
"created": 1780897477,
|
||||
"data": [{"b64_json": "abc123"}],
|
||||
"usage": {
|
||||
"num_output_tokens": 1024,
|
||||
"num_input_text_tokens": 22,
|
||||
},
|
||||
}
|
||||
|
||||
class MAIImageGenerationAzureChatCompletion(AzureChatCompletion):
|
||||
def make_sync_azure_httpx_request(self, **kwargs):
|
||||
return raw_response
|
||||
|
||||
logging_obj = MagicMock()
|
||||
image_response = MAIImageGenerationAzureChatCompletion().image_generation(
|
||||
prompt="A red fox",
|
||||
timeout=60.0,
|
||||
optional_params={"width": 1792, "height": 1024},
|
||||
logging_obj=logging_obj,
|
||||
headers={},
|
||||
model="MAI-Image-2.5",
|
||||
api_key="test-key",
|
||||
api_base="https://my-resource.services.ai.azure.com",
|
||||
api_version="preview",
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert image_response.data[0].b64_json == "abc123"
|
||||
assert image_response.usage.output_tokens == 1024
|
||||
assert image_response.usage.input_tokens == 22
|
||||
assert image_response.usage.total_tokens == 1046
|
||||
assert image_response.size == "1792x1024"
|
||||
|
||||
def test_mai_image_cost_calculator_token_based(self):
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "azure_ai/MAI-Image-2.5"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="azure_ai")
|
||||
input_text_tokens = 100
|
||||
output_image_tokens = 1024
|
||||
|
||||
image_response = ImageResponse(
|
||||
data=[ImageObject(b64_json="img1")],
|
||||
usage=ImageUsage(
|
||||
input_tokens=input_text_tokens,
|
||||
input_tokens_details=ImageUsageInputTokensDetails(
|
||||
text_tokens=input_text_tokens,
|
||||
image_tokens=0,
|
||||
),
|
||||
output_tokens=output_image_tokens,
|
||||
total_tokens=input_text_tokens + output_image_tokens,
|
||||
),
|
||||
)
|
||||
|
||||
cost = azure_ai_image_cost_calculator(
|
||||
model=model,
|
||||
image_response=image_response,
|
||||
)
|
||||
|
||||
expected_cost = (
|
||||
input_text_tokens * model_info["input_cost_per_token"]
|
||||
+ output_image_tokens * model_info["output_cost_per_image_token"]
|
||||
)
|
||||
assert round(cost, 10) == round(expected_cost, 10)
|
||||
|
||||
def test_mai_image_cost_calculator_falls_back_to_flat_image_pricing(self):
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "azure_ai/MAI-Image-2.5"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="azure_ai")
|
||||
image_response = ImageResponse(
|
||||
data=[ImageObject(b64_json="img1"), ImageObject(b64_json="img2")]
|
||||
)
|
||||
|
||||
cost = azure_ai_image_cost_calculator(
|
||||
model=model,
|
||||
image_response=image_response,
|
||||
)
|
||||
|
||||
assert (
|
||||
cost == len(image_response.data or []) * model_info["output_cost_per_image"]
|
||||
)
|
||||
assert cost > 0
|
||||
|
|
@ -1467,11 +1467,10 @@ def test_transform_request_with_function_tool():
|
|||
)
|
||||
|
||||
# Verify the structure
|
||||
assert "additionalModelRequestFields" in request_data
|
||||
additional_fields = request_data["additionalModelRequestFields"]
|
||||
# Function tools are not computer use tools, so they don't get anthropic_beta —
|
||||
# additionalModelRequestFields should be absent (not serialized as empty {})
|
||||
assert "additionalModelRequestFields" not in request_data
|
||||
|
||||
# Function tools are not computer use tools, so they don't get anthropic_beta
|
||||
# They are processed through the regular tool config
|
||||
assert "toolConfig" in request_data
|
||||
assert "tools" in request_data["toolConfig"]
|
||||
assert len(request_data["toolConfig"]["tools"]) == 1
|
||||
|
|
|
|||
|
|
@ -128,6 +128,70 @@ class TestBedrockFilesTransformation:
|
|||
# Must have messages
|
||||
assert "messages" in model_input
|
||||
|
||||
# Nova Pro rejects empty additionalModelRequestFields / system — they must be absent
|
||||
assert (
|
||||
"additionalModelRequestFields" not in model_input
|
||||
), "Nova: empty additionalModelRequestFields must be omitted, not serialized as {}"
|
||||
assert (
|
||||
"system" not in model_input
|
||||
), "Nova: empty system must be omitted, not serialized as []"
|
||||
|
||||
def test_nova_batch_jsonl_omits_empty_converse_fields(self):
|
||||
"""
|
||||
Regression test: Amazon Nova Pro returns 400 Malformed input request when
|
||||
additionalModelRequestFields or system are present but empty in the Converse
|
||||
API payload. The proxy must strip these keys when they carry no data.
|
||||
"""
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
config = BedrockFilesConfig()
|
||||
|
||||
openai_jsonl_content = [
|
||||
{
|
||||
"custom_id": "req-0",
|
||||
"method": "POST",
|
||||
"url": "/v1/chat/completions",
|
||||
"body": {
|
||||
"model": "us.amazon.nova-pro-v1:0",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What is 1 + 1? Answer with just the number.",
|
||||
}
|
||||
],
|
||||
"max_tokens": 16,
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
|
||||
openai_jsonl_content
|
||||
)
|
||||
|
||||
assert len(result) == 1
|
||||
model_input = result[0]["modelInput"]
|
||||
|
||||
assert (
|
||||
"additionalModelRequestFields" not in model_input
|
||||
or model_input["additionalModelRequestFields"]
|
||||
), "additionalModelRequestFields must be absent or non-empty — Nova rejects {}"
|
||||
assert (
|
||||
"system" not in model_input or model_input["system"]
|
||||
), "system must be absent or non-empty — Nova rejects []"
|
||||
|
||||
# Validate the exact shape AWS accepts
|
||||
assert model_input == {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"text": "What is 1 + 1? Answer with just the number."}
|
||||
],
|
||||
}
|
||||
],
|
||||
"inferenceConfig": {"maxTokens": 16},
|
||||
}
|
||||
|
||||
def test_nova_image_content_uses_converse_image_blocks(self):
|
||||
"""
|
||||
Test that image_url content blocks are converted to Bedrock Converse
|
||||
|
|
|
|||
|
|
@ -12,6 +12,11 @@ import sys
|
|||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
|
||||
import pytest
|
||||
from botocore.exceptions import (
|
||||
ConnectTimeoutError,
|
||||
PartialCredentialsError,
|
||||
ProfileNotFound,
|
||||
)
|
||||
|
||||
import litellm
|
||||
from litellm.llms.bedrock_mantle.responses.transformation import (
|
||||
|
|
@ -114,16 +119,15 @@ class TestBedrockMantleResponsesAuth:
|
|||
)
|
||||
assert headers["Authorization"] == "Bearer bearer-key"
|
||||
|
||||
def test_missing_key_raises(self, monkeypatch):
|
||||
def test_missing_bearer_does_not_raise_in_validate_environment(self, monkeypatch):
|
||||
# SigV4 may still apply, so validate_environment must defer instead of raising.
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
cfg = BedrockMantleResponsesAPIConfig()
|
||||
with pytest.raises(ValueError, match="Bedrock Mantle API key"):
|
||||
cfg.validate_environment(
|
||||
headers={},
|
||||
model="openai.gpt-5.5",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
)
|
||||
headers = cfg.validate_environment(
|
||||
headers={}, model="openai.gpt-5.5", litellm_params=GenericLiteLLMParams()
|
||||
)
|
||||
assert "Authorization" not in headers
|
||||
|
||||
def test_custom_llm_provider(self):
|
||||
cfg = BedrockMantleResponsesAPIConfig()
|
||||
|
|
@ -261,6 +265,386 @@ def local_cost_map(monkeypatch):
|
|||
litellm.get_model_info.cache_clear()
|
||||
|
||||
|
||||
class TestBedrockMantleResponsesSigV4:
|
||||
def test_bearer_short_circuits_without_credentials(self, monkeypatch):
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
|
||||
signer = BaseAWSLLM()
|
||||
signer.get_credentials = MagicMock(
|
||||
side_effect=AssertionError("get_credentials must not run for bearer auth")
|
||||
)
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer)
|
||||
|
||||
headers, signed_body = cfg.sign_request(
|
||||
headers={},
|
||||
optional_params={},
|
||||
request_data={"input": "hi"},
|
||||
api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses",
|
||||
api_key="bearer-from-config",
|
||||
)
|
||||
assert headers["Authorization"] == "Bearer bearer-from-config"
|
||||
assert signed_body == b'{"input": "hi"}'
|
||||
signer.get_credentials.assert_not_called()
|
||||
|
||||
def test_bearer_resolved_from_mantle_env_key(self, monkeypatch):
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "env-bearer")
|
||||
|
||||
signer = BaseAWSLLM()
|
||||
signer.get_credentials = MagicMock(
|
||||
side_effect=AssertionError("get_credentials must not run for bearer auth")
|
||||
)
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer)
|
||||
|
||||
headers, _ = cfg.sign_request(
|
||||
headers={},
|
||||
optional_params={},
|
||||
request_data={"input": "hi"},
|
||||
api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses",
|
||||
api_key=None,
|
||||
)
|
||||
assert headers["Authorization"] == "Bearer env-bearer"
|
||||
|
||||
def test_bearer_arg_takes_priority_over_mantle_env_key(self, monkeypatch):
|
||||
# The passed api_key (e.g. litellm_params.api_key) must win over the env
|
||||
# bearer; a reordered precedence chain would silently use the wrong token.
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "env-bearer")
|
||||
|
||||
signer = BaseAWSLLM()
|
||||
signer.get_credentials = MagicMock(
|
||||
side_effect=AssertionError("get_credentials must not run for bearer auth")
|
||||
)
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer)
|
||||
|
||||
headers, _ = cfg.sign_request(
|
||||
headers={},
|
||||
optional_params={},
|
||||
request_data={"input": "hi"},
|
||||
api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses",
|
||||
api_key="arg-bearer",
|
||||
)
|
||||
assert headers["Authorization"] == "Bearer arg-bearer"
|
||||
signer.get_credentials.assert_not_called()
|
||||
|
||||
def test_access_key_produces_sigv4_headers(self, monkeypatch):
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM())
|
||||
headers, signed_body = cfg.sign_request(
|
||||
headers={},
|
||||
optional_params={
|
||||
"aws_access_key_id": "AKIAEXAMPLE",
|
||||
"aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0",
|
||||
"aws_session_token": "session-token-test",
|
||||
"aws_region_name": "us-east-2",
|
||||
},
|
||||
request_data={"input": "hi"},
|
||||
api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses",
|
||||
api_key=None,
|
||||
)
|
||||
assert headers["Authorization"].startswith("AWS4-HMAC-SHA256")
|
||||
assert "Credential=AKIAEXAMPLE/" in headers["Authorization"]
|
||||
assert "/us-east-2/bedrock/aws4_request" in headers["Authorization"]
|
||||
assert "X-Amz-Date" in headers
|
||||
assert headers["X-Amz-Security-Token"] == "session-token-test"
|
||||
assert signed_body == b'{"input": "hi"}'
|
||||
|
||||
def test_assume_role_path_produces_sigv4_headers(self, monkeypatch):
|
||||
from unittest.mock import MagicMock
|
||||
from botocore.credentials import Credentials
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
|
||||
signer = BaseAWSLLM()
|
||||
signer.get_credentials = MagicMock(
|
||||
return_value=Credentials(
|
||||
access_key="ASIAEXAMPLE",
|
||||
secret_key="YXNzdW1lZC1yb2xlLXNlY3JldC1hc3N1bWVk",
|
||||
token="assumed-session-token",
|
||||
)
|
||||
)
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer)
|
||||
|
||||
headers, _ = cfg.sign_request(
|
||||
headers={},
|
||||
optional_params={
|
||||
"aws_role_name": "arn:aws:iam::000000000000:role/test-role",
|
||||
"aws_session_name": "litellm-test",
|
||||
"aws_region_name": "us-east-2",
|
||||
},
|
||||
request_data={"input": "hi"},
|
||||
api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses",
|
||||
api_key=None,
|
||||
)
|
||||
signer.get_credentials.assert_called_once()
|
||||
call = signer.get_credentials.call_args.kwargs
|
||||
assert call["aws_role_name"] == "arn:aws:iam::000000000000:role/test-role"
|
||||
assert call["aws_session_name"] == "litellm-test"
|
||||
assert headers["Authorization"].startswith("AWS4-HMAC-SHA256")
|
||||
assert "/us-east-2/bedrock/aws4_request" in headers["Authorization"]
|
||||
|
||||
def test_signed_body_matches_final_data_after_normalize(self, monkeypatch):
|
||||
"""Core regression: the signed bytes must equal the bytes actually sent.
|
||||
|
||||
Sign the *final* data dict and assert the returned signed_body decodes to
|
||||
exactly that dict, so a later change to the data would break the SigV4 hash.
|
||||
"""
|
||||
import json
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
|
||||
final_data = {"model": "openai.gpt-5.5", "input": "hi", "max_output_tokens": 16}
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM())
|
||||
_, signed_body = cfg.sign_request(
|
||||
headers={},
|
||||
optional_params={
|
||||
"aws_access_key_id": "AKIAEXAMPLE",
|
||||
"aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0",
|
||||
"aws_region_name": "us-east-2",
|
||||
},
|
||||
request_data=final_data,
|
||||
api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses",
|
||||
api_key=None,
|
||||
)
|
||||
assert signed_body is not None
|
||||
assert json.loads(signed_body) == final_data
|
||||
|
||||
def test_region_comes_from_optional_params(self, monkeypatch):
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
monkeypatch.delenv("AWS_REGION", raising=False)
|
||||
monkeypatch.delenv("AWS_REGION_NAME", raising=False)
|
||||
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM())
|
||||
headers, _ = cfg.sign_request(
|
||||
headers={},
|
||||
optional_params={
|
||||
"aws_access_key_id": "AKIAEXAMPLE",
|
||||
"aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0",
|
||||
"aws_region_name": "eu-west-1",
|
||||
},
|
||||
request_data={"input": "hi"},
|
||||
api_base="https://bedrock-mantle.eu-west-1.api.aws/openai/v1/responses",
|
||||
api_key=None,
|
||||
)
|
||||
assert "/eu-west-1/bedrock/aws4_request" in headers["Authorization"]
|
||||
|
||||
def test_url_region_and_sigv4_region_agree_from_litellm_params(self, monkeypatch):
|
||||
"""Adversarial-review regression: a caller-supplied aws_region_name (no region
|
||||
env set) must shape BOTH the URL host and the SigV4 credential scope, or the
|
||||
request is signed for one region and sent to another -> 401.
|
||||
"""
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_REGION", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
|
||||
monkeypatch.delenv("AWS_REGION", raising=False)
|
||||
monkeypatch.delenv("AWS_REGION_NAME", raising=False)
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
params = {
|
||||
"aws_region_name": "ap-southeast-2",
|
||||
"aws_access_key_id": "AKIAEXAMPLE",
|
||||
"aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0",
|
||||
}
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM())
|
||||
url = cfg.get_complete_url(api_base=None, litellm_params=params)
|
||||
assert (
|
||||
url == "https://bedrock-mantle.ap-southeast-2.api.aws/openai/v1/responses"
|
||||
)
|
||||
|
||||
headers, _ = cfg.sign_request(
|
||||
headers={},
|
||||
optional_params=params,
|
||||
request_data={"input": "hi"},
|
||||
api_base=url,
|
||||
api_key=None,
|
||||
)
|
||||
assert "/ap-southeast-2/bedrock/aws4_request" in headers["Authorization"]
|
||||
|
||||
def test_injected_default_region_base_does_not_override_aws_region_name(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""2nd-round adversarial regression: responses/main.py auto-injects
|
||||
litellm_params.api_base = https://bedrock-mantle.<DEFAULT>.api.aws/v1 (default
|
||||
region, ignoring aws_region_name). The config must still pin BOTH the URL host
|
||||
and the SigV4 scope to aws_region_name, or the IAM deployment 401s. A naive
|
||||
'resolve region only when api_base is None' fix would fail this test.
|
||||
"""
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_REGION", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
|
||||
monkeypatch.delenv("AWS_REGION", raising=False)
|
||||
monkeypatch.delenv("AWS_REGION_NAME", raising=False)
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
injected_base = "https://bedrock-mantle.us-east-1.api.aws/v1" # default region
|
||||
params = {
|
||||
"aws_region_name": "us-east-2", # what the caller actually wants
|
||||
"api_base": injected_base,
|
||||
"aws_access_key_id": "AKIAEXAMPLE",
|
||||
"aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0",
|
||||
}
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM())
|
||||
url = cfg.get_complete_url(api_base=injected_base, litellm_params=params)
|
||||
assert url == "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses"
|
||||
|
||||
headers, _ = cfg.sign_request(
|
||||
headers={},
|
||||
optional_params=params,
|
||||
request_data={"input": "hi"},
|
||||
api_base=url,
|
||||
api_key=None,
|
||||
)
|
||||
assert "/us-east-2/bedrock/aws4_request" in headers["Authorization"]
|
||||
assert "us-east-1" not in headers["Authorization"]
|
||||
|
||||
def test_custom_proxy_host_is_preserved(self, monkeypatch):
|
||||
"""A genuinely custom (non-Mantle) api_base host must be preserved, not rewritten
|
||||
to a bedrock-mantle host. Only standard Mantle hosts are region-pinned.
|
||||
"""
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
|
||||
cfg = BedrockMantleResponsesAPIConfig()
|
||||
url = cfg.get_complete_url(
|
||||
api_base="https://mantle-proxy.internal.example/openai/v1",
|
||||
litellm_params={"aws_region_name": "us-east-2"},
|
||||
)
|
||||
assert url == "https://mantle-proxy.internal.example/openai/v1/responses"
|
||||
|
||||
def test_caller_authorization_does_not_override_sigv4(self, monkeypatch):
|
||||
"""Adversarial-review regression: a caller-supplied Authorization header (e.g.
|
||||
from extra_headers, surviving the relaxed validate_environment) must not clobber
|
||||
the SigV4 Authorization that _sign_request would otherwise restore.
|
||||
"""
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM())
|
||||
headers, _ = cfg.sign_request(
|
||||
headers={"Authorization": "Bearer stale-caller-token"},
|
||||
optional_params={
|
||||
"aws_access_key_id": "AKIAEXAMPLE",
|
||||
"aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0",
|
||||
"aws_region_name": "us-east-2",
|
||||
},
|
||||
request_data={"input": "hi"},
|
||||
api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses",
|
||||
api_key=None,
|
||||
)
|
||||
assert headers["Authorization"].startswith("AWS4-HMAC-SHA256")
|
||||
assert "Bearer stale-caller-token" not in headers["Authorization"]
|
||||
|
||||
def test_no_bearer_and_no_credentials_raises_both_paths(self, monkeypatch):
|
||||
from unittest.mock import MagicMock
|
||||
from botocore.exceptions import NoCredentialsError
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
|
||||
signer = BaseAWSLLM()
|
||||
signer.get_credentials = MagicMock(side_effect=NoCredentialsError())
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer)
|
||||
|
||||
with pytest.raises(ValueError) as exc:
|
||||
cfg.sign_request(
|
||||
headers={},
|
||||
optional_params={"aws_region_name": "us-east-2"},
|
||||
request_data={"input": "hi"},
|
||||
api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses",
|
||||
api_key=None,
|
||||
)
|
||||
msg = str(exc.value)
|
||||
assert "Bearer" in msg
|
||||
assert "SigV4" in msg or "IAM" in msg
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cred_error",
|
||||
[
|
||||
PartialCredentialsError(provider="env", cred_var="aws_secret_access_key"),
|
||||
ProfileNotFound(profile="missing-profile"),
|
||||
],
|
||||
)
|
||||
def test_partial_credentials_raises_both_paths(self, monkeypatch, cred_error):
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
|
||||
signer = BaseAWSLLM()
|
||||
signer.get_credentials = MagicMock(side_effect=cred_error)
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer)
|
||||
|
||||
with pytest.raises(ValueError) as exc:
|
||||
cfg.sign_request(
|
||||
headers={},
|
||||
optional_params={"aws_region_name": "us-east-2"},
|
||||
request_data={"input": "hi"},
|
||||
api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses",
|
||||
api_key=None,
|
||||
)
|
||||
msg = str(exc.value)
|
||||
assert "Bearer" in msg
|
||||
assert "SigV4" in msg or "IAM" in msg
|
||||
|
||||
def test_sts_transport_error_is_not_masked_as_credentials(self, monkeypatch):
|
||||
# An AssumeRole / web-identity flow hits STS over the network, so a transient
|
||||
# connection error must surface as itself, not be rewritten into the
|
||||
# "no usable AWS credentials" message that would send the user to fix the
|
||||
# wrong thing.
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False)
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
|
||||
signer = BaseAWSLLM()
|
||||
signer.get_credentials = MagicMock(
|
||||
side_effect=ConnectTimeoutError(
|
||||
endpoint_url="https://sts.us-east-2.amazonaws.com"
|
||||
)
|
||||
)
|
||||
cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer)
|
||||
|
||||
with pytest.raises(ConnectTimeoutError):
|
||||
cfg.sign_request(
|
||||
headers={},
|
||||
optional_params={
|
||||
"aws_role_name": "arn:aws:iam::000000000000:role/test-role",
|
||||
"aws_region_name": "us-east-2",
|
||||
},
|
||||
request_data={"input": "hi"},
|
||||
api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses",
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
|
||||
class TestBedrockMantleResponsesPricing:
|
||||
def test_gpt_5_5_pricing_and_mode(self, local_cost_map):
|
||||
info = litellm.get_model_info("bedrock_mantle/openai.gpt-5.5")
|
||||
|
|
|
|||
|
|
@ -742,3 +742,241 @@ async def test_anthropic_post_retry_reserializes_mutated_body():
|
|||
assert first_sent == prebuilt # attempt 0 used prebuilt
|
||||
assert second_sent == _json.dumps(request_body) # attempt 1 re-serialized
|
||||
assert "MUTATED" in second_sent # ... the mutated body
|
||||
|
||||
|
||||
def test_base_responses_config_sign_request_is_noop_by_default():
|
||||
"""Default responses sign_request must be a no-op: unchanged headers, no signed body.
|
||||
|
||||
Guards the 15 existing responses providers from accidental signing when the
|
||||
handler starts calling sign_request.
|
||||
"""
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
|
||||
cfg = OpenAIResponsesAPIConfig()
|
||||
headers = {"Authorization": "Bearer sk-existing"}
|
||||
out_headers, signed_body = cfg.sign_request(
|
||||
headers=headers,
|
||||
optional_params={},
|
||||
request_data={"input": "hi"},
|
||||
api_base="https://api.openai.com/v1/responses",
|
||||
)
|
||||
assert out_headers == {"Authorization": "Bearer sk-existing"}
|
||||
assert signed_body is None
|
||||
|
||||
|
||||
def _make_responses_handler_call(signed_body):
|
||||
"""Drive BaseLLMHTTPHandler.response_api_handler with a fully mocked provider
|
||||
config + sync client, returning the kwargs the client.post was called with.
|
||||
|
||||
signed_body=None simulates a no-op (non-signing) provider; bytes simulates a
|
||||
signing provider (e.g. Bedrock Mantle).
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
provider_config = MagicMock()
|
||||
provider_config.validate_environment.return_value = {}
|
||||
provider_config.get_complete_url.return_value = (
|
||||
"https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses"
|
||||
)
|
||||
provider_config.transform_responses_api_request.return_value = {"input": "hi"}
|
||||
provider_config.should_fake_stream.return_value = False
|
||||
provider_config.sign_request.return_value = ({"X-Signed": "1"}, signed_body)
|
||||
|
||||
mock_client = MagicMock(spec=HTTPHandler)
|
||||
mock_client.post.return_value = MagicMock()
|
||||
|
||||
handler = BaseLLMHTTPHandler()
|
||||
handler.response_api_handler(
|
||||
model="openai.gpt-5.5",
|
||||
input="hi",
|
||||
responses_api_provider_config=provider_config,
|
||||
response_api_optional_request_params={},
|
||||
custom_llm_provider="bedrock_mantle",
|
||||
litellm_params=GenericLiteLLMParams(aws_region_name="us-east-2"),
|
||||
logging_obj=MagicMock(),
|
||||
client=mock_client,
|
||||
_is_async=False,
|
||||
)
|
||||
return mock_client.post.call_args.kwargs
|
||||
|
||||
|
||||
def test_responses_handler_sends_json_when_not_signed():
|
||||
"""No-op provider (signed_body is None) -> handler posts json=data, no data= bytes."""
|
||||
kwargs = _make_responses_handler_call(signed_body=None)
|
||||
assert kwargs.get("json") == {"input": "hi"}
|
||||
assert "data" not in kwargs
|
||||
|
||||
|
||||
def test_responses_handler_sends_signed_bytes_when_signed():
|
||||
"""Signing provider -> handler posts the exact signed bytes via data=, not json=."""
|
||||
kwargs = _make_responses_handler_call(signed_body=b'{"input": "hi"}')
|
||||
assert kwargs.get("data") == b'{"input": "hi"}'
|
||||
assert "json" not in kwargs
|
||||
assert kwargs["headers"] == {"X-Signed": "1"}
|
||||
|
||||
|
||||
def test_responses_handler_signs_after_fake_stream_prep_strips_stream():
|
||||
"""Fake-stream signing-order invariant: the bytes SIGNED must equal the bytes SENT.
|
||||
|
||||
In the streaming + fake-stream path the handler first runs
|
||||
_prepare_fake_stream_request, which pops "stream" out of the body, and only
|
||||
then calls sign_request. If signing ran before that pop, the signed body
|
||||
would still carry "stream" while the body sent over the wire would not,
|
||||
producing a SigV4 payload-hash mismatch (401) for a real Mantle deployment.
|
||||
We snapshot request_data at sign time and assert "stream" is already gone.
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
provider_config = MagicMock()
|
||||
provider_config.validate_environment.return_value = {}
|
||||
provider_config.get_complete_url.return_value = (
|
||||
"https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses"
|
||||
)
|
||||
provider_config.transform_responses_api_request.return_value = {
|
||||
"input": "hi",
|
||||
"stream": True,
|
||||
}
|
||||
provider_config.should_fake_stream.return_value = True
|
||||
provider_config.transform_response_api_response.return_value = ResponsesAPIResponse(
|
||||
id="resp_1",
|
||||
created_at=0,
|
||||
output=[],
|
||||
status="completed",
|
||||
model="openai.gpt-5.5",
|
||||
)
|
||||
|
||||
captured = {}
|
||||
|
||||
def _capture_sign(**kwargs):
|
||||
captured["request_data"] = dict(kwargs["request_data"])
|
||||
return ({"X-Signed": "1"}, b'{"input": "hi"}')
|
||||
|
||||
provider_config.sign_request.side_effect = _capture_sign
|
||||
|
||||
mock_client = MagicMock(spec=HTTPHandler)
|
||||
mock_client.post.return_value = MagicMock()
|
||||
|
||||
handler = BaseLLMHTTPHandler()
|
||||
handler.response_api_handler(
|
||||
model="openai.gpt-5.5",
|
||||
input="hi",
|
||||
responses_api_provider_config=provider_config,
|
||||
response_api_optional_request_params={"stream": True},
|
||||
custom_llm_provider="bedrock_mantle",
|
||||
litellm_params=GenericLiteLLMParams(aws_region_name="us-east-2"),
|
||||
logging_obj=MagicMock(),
|
||||
client=mock_client,
|
||||
_is_async=False,
|
||||
fake_stream=True,
|
||||
)
|
||||
|
||||
assert "stream" not in captured["request_data"]
|
||||
assert "input" in captured["request_data"]
|
||||
|
||||
post_kwargs = mock_client.post.call_args.kwargs
|
||||
assert post_kwargs.get("data") == b'{"input": "hi"}'
|
||||
assert "json" not in post_kwargs
|
||||
assert "stream" in post_kwargs
|
||||
|
||||
|
||||
def _make_compact_handler_call(signed_body, is_async):
|
||||
"""Drive (async_)compact_response_api_handler with a fully mocked provider config
|
||||
+ client, returning the kwargs the client.post was called with.
|
||||
|
||||
signed_body=None simulates a no-op (non-signing) provider; bytes simulates a
|
||||
signing provider (e.g. Bedrock Mantle SigV4 / bearer).
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
compact_url = "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses/compact"
|
||||
provider_config = MagicMock()
|
||||
provider_config.validate_environment.return_value = {}
|
||||
provider_config.get_complete_url.return_value = (
|
||||
"https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses"
|
||||
)
|
||||
provider_config.transform_compact_response_api_request.return_value = (
|
||||
compact_url,
|
||||
{"model": "openai.gpt-5.5", "input": "hi"},
|
||||
)
|
||||
provider_config.sign_request.return_value = ({"X-Signed": "1"}, signed_body)
|
||||
provider_config.transform_compact_response_api_response.return_value = "ok"
|
||||
|
||||
spec = AsyncHTTPHandler if is_async else HTTPHandler
|
||||
mock_client = MagicMock(spec=spec)
|
||||
if is_async:
|
||||
mock_client.post = AsyncMock(return_value=MagicMock())
|
||||
else:
|
||||
mock_client.post.return_value = MagicMock()
|
||||
|
||||
handler = BaseLLMHTTPHandler()
|
||||
result = handler.compact_response_api_handler(
|
||||
model="openai.gpt-5.5",
|
||||
input="hi",
|
||||
responses_api_provider_config=provider_config,
|
||||
response_api_optional_request_params={},
|
||||
custom_llm_provider="bedrock_mantle",
|
||||
litellm_params=GenericLiteLLMParams(aws_region_name="us-east-2"),
|
||||
logging_obj=MagicMock(),
|
||||
client=mock_client,
|
||||
_is_async=is_async,
|
||||
)
|
||||
if is_async:
|
||||
asyncio.run(result)
|
||||
return provider_config, mock_client.post.call_args.kwargs
|
||||
|
||||
|
||||
def test_compact_handler_sends_json_when_not_signed():
|
||||
"""No-op provider on compact (signed_body is None) -> posts json=data, no data= bytes."""
|
||||
provider_config, kwargs = _make_compact_handler_call(
|
||||
signed_body=None, is_async=False
|
||||
)
|
||||
provider_config.sign_request.assert_called_once()
|
||||
assert kwargs.get("json") == {"model": "openai.gpt-5.5", "input": "hi"}
|
||||
assert "data" not in kwargs
|
||||
|
||||
|
||||
def test_compact_handler_sends_signed_bytes_when_signed():
|
||||
"""Signing provider on compact -> posts the signed bytes via data=, not json=.
|
||||
|
||||
Regression for the adversarial-review finding that /responses/compact bypassed
|
||||
the SigV4 signing hook, so IAM-only Mantle callers sent unsigned bodies.
|
||||
"""
|
||||
provider_config, kwargs = _make_compact_handler_call(
|
||||
signed_body=b'{"model": "openai.gpt-5.5", "input": "hi"}', is_async=False
|
||||
)
|
||||
assert kwargs.get("data") == b'{"model": "openai.gpt-5.5", "input": "hi"}'
|
||||
assert "json" not in kwargs
|
||||
assert kwargs["headers"] == {"X-Signed": "1"}
|
||||
# signing must use the compact endpoint as api_base, not the create URL
|
||||
assert provider_config.sign_request.call_args.kwargs["api_base"].endswith(
|
||||
"/openai/v1/responses/compact"
|
||||
)
|
||||
|
||||
|
||||
def test_async_compact_handler_sends_signed_bytes_when_signed():
|
||||
"""Async compact must sign identically to sync (same omission in the async twin)."""
|
||||
provider_config, kwargs = _make_compact_handler_call(
|
||||
signed_body=b'{"model": "openai.gpt-5.5", "input": "hi"}', is_async=True
|
||||
)
|
||||
assert kwargs.get("data") == b'{"model": "openai.gpt-5.5", "input": "hi"}'
|
||||
assert "json" not in kwargs
|
||||
assert kwargs["headers"] == {"X-Signed": "1"}
|
||||
|
||||
|
||||
def test_async_compact_handler_sends_json_when_not_signed():
|
||||
"""Async no-op provider on compact -> posts json=data, no data= bytes."""
|
||||
_provider_config, kwargs = _make_compact_handler_call(
|
||||
signed_body=None, is_async=True
|
||||
)
|
||||
assert kwargs.get("json") == {"model": "openai.gpt-5.5", "input": "hi"}
|
||||
assert "data" not in kwargs
|
||||
|
|
|
|||
|
|
@ -127,12 +127,14 @@ def test_get_supported_openai_params_parallel_tool_calls():
|
|||
config = FireworksAIConfig()
|
||||
|
||||
supported_params = config.get_supported_openai_params(
|
||||
"fireworks_ai/accounts/fireworks/models/glm-4p6"
|
||||
"fireworks_ai/accounts/fireworks/models/glm-5p1"
|
||||
)
|
||||
assert "parallel_tool_calls" in supported_params
|
||||
assert "tools" in supported_params
|
||||
assert "tool_choice" in supported_params
|
||||
|
||||
unsupported_params = config.get_supported_openai_params(
|
||||
"fireworks_ai/accounts/fireworks/models/glm-5p1"
|
||||
"fireworks_ai/accounts/fireworks/models/llama-v3p1-8b-instruct"
|
||||
)
|
||||
assert "parallel_tool_calls" not in unsupported_params
|
||||
|
||||
|
|
@ -163,9 +165,9 @@ def test_get_model_info_respects_explicit_fireworks_capabilities():
|
|||
"""Test that get_model_info preserves explicit capability flags from the model map."""
|
||||
model_info = get_model_info("fireworks_ai/accounts/fireworks/models/glm-5p1")
|
||||
|
||||
assert model_info["supports_function_calling"] is False
|
||||
assert model_info["supports_function_calling"] is True
|
||||
assert model_info["supports_reasoning"] is True
|
||||
assert model_info["supports_tool_choice"] is False
|
||||
assert model_info["supports_tool_choice"] is True
|
||||
|
||||
|
||||
def test_get_provider_info_omits_false_supports_reasoning(monkeypatch):
|
||||
|
|
|
|||
|
|
@ -80,6 +80,46 @@ class TestGeminiTTSTransformation:
|
|||
assert "responseModalities" in result
|
||||
assert "AUDIO" in result["responseModalities"]
|
||||
|
||||
def test_gemini_tts_audio_parameter_mapping_with_language_code(self):
|
||||
config = GoogleAIStudioGeminiConfig()
|
||||
|
||||
non_default_params = {
|
||||
"audio": {"voice": "Kore", "format": "pcm16", "language_code": "en-US"}
|
||||
}
|
||||
optional_params = {}
|
||||
|
||||
result = config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model="gemini-2.5-flash-preview-tts",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert "speechConfig" in result
|
||||
assert result["speechConfig"]["languageCode"] == "en-US"
|
||||
assert (
|
||||
result["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"]
|
||||
== "Kore"
|
||||
)
|
||||
|
||||
def test_map_audio_params_language_code(self):
|
||||
config = GoogleAIStudioGeminiConfig()
|
||||
|
||||
result = config._map_audio_params(
|
||||
{"voice": "Kore", "format": "pcm16", "language_code": "de-DE"}
|
||||
)
|
||||
|
||||
assert result["languageCode"] == "de-DE"
|
||||
assert result["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore"
|
||||
|
||||
def test_map_audio_params_no_language_code(self):
|
||||
config = GoogleAIStudioGeminiConfig()
|
||||
|
||||
result = config._map_audio_params({"voice": "Kore", "format": "pcm16"})
|
||||
|
||||
assert "languageCode" not in result
|
||||
assert result["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore"
|
||||
|
||||
def test_gemini_tts_audio_parameter_with_existing_modalities(self):
|
||||
"""Test audio parameter mapping when modalities already exist"""
|
||||
config = GoogleAIStudioGeminiConfig()
|
||||
|
|
@ -328,5 +368,57 @@ class TestGeminiTTSSpeechConfigInRequestBody:
|
|||
assert "AUDIO" in generation_config["responseModalities"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,custom_llm_provider",
|
||||
[
|
||||
("gemini-2.5-flash-tts", "vertex_ai"),
|
||||
("gemini-2.5-flash-tts", "gemini"),
|
||||
("gemini-2.5-flash-preview-tts", "vertex_ai"),
|
||||
],
|
||||
)
|
||||
def test_language_code_end_to_end_mapping(self, model, custom_llm_provider):
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
from litellm.llms.vertex_ai.gemini.transformation import (
|
||||
_transform_request_body,
|
||||
)
|
||||
|
||||
config = VertexGeminiConfig()
|
||||
|
||||
non_default_params = {
|
||||
"audio": {"voice": "Puck", "format": "pcm16", "language_code": "pt-BR"}
|
||||
}
|
||||
optional_params = {}
|
||||
|
||||
mapped_params = config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert mapped_params["speechConfig"]["languageCode"] == "pt-BR"
|
||||
|
||||
request_body = _transform_request_body(
|
||||
messages=[{"role": "user", "content": "Hello world"}],
|
||||
model=model,
|
||||
optional_params=mapped_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params={},
|
||||
cached_content=None,
|
||||
)
|
||||
|
||||
generation_config = request_body["generationConfig"]
|
||||
assert generation_config["speechConfig"]["languageCode"] == "pt-BR"
|
||||
assert (
|
||||
generation_config["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"][
|
||||
"voiceName"
|
||||
]
|
||||
== "Puck"
|
||||
)
|
||||
assert "AUDIO" in generation_config["responseModalities"]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ from unittest.mock import patch, MagicMock
|
|||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
|
||||
import pytest
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
from litellm.llms.github_copilot.responses.transformation import (
|
||||
|
|
@ -22,13 +24,26 @@ from litellm.llms.github_copilot.responses.transformation import (
|
|||
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def use_local_model_cost_map(monkeypatch: pytest.MonkeyPatch):
|
||||
"""Pin litellm.model_cost to the bundled local backup so tests don't depend
|
||||
on remote catalog fetches (and don't change behavior across remote refreshes)."""
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
monkeypatch.setattr(
|
||||
litellm, "model_cost", get_model_cost_map(url=litellm.model_cost_map_url)
|
||||
)
|
||||
litellm.add_known_models(model_cost_map=litellm.model_cost)
|
||||
|
||||
|
||||
class TestGithubCopilotResponsesAPITransformation:
|
||||
"""Test GitHub Copilot Responses API configuration and transformations"""
|
||||
|
||||
def test_github_copilot_provider_config_registration(self):
|
||||
"""Test that GitHub Copilot provider returns GithubCopilotResponsesAPIConfig"""
|
||||
"""Test that GitHub Copilot provider returns the native Responses API
|
||||
config for a Responses-capable catalog model. Exercises the full stack:
|
||||
catalog lookup -> github_copilot_supports_responses_api -> native config."""
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model="github_copilot/gpt-5.1-codex",
|
||||
model="github_copilot/gpt-5.3-codex",
|
||||
provider=LlmProviders.GITHUB_COPILOT,
|
||||
)
|
||||
|
||||
|
|
@ -373,3 +388,200 @@ class TestGithubCopilotResponsesAPITransformation:
|
|||
|
||||
# Non-reasoning items should pass through unchanged
|
||||
assert result == message_item
|
||||
|
||||
|
||||
class TestGithubCopilotResponsesAPIRouting:
|
||||
"""``ProviderConfigManager.get_provider_responses_api_config`` for github_copilot
|
||||
returns the native Responses config only when the model has ``mode=responses``
|
||||
in the (already-merged) model info; otherwise returns None so the dispatcher
|
||||
routes through the chat-completions translation bridge."""
|
||||
|
||||
@patch(
|
||||
"litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper"
|
||||
)
|
||||
def test_returns_config_when_mode_is_responses(self, mock_get_info):
|
||||
"""``mode=responses`` returns native config."""
|
||||
mock_get_info.return_value = {"mode": "responses"}
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model="github_copilot/some-responses-model",
|
||||
provider=LlmProviders.GITHUB_COPILOT,
|
||||
)
|
||||
assert isinstance(config, GithubCopilotResponsesAPIConfig)
|
||||
|
||||
@patch(
|
||||
"litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper"
|
||||
)
|
||||
def test_returns_none_when_mode_is_chat(self, mock_get_info):
|
||||
"""``mode=chat`` returns None so dispatcher uses bridge."""
|
||||
mock_get_info.return_value = {"mode": "chat"}
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model="github_copilot/some-chat-only-model",
|
||||
provider=LlmProviders.GITHUB_COPILOT,
|
||||
)
|
||||
assert config is None
|
||||
|
||||
@patch(
|
||||
"litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper"
|
||||
)
|
||||
def test_returns_none_when_mode_is_unset_and_no_endpoints(self, mock_get_info):
|
||||
"""Entry without ``mode`` and without ``supported_endpoints`` returns None
|
||||
(conservative default)."""
|
||||
mock_get_info.return_value = {}
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model="github_copilot/some-model",
|
||||
provider=LlmProviders.GITHUB_COPILOT,
|
||||
)
|
||||
assert config is None
|
||||
|
||||
def test_returns_config_when_mode_unset_but_endpoints_have_responses(self):
|
||||
"""``mode`` unset but ``supported_endpoints`` declaring /v1/responses
|
||||
returns native config (endpoint-list fallback for stale-but-correct
|
||||
catalog entries that lack ``mode``).
|
||||
|
||||
Exercises the real ``_cached_get_model_info_helper`` plumbing via
|
||||
``register_model`` (no mock). ``supported_endpoints`` is not carried on
|
||||
the normalized ``ModelInfoBase`` the helper returns, so the gate must
|
||||
read it from the raw ``litellm.model_cost`` entry; a mock-based test
|
||||
would mask that.
|
||||
"""
|
||||
litellm.register_model(
|
||||
{
|
||||
"github_copilot/test-endpoints-only-model": {
|
||||
"litellm_provider": "github_copilot",
|
||||
"max_tokens": 1,
|
||||
"input_cost_per_token": 0,
|
||||
"output_cost_per_token": 0,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses",
|
||||
],
|
||||
}
|
||||
}
|
||||
)
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model="github_copilot/test-endpoints-only-model",
|
||||
provider=LlmProviders.GITHUB_COPILOT,
|
||||
)
|
||||
assert isinstance(config, GithubCopilotResponsesAPIConfig)
|
||||
|
||||
def test_mode_chat_overrides_endpoints_with_responses(self):
|
||||
"""``mode=chat`` is a hard opt-out: forces bridge even when
|
||||
``supported_endpoints`` includes /v1/responses. Lets users force the
|
||||
bridge for dual-endpoint models without clearing endpoint metadata.
|
||||
|
||||
Exercises the real ``_cached_get_model_info_helper`` plumbing via
|
||||
``register_model`` (no mock) so the ``mode``-over-endpoints precedence
|
||||
is verified against the actual model-info resolution.
|
||||
"""
|
||||
litellm.register_model(
|
||||
{
|
||||
"github_copilot/test-chat-override-model": {
|
||||
"litellm_provider": "github_copilot",
|
||||
"max_tokens": 1,
|
||||
"input_cost_per_token": 0,
|
||||
"output_cost_per_token": 0,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses",
|
||||
],
|
||||
}
|
||||
}
|
||||
)
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model="github_copilot/test-chat-override-model",
|
||||
provider=LlmProviders.GITHUB_COPILOT,
|
||||
)
|
||||
assert config is None
|
||||
|
||||
def test_returns_config_when_model_is_none(self):
|
||||
"""Follow-up GET/DELETE operations pass model=None and keep the native
|
||||
config path (no per-model lookup is possible)."""
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=LlmProviders.GITHUB_COPILOT,
|
||||
)
|
||||
assert isinstance(config, GithubCopilotResponsesAPIConfig)
|
||||
|
||||
@patch(
|
||||
"litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper"
|
||||
)
|
||||
def test_returns_none_when_get_model_info_raises(self, mock_get_info):
|
||||
"""Catalog lookup failure (model not registered) returns None
|
||||
(conservative default; bridge handles unknown models safely)."""
|
||||
mock_get_info.side_effect = Exception("model not in catalog")
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model="github_copilot/never-seen-model",
|
||||
provider=LlmProviders.GITHUB_COPILOT,
|
||||
)
|
||||
assert config is None
|
||||
|
||||
@patch(
|
||||
"litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper"
|
||||
)
|
||||
def test_user_override_via_register_model(self, mock_get_info):
|
||||
"""User-supplied per-deployment ``model_info`` flows through
|
||||
``litellm.register_model`` (called by the router) into the merged
|
||||
catalog read by ``_cached_get_model_info_helper``. Setting ``mode=responses``
|
||||
for a model whose catalog entry says ``mode=chat`` therefore opts in
|
||||
to native dispatch without any per-call argument plumbing."""
|
||||
mock_get_info.return_value = {"mode": "responses"}
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model="github_copilot/some-chat-only-model",
|
||||
provider=LlmProviders.GITHUB_COPILOT,
|
||||
)
|
||||
assert isinstance(config, GithubCopilotResponsesAPIConfig)
|
||||
|
||||
@patch(
|
||||
"litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper"
|
||||
)
|
||||
def test_realistic_chat_only_entry_returns_none(self, mock_get_info):
|
||||
"""Realistic ``model_prices_and_context_window.json`` shape for a
|
||||
chat-only Copilot model (e.g. github_copilot/gemini-3.1-pro-preview)
|
||||
returns None so /v1/responses calls fall back to the bridge."""
|
||||
mock_get_info.return_value = {
|
||||
"litellm_provider": "github_copilot",
|
||||
"max_input_tokens": 136000,
|
||||
"max_output_tokens": 64000,
|
||||
"max_tokens": 64000,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions"],
|
||||
"supports_function_calling": True,
|
||||
"supports_tool_choice": True,
|
||||
"supports_parallel_function_calling": True,
|
||||
"supports_vision": True,
|
||||
"supports_reasoning": True,
|
||||
}
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model="github_copilot/some-chat-only-model",
|
||||
provider=LlmProviders.GITHUB_COPILOT,
|
||||
)
|
||||
assert config is None
|
||||
|
||||
@patch(
|
||||
"litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper"
|
||||
)
|
||||
def test_realistic_responses_only_entry_returns_config(self, mock_get_info):
|
||||
"""Realistic catalog entry for a Responses-only Copilot model
|
||||
(e.g. github_copilot/gpt-5.5) returns the native config."""
|
||||
mock_get_info.return_value = {
|
||||
"litellm_provider": "github_copilot",
|
||||
"max_input_tokens": 272000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "responses",
|
||||
"supported_endpoints": ["/v1/responses"],
|
||||
"supports_function_calling": True,
|
||||
"supports_tool_choice": True,
|
||||
"supports_parallel_function_calling": True,
|
||||
"supports_response_schema": True,
|
||||
"supports_vision": True,
|
||||
"supports_reasoning": True,
|
||||
"supports_none_reasoning_effort": True,
|
||||
"supports_xhigh_reasoning_effort": True,
|
||||
}
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model="github_copilot/some-responses-only-model",
|
||||
provider=LlmProviders.GITHUB_COPILOT,
|
||||
)
|
||||
assert isinstance(config, GithubCopilotResponsesAPIConfig)
|
||||
|
|
|
|||
172
tests/test_litellm/llms/parasail/test_parasail.py
Normal file
172
tests/test_litellm/llms/parasail/test_parasail.py
Normal file
|
|
@ -0,0 +1,172 @@
|
|||
import os
|
||||
from unittest.mock import patch
|
||||
|
||||
PARASAIL_API_BASE = "https://api.parasail.io/v1"
|
||||
PARASAIL_RESPONSES_GATEWAY = "https://api-webflux.saas.parasail.io/v1"
|
||||
|
||||
|
||||
def test_parasail_json_registry():
|
||||
import litellm
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
assert litellm.LlmProviders.PARASAIL.value == "parasail"
|
||||
assert litellm.LlmProviders("parasail") == litellm.LlmProviders.PARASAIL
|
||||
assert JSONProviderRegistry.exists("parasail")
|
||||
config = JSONProviderRegistry.get("parasail")
|
||||
assert config is not None
|
||||
assert config.base_url == PARASAIL_API_BASE
|
||||
assert config.api_key_env == "PARASAIL_API_KEY"
|
||||
assert config.api_base_env == "PARASAIL_API_BASE"
|
||||
assert "/v1/chat/completions" in config.supported_endpoints
|
||||
assert "/v1/responses" in config.supported_endpoints
|
||||
assert config.special_handling.get("force_store_false") is True
|
||||
|
||||
|
||||
def test_parasail_listed_in_openai_compatible_providers():
|
||||
from litellm.constants import openai_compatible_providers
|
||||
|
||||
assert "parasail" in openai_compatible_providers
|
||||
|
||||
|
||||
def test_parasail_dynamic_config_env_vars():
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
config = create_config_class(JSONProviderRegistry.get("parasail"))()
|
||||
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"PARASAIL_API_KEY": "test-key",
|
||||
"PARASAIL_API_BASE": PARASAIL_RESPONSES_GATEWAY,
|
||||
},
|
||||
):
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
|
||||
|
||||
assert api_base == PARASAIL_RESPONSES_GATEWAY
|
||||
assert api_key == "test-key"
|
||||
|
||||
|
||||
def test_parasail_provider_detection_by_prefix():
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
model, provider, _, api_base = get_llm_provider(
|
||||
"parasail/parasail-llama-33-70b-fp8"
|
||||
)
|
||||
|
||||
assert model == "parasail-llama-33-70b-fp8"
|
||||
assert provider == "parasail"
|
||||
assert api_base == PARASAIL_API_BASE
|
||||
|
||||
|
||||
def test_parasail_chat_complete_url():
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
config = create_config_class(JSONProviderRegistry.get("parasail"))()
|
||||
|
||||
assert (
|
||||
config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="parasail-llama-33-70b-fp8",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
== f"{PARASAIL_API_BASE}/chat/completions"
|
||||
)
|
||||
|
||||
|
||||
def test_parasail_responses_api_config():
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="parasail",
|
||||
model="parasail-kimi-k25-elicit",
|
||||
)
|
||||
|
||||
assert isinstance(config, OpenAIResponsesAPIConfig)
|
||||
assert config.custom_llm_provider == "parasail"
|
||||
assert (
|
||||
config.get_complete_url(api_base=None, litellm_params={})
|
||||
== f"{PARASAIL_API_BASE}/responses"
|
||||
)
|
||||
|
||||
|
||||
def test_parasail_responses_api_honors_api_base_override():
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="parasail",
|
||||
model="parasail-kimi-k25-elicit",
|
||||
)
|
||||
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{"PARASAIL_API_BASE": PARASAIL_RESPONSES_GATEWAY},
|
||||
):
|
||||
url = config.get_complete_url(api_base=None, litellm_params={})
|
||||
|
||||
assert url == f"{PARASAIL_RESPONSES_GATEWAY}/responses"
|
||||
|
||||
|
||||
def test_parasail_responses_api_forces_store_false_when_caller_sets_true():
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="parasail",
|
||||
model="parasail-kimi-k25-elicit",
|
||||
)
|
||||
|
||||
request_params: dict = {"store": True, "temperature": 0.2}
|
||||
transformed = config.transform_responses_api_request(
|
||||
model="parasail-kimi-k25-elicit",
|
||||
input="hello",
|
||||
response_api_optional_request_params=request_params,
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert transformed["store"] is False
|
||||
assert transformed["temperature"] == 0.2
|
||||
|
||||
|
||||
def test_parasail_responses_api_forces_store_false_when_caller_omits_store():
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="parasail",
|
||||
model="parasail-kimi-k25-elicit",
|
||||
)
|
||||
|
||||
transformed = config.transform_responses_api_request(
|
||||
model="parasail-kimi-k25-elicit",
|
||||
input="hello",
|
||||
response_api_optional_request_params={},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert transformed["store"] is False
|
||||
|
||||
|
||||
def test_parasail_responses_api_validate_environment_sets_bearer_token():
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="parasail",
|
||||
model="parasail-kimi-k25-elicit",
|
||||
)
|
||||
|
||||
with patch.dict(os.environ, {"PARASAIL_API_KEY": "secret-from-env"}):
|
||||
headers = config.validate_environment(
|
||||
headers={},
|
||||
model="parasail-kimi-k25-elicit",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
)
|
||||
|
||||
assert headers["Authorization"] == "Bearer secret-from-env"
|
||||
|
|
@ -1459,6 +1459,26 @@ def test_vertex_ai_process_candidates_with_grounding_metadata():
|
|||
assert len(result[0]) == 1
|
||||
|
||||
|
||||
def test_set_stream_metadata_mirrors_non_streaming_safety_field_names():
|
||||
safety_ratings = [
|
||||
[{"category": "HARM_CATEGORY_HATE_SPEECH", "probability": "NEGLIGIBLE"}]
|
||||
]
|
||||
|
||||
model_response = ModelResponse()
|
||||
VertexGeminiConfig._set_stream_metadata_on_response(
|
||||
model_response=model_response,
|
||||
grounding_metadata=[],
|
||||
url_context_metadata=[],
|
||||
safety_ratings=safety_ratings,
|
||||
citation_metadata=[],
|
||||
)
|
||||
|
||||
assert getattr(model_response, "vertex_ai_safety_ratings") == safety_ratings
|
||||
assert getattr(model_response, "vertex_ai_safety_results") == safety_ratings
|
||||
assert model_response._hidden_params["vertex_ai_safety_ratings"] == safety_ratings
|
||||
assert model_response._hidden_params["vertex_ai_safety_results"] == safety_ratings
|
||||
|
||||
|
||||
def test_vertex_ai_tool_call_id_format():
|
||||
"""
|
||||
Test that tool call IDs have the correct format and length.
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Tests for MCP OAuth discoverable endpoints"""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -2661,3 +2662,74 @@ async def test_token_endpoint_sets_no_store_cache_control():
|
|||
|
||||
assert response.headers["cache-control"] == "no-store"
|
||||
assert response.headers["pragma"] == "no-cache"
|
||||
|
||||
|
||||
async def _exchange_with_upstream_token_response(upstream_body):
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
exchange_token_with_server,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
server = MCPServer(
|
||||
server_id="t",
|
||||
name="t",
|
||||
server_name="t",
|
||||
alias="t",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="cid",
|
||||
client_secret="cs",
|
||||
authorization_url="https://provider.com/oauth/authorize",
|
||||
token_url="https://provider.com/oauth/token",
|
||||
)
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
fake_http_response = MagicMock()
|
||||
fake_http_response.json.return_value = upstream_body
|
||||
fake_http_response.raise_for_status = MagicMock()
|
||||
fake_http_client = MagicMock()
|
||||
fake_http_client.post = AsyncMock(return_value=fake_http_response)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
|
||||
return_value=fake_http_client,
|
||||
):
|
||||
response = await exchange_token_with_server(
|
||||
request=mock_request,
|
||||
mcp_server=server,
|
||||
grant_type="authorization_code",
|
||||
code="c",
|
||||
redirect_uri="http://127.0.0.1:3000/cb",
|
||||
client_id="cid",
|
||||
client_secret=None,
|
||||
code_verifier=None,
|
||||
)
|
||||
return json.loads(response.body)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_exchange_omits_expires_in_when_upstream_omits_it():
|
||||
"""A provider that issues a non-expiring token (e.g. Slack without token
|
||||
rotation) returns no ``expires_in``. The exchange must mirror that and omit
|
||||
``expires_in`` rather than fabricate a 1-hour TTL, so the stored credential
|
||||
is treated as non-expiring instead of dying after an hour."""
|
||||
body = await _exchange_with_upstream_token_response(
|
||||
{"access_token": "tok", "token_type": "Bearer"}
|
||||
)
|
||||
assert "expires_in" not in body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_exchange_passes_through_upstream_expires_in():
|
||||
"""When the provider does send ``expires_in`` (e.g. Slack with token
|
||||
rotation), the exchange forwards the real value unchanged."""
|
||||
body = await _exchange_with_upstream_token_response(
|
||||
{"access_token": "tok", "token_type": "Bearer", "expires_in": 43200}
|
||||
)
|
||||
assert body["expires_in"] == 43200
|
||||
|
|
|
|||
|
|
@ -501,6 +501,7 @@ class TestListToolsRestAPI:
|
|||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
extra_headers=None,
|
||||
apply_tool_filters=True,
|
||||
):
|
||||
captured["called"] = True
|
||||
captured["server"] = server
|
||||
|
|
@ -545,6 +546,78 @@ class TestListToolsRestAPI:
|
|||
assert result["error"] is None
|
||||
assert result["message"] == "Successfully retrieved tools"
|
||||
|
||||
async def test_include_disabled_tools_is_admin_only(self, monkeypatch):
|
||||
"""include_disabled_tools skips the allowlist filter only for PROXY_ADMIN;
|
||||
a non-admin passing it stays filtered so the REST endpoint can't be used
|
||||
to enumerate deliberately-disabled tools."""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
async def fake_contexts(user_api_key_auth):
|
||||
return [user_api_key_auth]
|
||||
|
||||
async def fake_get_allowed_mcp_servers(*args, **kwargs):
|
||||
return ["server-1"]
|
||||
|
||||
class StubServer:
|
||||
alias = "server-1"
|
||||
server_name = "server-1"
|
||||
name = "stub"
|
||||
allowed_tools = ["tool1"]
|
||||
mcp_info = {"server_name": "stub"}
|
||||
available_on_public_internet = True
|
||||
|
||||
stub_server = StubServer()
|
||||
captured = {}
|
||||
|
||||
async def fake_get_tools(
|
||||
server, server_auth_header, *args, apply_tool_filters=True, **kwargs
|
||||
):
|
||||
captured["apply_tool_filters"] = apply_tool_filters
|
||||
return ["tool-1"]
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints,
|
||||
"build_effective_auth_contexts",
|
||||
fake_contexts,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_allowed_mcp_servers",
|
||||
fake_get_allowed_mcp_servers,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_mcp_server_by_id",
|
||||
lambda server_id: stub_server if server_id == "server-1" else None,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints,
|
||||
"_get_tools_for_single_server",
|
||||
fake_get_tools,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
request = _build_request(path="/mcp-rest/tools/list", method="GET")
|
||||
|
||||
await rest_endpoints.list_tool_rest_api(
|
||||
request,
|
||||
server_id="server-1",
|
||||
include_disabled_tools=True,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
assert captured["apply_tool_filters"] is False
|
||||
|
||||
await rest_endpoints.list_tool_rest_api(
|
||||
request,
|
||||
server_id="server-1",
|
||||
include_disabled_tools=True,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
assert captured["apply_tool_filters"] is True
|
||||
|
||||
@pytest.mark.parametrize("upstream_status", [401, 403])
|
||||
async def test_upstream_auth_failure_surfaces_status_and_challenge(
|
||||
self, monkeypatch, upstream_status
|
||||
|
|
@ -649,6 +722,7 @@ class TestListToolsRestAPI:
|
|||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
extra_headers=None,
|
||||
apply_tool_filters=True,
|
||||
):
|
||||
captured["called"] = True
|
||||
captured["server_arg"] = server
|
||||
|
|
@ -792,6 +866,7 @@ class TestListToolsRestAPI:
|
|||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
extra_headers=None,
|
||||
apply_tool_filters=True,
|
||||
):
|
||||
captured["server"] = server
|
||||
captured["auth_header"] = server_auth_header
|
||||
|
|
@ -1284,6 +1359,56 @@ class TestGetToolsForSingleServer:
|
|||
assert "tool1" not in tool_names
|
||||
assert "tool4" not in tool_names
|
||||
|
||||
async def test_apply_tool_filters_false_returns_full_catalog(self, monkeypatch):
|
||||
"""apply_tool_filters=False returns the raw catalog without the server
|
||||
allowed_tools gate, so the config UI can render disabled tools as off."""
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
class MockTool:
|
||||
def __init__(self, name):
|
||||
self.name = name
|
||||
self.description = name
|
||||
self.inputSchema = {}
|
||||
|
||||
mock_tools = [MockTool("tool1"), MockTool("tool2"), MockTool("tool3")]
|
||||
|
||||
async def fake_get_tools_from_server(**kwargs):
|
||||
return mock_tools
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_get_tools_from_server",
|
||||
fake_get_tools_from_server,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
# Server enforces an allowlist of just tool1.
|
||||
server = MCPServer(
|
||||
server_id="test-server-id",
|
||||
name="test-server",
|
||||
transport=MCPTransport.sse,
|
||||
allowed_tools=["tool1"],
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test-key", object_permission=None)
|
||||
|
||||
# Runtime default: only the allowed tool comes back.
|
||||
filtered = await rest_endpoints._get_tools_for_single_server(
|
||||
server=server,
|
||||
server_auth_header=None,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
)
|
||||
assert [t.name for t in filtered] == ["tool1"]
|
||||
|
||||
# Config view: full catalog, including the disabled tools.
|
||||
full = await rest_endpoints._get_tools_for_single_server(
|
||||
server=server,
|
||||
server_auth_header=None,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
apply_tool_filters=False,
|
||||
)
|
||||
assert {t.name for t in full} == {"tool1", "tool2", "tool3"}
|
||||
|
||||
|
||||
class TestStdioCommandAllowlist:
|
||||
"""Tests for MCP stdio command allowlist validation."""
|
||||
|
|
|
|||
|
|
@ -3182,6 +3182,310 @@ def test_build_decode_kwargs_no_warning_when_scoped(
|
|||
assert matching == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Defer to single-team DB fallback (PR #26418) when JWT claims are present
|
||||
# but do not resolve to a LiteLLM team.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_and_validate_specific_team_id_unresolved_claim_returns_none():
|
||||
"""With `team_claim_fallback=True`: team_id claim is present in the JWT
|
||||
but the team is missing in the DB — return (None, None) so the
|
||||
auth_builder single-team fallback can run, instead of raising and
|
||||
failing auth."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
team_id_jwt_field="team_id",
|
||||
team_claim_fallback=True,
|
||||
)
|
||||
token = {"sub": "user-1", "team_id": "claim-team-not-in-db"}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_team:
|
||||
mock_get_team.side_effect = HTTPException(status_code=404, detail="missing")
|
||||
|
||||
team_id, team_object = await JWTAuthManager.find_and_validate_specific_team_id(
|
||||
jwt_handler=jwt_handler,
|
||||
jwt_valid_token=token,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=None,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert team_id is None
|
||||
assert team_object is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_team_with_model_access_unresolved_group_claim_returns_none(
|
||||
monkeypatch,
|
||||
):
|
||||
"""With `team_claim_fallback=True`: group claim resolves to team_ids that
|
||||
don't exist in the DB — return (None, None) instead of raising 403, so
|
||||
the single-team fallback can run."""
|
||||
import sys
|
||||
import types
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.router import Router
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}
|
||||
]
|
||||
)
|
||||
proxy_server_module = types.ModuleType("proxy_server")
|
||||
proxy_server_module.llm_router = router
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module)
|
||||
|
||||
async def raise_404(*_args, **_kwargs):
|
||||
raise HTTPException(status_code=404, detail="missing")
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", raise_404)
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_claim_fallback=True)
|
||||
|
||||
team_id, team_object = await JWTAuthManager.find_team_with_model_access(
|
||||
team_ids={"idp-group-a", "idp-group-b"},
|
||||
requested_model="gpt-4o-mini",
|
||||
route="/chat/completions",
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=None,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert team_id is None
|
||||
assert team_object is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_and_validate_specific_team_id_non_http_exception_still_propagates():
|
||||
"""Regression guard: only the 404 HTTPException raised by
|
||||
`get_team_object` ("team doesn't exist in db") is softened. Other
|
||||
errors — e.g. "No DB Connected" — must still propagate so operator-side
|
||||
problems are loud."""
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id")
|
||||
token = {"sub": "user-1", "team_id": "some-claim-team"}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_team:
|
||||
mock_get_team.side_effect = RuntimeError("simulated infrastructure error")
|
||||
|
||||
with pytest.raises(RuntimeError, match="simulated infrastructure error"):
|
||||
await JWTAuthManager.find_and_validate_specific_team_id(
|
||||
jwt_handler=jwt_handler,
|
||||
jwt_valid_token=token,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=None,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_and_validate_specific_team_id_non_404_http_exception_propagates():
|
||||
"""Regression guard: only 404 HTTPException is softened. If
|
||||
`get_team_object` is ever updated to raise a different HTTP status code
|
||||
(e.g. 403 for a blocked team), that error must still propagate rather
|
||||
than silently fall through to the single-team DB fallback."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id")
|
||||
token = {"sub": "user-1", "team_id": "some-claim-team"}
|
||||
|
||||
for status_code in (400, 403, 500):
|
||||
with patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_team:
|
||||
mock_get_team.side_effect = HTTPException(
|
||||
status_code=status_code, detail="non-404 failure"
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await JWTAuthManager.find_and_validate_specific_team_id(
|
||||
jwt_handler=jwt_handler,
|
||||
jwt_valid_token=token,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=None,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
assert exc_info.value.status_code == status_code
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_team_with_model_access_enforce_team_based_access_still_raises():
|
||||
"""Regression guard: when no group claims are present and
|
||||
`enforce_team_based_model_access` is on, the original 403 still fires —
|
||||
the new soft-fail only applies to the unresolved-claim path inside the
|
||||
loop, not to the no-team-claims-at-all path at the top."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(enforce_team_based_model_access=True)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await JWTAuthManager.find_team_with_model_access(
|
||||
team_ids=set(),
|
||||
requested_model="gpt-4o-mini",
|
||||
route="/chat/completions",
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=None,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "enforce_team_based_model_access" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_team_with_model_access_resolved_team_without_model_still_raises_403(
|
||||
monkeypatch,
|
||||
):
|
||||
"""Regression guard: when the JWT group claim DOES resolve to a real
|
||||
LiteLLM team but that team does not grant the requested model, keep the
|
||||
original 403. Only the unresolved-claim case is softened."""
|
||||
import sys
|
||||
import types
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.router import Router
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}},
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {"model": "gpt-3.5-turbo"},
|
||||
},
|
||||
]
|
||||
)
|
||||
proxy_server_module = types.ModuleType("proxy_server")
|
||||
proxy_server_module.llm_router = router
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module)
|
||||
|
||||
team = LiteLLM_TeamTable(team_id="real-team", models=["gpt-3.5-turbo"])
|
||||
|
||||
async def mock_get_team_object(*_args, **_kwargs):
|
||||
return team
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object
|
||||
)
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await JWTAuthManager.find_team_with_model_access(
|
||||
team_ids={"real-team"},
|
||||
requested_model="gpt-4o-mini",
|
||||
route="/chat/completions",
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=None,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "No team has access to the requested model" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_and_validate_specific_team_id_unresolved_claim_default_raises():
|
||||
"""Default `team_claim_fallback=False`: unresolved team_id claim must
|
||||
still raise — preserves the strict claim-based authorization boundary
|
||||
when the operator has not opted in to the fallback."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id")
|
||||
token = {"sub": "user-1", "team_id": "claim-team-not-in-db"}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_team:
|
||||
mock_get_team.side_effect = HTTPException(status_code=404, detail="missing")
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await JWTAuthManager.find_and_validate_specific_team_id(
|
||||
jwt_handler=jwt_handler,
|
||||
jwt_valid_token=token,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=None,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_team_with_model_access_unresolved_group_claim_default_raises(
|
||||
monkeypatch,
|
||||
):
|
||||
"""Default `team_claim_fallback=False`: group claims that don't resolve
|
||||
to any LiteLLM team must still raise 403 — preserves the strict
|
||||
claim-based authorization boundary."""
|
||||
import sys
|
||||
import types
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.router import Router
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}
|
||||
]
|
||||
)
|
||||
proxy_server_module = types.ModuleType("proxy_server")
|
||||
proxy_server_module.llm_router = router
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module)
|
||||
|
||||
async def raise_404(*_args, **_kwargs):
|
||||
raise HTTPException(status_code=404, detail="missing")
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", raise_404)
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await JWTAuthManager.find_team_with_model_access(
|
||||
team_ids={"idp-group-a", "idp-group-b"},
|
||||
requested_model="gpt-4o-mini",
|
||||
route="/chat/completions",
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=None,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
# GH #26789: JWT claim user_id must rebind to legacy DB row after fuzzy match.
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -112,11 +112,71 @@ async def test_should_clear_stale_budget_reservation_when_budget_checks_skip():
|
|||
user_api_key_cache=MagicMock(),
|
||||
proxy_logging_obj=MagicMock(),
|
||||
skip_budget_checks=True,
|
||||
general_settings={},
|
||||
)
|
||||
|
||||
assert user_api_key_auth_obj.budget_reservation is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disable_budget_reservation_skips_reservation():
|
||||
"""#27639: general_settings.disable_budget_reservation turns off the optimistic Redis
|
||||
reservation so operators hit by phantom BudgetExceededError can opt out of it."""
|
||||
user_api_key_auth_obj = UserAPIKeyAuth(token="test_token")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
|
||||
new=AsyncMock(return_value={"reserved_cost": 0.5, "entries": []}),
|
||||
) as mock_reserve:
|
||||
await _reserve_budget_after_common_checks(
|
||||
user_api_key_auth_obj=user_api_key_auth_obj,
|
||||
request_data={"model": "gpt-4o"},
|
||||
route="/v1/chat/completions",
|
||||
llm_router=None,
|
||||
team_object=None,
|
||||
user_object=None,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=MagicMock(),
|
||||
proxy_logging_obj=MagicMock(),
|
||||
skip_budget_checks=False,
|
||||
general_settings={"disable_budget_reservation": True},
|
||||
)
|
||||
|
||||
mock_reserve.assert_not_called()
|
||||
assert user_api_key_auth_obj.budget_reservation is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_budget_reservation_runs_when_not_disabled():
|
||||
"""Control for #27639: with the flag absent, the reservation still runs and is stored."""
|
||||
user_api_key_auth_obj = UserAPIKeyAuth(token="test_token")
|
||||
reservation = {
|
||||
"reserved_cost": 0.5,
|
||||
"entries": [{"counter_key": "spend:key:test_token"}],
|
||||
}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
|
||||
new=AsyncMock(return_value=reservation),
|
||||
) as mock_reserve:
|
||||
await _reserve_budget_after_common_checks(
|
||||
user_api_key_auth_obj=user_api_key_auth_obj,
|
||||
request_data={"model": "gpt-4o"},
|
||||
route="/v1/chat/completions",
|
||||
llm_router=None,
|
||||
team_object=None,
|
||||
user_object=None,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=MagicMock(),
|
||||
proxy_logging_obj=MagicMock(),
|
||||
skip_budget_checks=False,
|
||||
general_settings={},
|
||||
)
|
||||
|
||||
mock_reserve.assert_awaited_once()
|
||||
assert user_api_key_auth_obj.budget_reservation == reservation
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_not_reuse_cached_key_object_for_request_state():
|
||||
key_cache = DualCache()
|
||||
|
|
|
|||
|
|
@ -41,7 +41,8 @@ def test_crowdstrike_aidr_guardrail_config() -> None:
|
|||
)
|
||||
|
||||
|
||||
def test_crowdstrike_aidr_guardrail_config_no_api_key() -> None:
|
||||
def test_crowdstrike_aidr_guardrail_config_no_api_key(monkeypatch) -> None:
|
||||
monkeypatch.delenv("CS_AIDR_TOKEN", raising=False)
|
||||
with pytest.raises(CrowdStrikeAIDRGuardrailMissingSecrets):
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
|
|
@ -59,7 +60,8 @@ def test_crowdstrike_aidr_guardrail_config_no_api_key() -> None:
|
|||
)
|
||||
|
||||
|
||||
def test_crowdstrike_aidr_guardrail_config_no_api_base() -> None:
|
||||
def test_crowdstrike_aidr_guardrail_config_no_api_base(monkeypatch) -> None:
|
||||
monkeypatch.delenv("CS_AIDR_BASE_URL", raising=False)
|
||||
with pytest.raises(CrowdStrikeAIDRGuardrailMissingSecrets):
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
|
|
@ -412,6 +414,171 @@ async def test_apply_guardrail_response_ok(
|
|||
assert result["texts"] == inputs["texts"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_sends_user_id_model_and_extra_info(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Hello"],
|
||||
"structured_messages": [{"role": "user", "content": "Hello"}],
|
||||
"model": "gpt-4o",
|
||||
}
|
||||
request_data = {
|
||||
"messages": inputs["structured_messages"],
|
||||
"model": "gpt-4o",
|
||||
"litellm_metadata": {
|
||||
"user_api_key_user_id": "uid-abc",
|
||||
"user_api_key_user_email": "alice@example.com",
|
||||
},
|
||||
}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
json={"result": {"blocked": False, "transformed": False}},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
) as mock_method:
|
||||
await crowdstrike_aidr_guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
payload = mock_method.call_args.kwargs["json"]
|
||||
assert payload["user_id"] == "uid-abc"
|
||||
assert payload["model"] == "gpt-4o"
|
||||
assert payload["extra_info"] == {"user_name": "alice@example.com"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_empty_extra_info_when_no_email(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Hello"],
|
||||
"structured_messages": [{"role": "user", "content": "Hello"}],
|
||||
"model": "gemini-flash",
|
||||
}
|
||||
request_data = {
|
||||
"messages": inputs["structured_messages"],
|
||||
"model": "gemini-flash",
|
||||
"litellm_metadata": {
|
||||
"user_api_key_user_id": "uid-no-email",
|
||||
"user_api_key_user_email": None,
|
||||
},
|
||||
}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
json={"result": {"blocked": False, "transformed": False}},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
) as mock_method:
|
||||
await crowdstrike_aidr_guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
payload = mock_method.call_args.kwargs["json"]
|
||||
assert payload["user_id"] == "uid-no-email"
|
||||
assert payload["model"] == "gemini-flash"
|
||||
assert payload["extra_info"] == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_no_metadata_skips_user_fields(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Hello"],
|
||||
"structured_messages": [{"role": "user", "content": "Hello"}],
|
||||
}
|
||||
request_data = {"messages": inputs["structured_messages"]}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
json={"result": {"blocked": False, "transformed": False}},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
) as mock_method:
|
||||
await crowdstrike_aidr_guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
payload = mock_method.call_args.kwargs["json"]
|
||||
assert "user_id" not in payload
|
||||
assert "model" not in payload
|
||||
assert "extra_info" not in payload
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"litellm_metadata, metadata",
|
||||
[
|
||||
(None, {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}),
|
||||
({"trace_id": "t1"}, {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}),
|
||||
(["unexpected"], {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}),
|
||||
({"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}, {"trace_id": "t1"}),
|
||||
],
|
||||
ids=["identity_in_metadata_llm_none", "identity_in_metadata_llm_user_dict", "identity_in_metadata_llm_non_mapping", "identity_in_litellm_metadata"],
|
||||
)
|
||||
async def test_apply_guardrail_reads_identity_from_either_metadata_bag(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
litellm_metadata,
|
||||
metadata,
|
||||
) -> None:
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Hello"],
|
||||
"structured_messages": [{"role": "user", "content": "Hello"}],
|
||||
"model": "gpt-4o",
|
||||
}
|
||||
request_data = {
|
||||
"messages": inputs["structured_messages"],
|
||||
"model": "gpt-4o",
|
||||
"litellm_metadata": litellm_metadata,
|
||||
"metadata": metadata,
|
||||
}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
json={"result": {"blocked": False, "transformed": False}},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
) as mock_method:
|
||||
await crowdstrike_aidr_guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
payload = mock_method.call_args.kwargs["json"]
|
||||
assert payload["user_id"] == "uid-abc"
|
||||
assert payload["extra_info"] == {"user_name": "alice@example.com"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_skipped_messages_stay_aligned(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
|
|
|
|||
|
|
@ -696,6 +696,33 @@ async def test_test_model_connection_falls_back_to_deployments_zero_without_id()
|
|||
assert model_params.get("api_key") == "fake-key-A"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"status,error_message",
|
||||
[
|
||||
("healthy", ""),
|
||||
("unhealthy", "Galileo authentication failed"),
|
||||
],
|
||||
)
|
||||
async def test_health_services_endpoint_galileo(status, error_message):
|
||||
with patch("litellm.integrations.galileo.GalileoObserve") as MockGalileoObserve:
|
||||
mock_instance = MagicMock()
|
||||
mock_instance.async_health_check = AsyncMock(
|
||||
return_value={"status": status, "error_message": error_message}
|
||||
)
|
||||
MockGalileoObserve.return_value = mock_instance
|
||||
|
||||
result = await health_services_endpoint(service="galileo")
|
||||
|
||||
if status == "healthy":
|
||||
assert result["status"] == "healthy"
|
||||
assert result["message"] == "Galileo is healthy"
|
||||
else:
|
||||
assert result["status"] == "unhealthy"
|
||||
assert result["message"] == error_message
|
||||
mock_instance.async_health_check.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_services_endpoint_datadog_llm_observability():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from litellm.proxy.management_endpoints.model_management_endpoints import (
|
|||
ModelManagementAuthChecks,
|
||||
_get_team_deployments,
|
||||
clear_cache,
|
||||
delete_team_models,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, updateDeployment
|
||||
|
|
@ -1915,6 +1916,190 @@ class TestDeleteTeamBYOKModelGhost:
|
|||
mock_refresh.assert_not_awaited()
|
||||
|
||||
|
||||
class TestDeleteModelTeamAuth:
|
||||
"""Team auth on the /model/delete path.
|
||||
|
||||
A model added via /model/new with model_info.team_id is orphaned once its
|
||||
team is deleted: can_user_make_model_call looked the team up and raised
|
||||
'Team id=... does not exist in db' before the delete could run, so the model
|
||||
was undeletable from the Models + Endpoints page. Without the team, team-admin
|
||||
membership can't be verified, so a proxy admin (and only a proxy admin) may
|
||||
delete the orphan; a missing team must never let a non-admin through. The team
|
||||
is also looked up exactly once -- the auth check must not add a second query.
|
||||
"""
|
||||
|
||||
def _orphaned_model_mocks(self, team_id, model_id):
|
||||
db_row = LiteLLM_ProxyModelTable(
|
||||
model_id=model_id,
|
||||
model_name=f"model_name_{team_id}_abc-uuid",
|
||||
litellm_params={"model": "openai/gpt-4.1-nano"},
|
||||
model_info={
|
||||
"id": model_id,
|
||||
"team_id": team_id,
|
||||
"team_public_model_name": "orphaned-gpt",
|
||||
},
|
||||
created_by="admin",
|
||||
updated_by="admin",
|
||||
)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = MagicMock()
|
||||
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
|
||||
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
|
||||
return_value=db_row
|
||||
)
|
||||
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
|
||||
mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||||
# The team is gone -> every team lookup returns None.
|
||||
mock_prisma.db.litellm_teamtable = AsyncMock()
|
||||
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_teamtable.update = AsyncMock()
|
||||
mock_prisma.db.litellm_modeltable = AsyncMock()
|
||||
mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[])
|
||||
return mock_prisma
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_admin_can_delete_model_when_team_deleted(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
ModelInfoDelete,
|
||||
delete_model as delete_model_endpoint,
|
||||
)
|
||||
|
||||
team_id = "deleted-team-xyz"
|
||||
model_id = "orphaned-byok-1"
|
||||
mock_prisma = self._orphaned_model_mocks(team_id, model_id)
|
||||
|
||||
admin_user = UserAPIKeyAuth(
|
||||
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
_PS = "litellm.proxy.proxy_server"
|
||||
_MOD = "litellm.proxy.management_endpoints.model_management_endpoints"
|
||||
with (
|
||||
patch(f"{_PS}.prisma_client", mock_prisma),
|
||||
patch(f"{_PS}.store_model_in_db", True),
|
||||
patch(f"{_PS}.premium_user", True),
|
||||
patch(f"{_PS}.llm_router", MagicMock()),
|
||||
patch(f"{_PS}.proxy_logging_obj", MagicMock()),
|
||||
patch(f"{_PS}.user_api_key_cache", MagicMock()),
|
||||
patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()),
|
||||
):
|
||||
result = await delete_model_endpoint(
|
||||
model_info=ModelInfoDelete(id=model_id),
|
||||
user_api_key_dict=admin_user,
|
||||
)
|
||||
|
||||
assert "deleted successfully" in result["message"]
|
||||
mock_prisma.db.litellm_proxymodeltable.delete.assert_awaited_once()
|
||||
# Team is gone -> no team.models cleanup to do.
|
||||
mock_prisma.db.litellm_teamtable.update.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_admin_cannot_delete_model_when_team_deleted(self):
|
||||
"""A missing team must never let a non-admin delete the orphan (no fail-open)."""
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
ModelInfoDelete,
|
||||
delete_model as delete_model_endpoint,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ProxyException
|
||||
|
||||
team_id = "deleted-team-abc"
|
||||
model_id = "orphaned-byok-2"
|
||||
mock_prisma = self._orphaned_model_mocks(team_id, model_id)
|
||||
|
||||
non_admin = UserAPIKeyAuth(
|
||||
user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
_PS = "litellm.proxy.proxy_server"
|
||||
_MOD = "litellm.proxy.management_endpoints.model_management_endpoints"
|
||||
with (
|
||||
patch(f"{_PS}.prisma_client", mock_prisma),
|
||||
patch(f"{_PS}.store_model_in_db", True),
|
||||
patch(f"{_PS}.premium_user", True),
|
||||
patch(f"{_PS}.llm_router", MagicMock()),
|
||||
patch(f"{_PS}.proxy_logging_obj", MagicMock()),
|
||||
patch(f"{_PS}.user_api_key_cache", MagicMock()),
|
||||
patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await delete_model_endpoint(
|
||||
model_info=ModelInfoDelete(id=model_id),
|
||||
user_api_key_dict=non_admin,
|
||||
)
|
||||
|
||||
assert str(exc_info.value.code) == "403"
|
||||
mock_prisma.db.litellm_proxymodeltable.delete.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_live_team_delete_looks_up_team_once(self):
|
||||
"""The auth check must not add a redundant team query on the live-team path."""
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
ModelInfoDelete,
|
||||
delete_model as delete_model_endpoint,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ProxyException
|
||||
|
||||
team_id = "live-team-1"
|
||||
model_id = "live-byok-1"
|
||||
db_row = LiteLLM_ProxyModelTable(
|
||||
model_id=model_id,
|
||||
model_name=f"model_name_{team_id}_abc-uuid",
|
||||
litellm_params={"model": "openai/gpt-4.1-nano"},
|
||||
model_info={
|
||||
"id": model_id,
|
||||
"team_id": team_id,
|
||||
"team_public_model_name": "live-gpt",
|
||||
},
|
||||
created_by="admin",
|
||||
updated_by="admin",
|
||||
)
|
||||
team_row = LiteLLM_TeamTable(
|
||||
team_id=team_id,
|
||||
team_alias="live-team",
|
||||
members_with_roles=[Member(user_id="admin", role="admin")],
|
||||
models=["live-gpt"],
|
||||
)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = MagicMock()
|
||||
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
|
||||
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(
|
||||
return_value=db_row
|
||||
)
|
||||
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
|
||||
mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.db.litellm_teamtable = AsyncMock()
|
||||
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
|
||||
mock_prisma.db.litellm_modeltable = AsyncMock()
|
||||
mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[])
|
||||
|
||||
# A team member who is not the team admin: rejected before the delete runs,
|
||||
# so the only team lookup is the single one inside the auth check.
|
||||
non_admin = UserAPIKeyAuth(
|
||||
user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
_PS = "litellm.proxy.proxy_server"
|
||||
_MOD = "litellm.proxy.management_endpoints.model_management_endpoints"
|
||||
with (
|
||||
patch(f"{_PS}.prisma_client", mock_prisma),
|
||||
patch(f"{_PS}.store_model_in_db", True),
|
||||
patch(f"{_PS}.premium_user", True),
|
||||
patch(f"{_PS}.llm_router", MagicMock()),
|
||||
patch(f"{_PS}.proxy_logging_obj", MagicMock()),
|
||||
patch(f"{_PS}.user_api_key_cache", MagicMock()),
|
||||
patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await delete_model_endpoint(
|
||||
model_info=ModelInfoDelete(id=model_id),
|
||||
user_api_key_dict=non_admin,
|
||||
)
|
||||
|
||||
assert str(exc_info.value.code) == "403"
|
||||
assert mock_prisma.db.litellm_teamtable.find_unique.await_count == 1
|
||||
mock_prisma.db.litellm_proxymodeltable.delete.assert_not_awaited()
|
||||
|
||||
|
||||
class TestGetTeamDeployments:
|
||||
"""Tests for _get_team_deployments which filters by model_name prefix + Python-side team_id check."""
|
||||
|
||||
|
|
@ -2000,6 +2185,148 @@ class TestGetTeamDeployments:
|
|||
assert result[0] is dep1
|
||||
|
||||
|
||||
def _model_row(model_id: str, team_id: str):
|
||||
row = MagicMock()
|
||||
row.model_id = model_id
|
||||
row.model_name = f"model_name_{team_id}_{model_id}"
|
||||
row.model_info = {"team_id": team_id}
|
||||
return row
|
||||
|
||||
|
||||
class _TxProxyModelTable:
|
||||
"""Transactional proxy-model table that records the order of DB writes."""
|
||||
|
||||
def __init__(self, rows, events):
|
||||
self._rows = list(rows)
|
||||
self.events = events
|
||||
|
||||
async def find_many(self, where):
|
||||
prefix = where["model_name"]["startswith"]
|
||||
return [r for r in self._rows if r.model_name.startswith(prefix)]
|
||||
|
||||
async def delete_many(self, where):
|
||||
ids = list(where["model_id"]["in"])
|
||||
self.events.append(("delete_many", tuple(ids)))
|
||||
self._rows = [r for r in self._rows if r.model_id not in ids]
|
||||
return len(ids)
|
||||
|
||||
|
||||
class _TxPrismaClient:
|
||||
"""Minimal prisma stub whose ``db.tx()`` yields a transaction and records commit."""
|
||||
|
||||
def __init__(self, rows):
|
||||
self.events: list = []
|
||||
self._table = _TxProxyModelTable(rows, self.events)
|
||||
tx = MagicMock()
|
||||
tx.litellm_proxymodeltable = self._table
|
||||
outer = self
|
||||
|
||||
class _TxCM:
|
||||
async def __aenter__(self):
|
||||
return tx
|
||||
|
||||
async def __aexit__(self, *exc):
|
||||
outer.events.append(("commit",))
|
||||
return False
|
||||
|
||||
self.db = MagicMock()
|
||||
self.db.tx = MagicMock(return_value=_TxCM())
|
||||
|
||||
|
||||
class _RecordingRouter:
|
||||
def __init__(self, events):
|
||||
self.events = events
|
||||
self.deleted: list = []
|
||||
|
||||
def delete_deployment(self, id): # noqa: A002 - matches router signature
|
||||
self.events.append(("router", id))
|
||||
self.deleted.append(id)
|
||||
|
||||
|
||||
class TestDeleteTeamModels:
|
||||
"""delete_team_models must remove every team's BYOK models in one transaction
|
||||
and sync the in-memory router only after that transaction commits."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deletes_all_teams_models_and_syncs_router(self):
|
||||
rows = [_model_row("a1", "team_a"), _model_row("b1", "team_b")]
|
||||
prisma = _TxPrismaClient(rows)
|
||||
router = _RecordingRouter(prisma.events)
|
||||
|
||||
deleted = await delete_team_models(
|
||||
team_ids=["team_a", "team_b"],
|
||||
prisma_client=prisma,
|
||||
llm_router=router,
|
||||
)
|
||||
|
||||
assert sorted(deleted) == ["a1", "b1"]
|
||||
assert sorted(router.deleted) == ["a1", "b1"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_sync_happens_after_commit(self):
|
||||
"""Race-safety: the router is touched only once the DB transaction has
|
||||
committed, so a rollback can never leave a deployment without its row."""
|
||||
rows = [_model_row("a1", "team_a"), _model_row("b1", "team_b")]
|
||||
prisma = _TxPrismaClient(rows)
|
||||
router = _RecordingRouter(prisma.events)
|
||||
|
||||
await delete_team_models(
|
||||
team_ids=["team_a", "team_b"], prisma_client=prisma, llm_router=router
|
||||
)
|
||||
|
||||
commit_idx = prisma.events.index(("commit",))
|
||||
router_indices = [i for i, e in enumerate(prisma.events) if e[0] == "router"]
|
||||
delete_indices = [
|
||||
i for i, e in enumerate(prisma.events) if e[0] == "delete_many"
|
||||
]
|
||||
assert router_indices, "router was never synced"
|
||||
assert all(i > commit_idx for i in router_indices)
|
||||
assert all(i < commit_idx for i in delete_indices)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_only_owning_team_models_deleted(self):
|
||||
"""A row sharing the prefix but a different model_info.team_id is left alone."""
|
||||
mine = _model_row("a1", "team_a")
|
||||
intruder = MagicMock()
|
||||
intruder.model_id = "x9"
|
||||
intruder.model_name = "model_name_team_a_x9"
|
||||
intruder.model_info = {"team_id": "someone_else"}
|
||||
prisma = _TxPrismaClient([mine, intruder])
|
||||
router = _RecordingRouter(prisma.events)
|
||||
|
||||
deleted = await delete_team_models(
|
||||
team_ids=["team_a"], prisma_client=prisma, llm_router=router
|
||||
)
|
||||
|
||||
assert deleted == ["a1"]
|
||||
assert router.deleted == ["a1"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_models_no_writes(self):
|
||||
prisma = _TxPrismaClient([])
|
||||
router = _RecordingRouter(prisma.events)
|
||||
|
||||
deleted = await delete_team_models(
|
||||
team_ids=["team_a"], prisma_client=prisma, llm_router=router
|
||||
)
|
||||
|
||||
assert deleted == []
|
||||
assert router.deleted == []
|
||||
assert not any(e[0] == "delete_many" for e in prisma.events)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_router_is_safe(self):
|
||||
rows = [_model_row("a1", "team_a")]
|
||||
prisma = _TxPrismaClient(rows)
|
||||
|
||||
deleted = await delete_team_models(
|
||||
team_ids=["team_a"], prisma_client=prisma, llm_router=None
|
||||
)
|
||||
|
||||
assert deleted == ["a1"]
|
||||
assert any(e[0] == "delete_many" for e in prisma.events)
|
||||
|
||||
|
||||
def _build_db_model_for_blocked_test():
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
|
||||
|
|
|
|||
|
|
@ -4531,6 +4531,199 @@ async def test_update_team_standalone_budget_exceeds_user_limit():
|
|||
assert "budget" in str(exc_info.value.message).lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_standalone_unchanged_budget_allowed():
|
||||
"""
|
||||
Test that /team/update for a standalone team does NOT compare against the
|
||||
caller's personal max_budget when the budget is unchanged.
|
||||
|
||||
This is the LiteLLM UI scenario: the UI sends the full team object on every
|
||||
update (including the unchanged max_budget). A team admin only changing
|
||||
tpm_limit should not be blocked by a budget the team already has.
|
||||
|
||||
Scenario:
|
||||
- User (team admin) has personal max_budget=$100
|
||||
- Standalone team exists with current budget=$500
|
||||
- User updates tpm_limit and re-sends the unchanged max_budget=$500
|
||||
- Expected: Should succeed (budget unchanged, not an increase)
|
||||
"""
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_UserTable,
|
||||
UpdateTeamRequest,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.team_endpoints import update_team
|
||||
|
||||
team_admin_user = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="standalone-unchanged-budget-admin",
|
||||
models=[],
|
||||
)
|
||||
|
||||
# UI re-sends the unchanged max_budget alongside the tpm_limit change.
|
||||
update_request = UpdateTeamRequest(
|
||||
team_id="standalone-unchanged-budget-123",
|
||||
max_budget=500.0, # Unchanged from the team's current budget
|
||||
tpm_limit=50000,
|
||||
)
|
||||
|
||||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()
|
||||
) as mock_audit,
|
||||
):
|
||||
# Mock existing standalone team (no organization_id) with budget=$500
|
||||
mock_existing_team = MagicMock()
|
||||
mock_existing_team.team_id = "standalone-unchanged-budget-123"
|
||||
mock_existing_team.organization_id = None
|
||||
mock_existing_team.max_budget = 500.0
|
||||
mock_existing_team.model_id = None
|
||||
mock_existing_team.model_dump.return_value = {
|
||||
"team_id": "standalone-unchanged-budget-123",
|
||||
"organization_id": None,
|
||||
"max_budget": 500.0,
|
||||
"members_with_roles": [
|
||||
{"user_id": "standalone-unchanged-budget-admin", "role": "admin"}
|
||||
],
|
||||
}
|
||||
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(
|
||||
return_value=mock_existing_team
|
||||
)
|
||||
mock_prisma.jsonify_team_object = lambda db_data: db_data
|
||||
|
||||
# User has a restrictive personal budget that is lower than the team's.
|
||||
mock_user_obj = LiteLLM_UserTable(
|
||||
user_id="standalone-unchanged-budget-admin",
|
||||
max_budget=100.0,
|
||||
)
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=mock_user_obj)
|
||||
mock_cache.async_set_cache = AsyncMock()
|
||||
|
||||
mock_updated_team = MagicMock()
|
||||
mock_updated_team.team_id = "standalone-unchanged-budget-123"
|
||||
mock_updated_team.organization_id = None
|
||||
mock_updated_team.max_budget = 500.0
|
||||
mock_updated_team.litellm_model_table = None
|
||||
mock_updated_team.model_dump.return_value = {
|
||||
"team_id": "standalone-unchanged-budget-123",
|
||||
"organization_id": None,
|
||||
"max_budget": 500.0,
|
||||
"tpm_limit": 50000,
|
||||
}
|
||||
mock_prisma.db.litellm_teamtable.update = AsyncMock(
|
||||
return_value=mock_updated_team
|
||||
)
|
||||
|
||||
# Should NOT raise - unchanged budget skips the personal-budget check.
|
||||
result = await update_team(
|
||||
data=update_request,
|
||||
http_request=dummy_request,
|
||||
user_api_key_dict=team_admin_user,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["data"].max_budget == 500.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_standalone_lower_budget_allowed():
|
||||
"""
|
||||
Test that /team/update for a standalone team allows lowering the budget
|
||||
below the team's current value even when the new value still exceeds the
|
||||
caller's personal max_budget.
|
||||
|
||||
Scenario:
|
||||
- User (team admin) has personal max_budget=$100
|
||||
- Standalone team exists with current budget=$500
|
||||
- User lowers team budget to $300 (a decrease, still above user's $100)
|
||||
- Expected: Should succeed (decrease is not an increase above team budget)
|
||||
"""
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_UserTable,
|
||||
UpdateTeamRequest,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.team_endpoints import update_team
|
||||
|
||||
team_admin_user = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="standalone-lower-budget-admin",
|
||||
models=[],
|
||||
)
|
||||
|
||||
update_request = UpdateTeamRequest(
|
||||
team_id="standalone-lower-budget-123",
|
||||
max_budget=300.0, # Lower than current $500, still above user's $100
|
||||
)
|
||||
|
||||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()
|
||||
) as mock_audit,
|
||||
):
|
||||
mock_existing_team = MagicMock()
|
||||
mock_existing_team.team_id = "standalone-lower-budget-123"
|
||||
mock_existing_team.organization_id = None
|
||||
mock_existing_team.max_budget = 500.0
|
||||
mock_existing_team.model_id = None
|
||||
mock_existing_team.model_dump.return_value = {
|
||||
"team_id": "standalone-lower-budget-123",
|
||||
"organization_id": None,
|
||||
"max_budget": 500.0,
|
||||
"members_with_roles": [
|
||||
{"user_id": "standalone-lower-budget-admin", "role": "admin"}
|
||||
],
|
||||
}
|
||||
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(
|
||||
return_value=mock_existing_team
|
||||
)
|
||||
mock_prisma.jsonify_team_object = lambda db_data: db_data
|
||||
|
||||
mock_user_obj = LiteLLM_UserTable(
|
||||
user_id="standalone-lower-budget-admin",
|
||||
max_budget=100.0,
|
||||
)
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=mock_user_obj)
|
||||
mock_cache.async_set_cache = AsyncMock()
|
||||
|
||||
mock_updated_team = MagicMock()
|
||||
mock_updated_team.team_id = "standalone-lower-budget-123"
|
||||
mock_updated_team.organization_id = None
|
||||
mock_updated_team.max_budget = 300.0
|
||||
mock_updated_team.litellm_model_table = None
|
||||
mock_updated_team.model_dump.return_value = {
|
||||
"team_id": "standalone-lower-budget-123",
|
||||
"organization_id": None,
|
||||
"max_budget": 300.0,
|
||||
}
|
||||
mock_prisma.db.litellm_teamtable.update = AsyncMock(
|
||||
return_value=mock_updated_team
|
||||
)
|
||||
|
||||
result = await update_team(
|
||||
data=update_request,
|
||||
http_request=dummy_request,
|
||||
user_api_key_dict=team_admin_user,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["data"].max_budget == 300.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_org_scoped_budget_exceeds_org_limit():
|
||||
"""
|
||||
|
|
@ -6157,6 +6350,14 @@ async def test_delete_team_persists_deleted_teams(monkeypatch):
|
|||
mock_find_many_keys = AsyncMock(return_value=[])
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many_keys
|
||||
|
||||
# delete_team now deletes team BYOK models inside a transaction; this team has none.
|
||||
mock_tx = AsyncMock()
|
||||
mock_tx.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||||
mock_tx_cm = MagicMock()
|
||||
mock_tx_cm.__aenter__ = AsyncMock(return_value=mock_tx)
|
||||
mock_tx_cm.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_prisma_client.db.tx = MagicMock(return_value=mock_tx_cm)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.prisma_client",
|
||||
mock_prisma_client,
|
||||
|
|
@ -8306,6 +8507,8 @@ async def test_new_team_encrypts_callback_vars(
|
|||
assert cv["langfuse_secret_key"] != "sk-real"
|
||||
recovered = decrypt_callback_vars(metadata)["logging"][0]["callback_vars"]
|
||||
assert recovered["langfuse_secret_key"] == "sk-real"
|
||||
|
||||
|
||||
def _non_admin_auth():
|
||||
return UserAPIKeyAuth(
|
||||
user_id="u-team-admin", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
|
|
|
|||
|
|
@ -321,6 +321,44 @@ class TestAzureAnthropicCostCalculation:
|
|||
assert call_kwargs["model"] == "azure_ai/claude-sonnet-4-5_gb_20250929"
|
||||
assert call_kwargs["custom_llm_provider"] == "azure_ai"
|
||||
|
||||
def test_passthrough_logging_sets_response_cost_with_server_tool_use_dict(self):
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
logging_obj = self._create_mock_logging_obj(model="claude-3-7-sonnet-20250219")
|
||||
logging_obj.get_router_model_id.return_value = None
|
||||
logging_obj.litellm_params = {}
|
||||
|
||||
response = ModelResponse(
|
||||
id="test-id",
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(content="test", role="assistant"),
|
||||
)
|
||||
],
|
||||
created=1234567890,
|
||||
model="claude-3-7-sonnet-20250219",
|
||||
usage={
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
"server_tool_use": {"web_search_requests": 1},
|
||||
},
|
||||
)
|
||||
|
||||
kwargs = AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload(
|
||||
litellm_model_response=response,
|
||||
model="claude-3-7-sonnet-20250219",
|
||||
kwargs={},
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
assert "response_cost" in kwargs
|
||||
assert kwargs["response_cost"] > 0
|
||||
|
||||
|
||||
class TestAnthropicBatchPassthroughCostTracking:
|
||||
"""Test cases for Anthropic batch passthrough cost tracking functionality"""
|
||||
|
|
|
|||
|
|
@ -257,6 +257,64 @@ class TestOpenAIPassthroughLoggingHandler:
|
|||
)
|
||||
assert OpenAIPassthroughLoggingHandler.is_openai_responses_route("") == False
|
||||
|
||||
def test_is_openai_route_recognizes_cognitiveservices_azure_com(self):
|
||||
"""Azure OpenAI resources created via the newer "Azure AI Foundry" /
|
||||
Cognitive Services pathway live on `*.cognitiveservices.azure.com`
|
||||
subdomains rather than the older `openai.azure.com`. All four
|
||||
is_openai_*_route methods must recognize both Azure subdomains so
|
||||
cost tracking applies regardless of which Azure naming the user's
|
||||
resource happens to be on.
|
||||
"""
|
||||
cognitive_chat = (
|
||||
"https://my-resource.cognitiveservices.azure.com/v1/chat/completions"
|
||||
)
|
||||
cognitive_images_gen = (
|
||||
"https://my-resource.cognitiveservices.azure.com/v1/images/generations"
|
||||
)
|
||||
cognitive_images_edit = (
|
||||
"https://my-resource.cognitiveservices.azure.com/v1/images/edits"
|
||||
)
|
||||
cognitive_responses = (
|
||||
"https://my-resource.cognitiveservices.azure.com/v1/responses"
|
||||
)
|
||||
|
||||
assert (
|
||||
OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route(
|
||||
cognitive_chat
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
OpenAIPassthroughLoggingHandler.is_openai_image_generation_route(
|
||||
cognitive_images_gen
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
OpenAIPassthroughLoggingHandler.is_openai_image_editing_route(
|
||||
cognitive_images_edit
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
OpenAIPassthroughLoggingHandler.is_openai_responses_route(
|
||||
cognitive_responses
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
# Cross-route negatives still hold for cognitiveservices hosts.
|
||||
assert (
|
||||
OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route(
|
||||
cognitive_responses
|
||||
)
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
OpenAIPassthroughLoggingHandler.is_openai_responses_route(cognitive_chat)
|
||||
is False
|
||||
)
|
||||
|
||||
@patch("litellm.completion_cost")
|
||||
@patch(
|
||||
"litellm.litellm_core_utils.litellm_logging.get_standard_logging_object_payload"
|
||||
|
|
@ -766,6 +824,14 @@ class TestOpenAIPassthroughIntegration:
|
|||
== True
|
||||
)
|
||||
assert self.handler.is_openai_route("https://api.openai.com/v1/models") == True
|
||||
# Azure OpenAI on the shared Cognitive Services domain, identified by an
|
||||
# OpenAI-style path segment.
|
||||
assert (
|
||||
self.handler.is_openai_route(
|
||||
"https://my-resource.cognitiveservices.azure.com/v1/chat/completions"
|
||||
)
|
||||
== True
|
||||
)
|
||||
|
||||
# Negative cases
|
||||
assert (
|
||||
|
|
@ -782,6 +848,28 @@ class TestOpenAIPassthroughIntegration:
|
|||
self.handler.is_openai_route("https://api.assemblyai.com/v2/transcript")
|
||||
== False
|
||||
)
|
||||
# Non-OpenAI Azure Cognitive Services share the `cognitiveservices.azure.com`
|
||||
# domain but must NOT be classified as OpenAI routes (no OpenAI path segment).
|
||||
assert (
|
||||
self.handler.is_openai_route(
|
||||
"https://my-resource.cognitiveservices.azure.com/speechtotext/v3.1/recognize"
|
||||
)
|
||||
== False
|
||||
)
|
||||
assert (
|
||||
self.handler.is_openai_route(
|
||||
"https://my-resource.cognitiveservices.azure.com/vision/v3.2/analyze"
|
||||
)
|
||||
== False
|
||||
)
|
||||
# A look-alike domain that merely contains an OpenAI host as a substring
|
||||
# must be rejected by the suffix-based hostname match.
|
||||
assert (
|
||||
self.handler.is_openai_route(
|
||||
"https://cognitiveservices.azure.com.attacker.example/v1/chat/completions"
|
||||
)
|
||||
== False
|
||||
)
|
||||
assert self.handler.is_openai_route("") == False
|
||||
|
||||
@patch(
|
||||
|
|
|
|||
|
|
@ -60,6 +60,15 @@ def test_v2_model_info_invalid_page_returns_422(client, auth_as, empty_router):
|
|||
assert "detail" in response.json()
|
||||
|
||||
|
||||
def test_v2_model_info_in_openapi_schema():
|
||||
"""``GET /v2/model/info`` is published in the proxy OpenAPI/Swagger spec."""
|
||||
from litellm.proxy.proxy_server import get_openapi_schema
|
||||
|
||||
schema = get_openapi_schema()
|
||||
assert "/v2/model/info" in schema["paths"]
|
||||
assert "get" in schema["paths"]["/v2/model/info"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /v1/model/info, GET /model/info
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -949,6 +949,28 @@ class TestToolChoiceTransformation:
|
|||
result = LiteLLMCompletionResponsesConfig._transform_tool_choice(tool_choice)
|
||||
assert result == tool_choice
|
||||
|
||||
def test_transform_tool_choice_responses_flat_function_name(self):
|
||||
"""Responses-API forced-function with a top-level name maps to the nested Chat
|
||||
Completions shape instead of degrading to required and dropping the name"""
|
||||
result = LiteLLMCompletionResponsesConfig._transform_tool_choice(
|
||||
{"type": "function", "name": "get_weather"}
|
||||
)
|
||||
assert result == {"type": "function", "function": {"name": "get_weather"}}
|
||||
|
||||
def test_transform_tool_choice_function_without_name_falls_back_to_required(self):
|
||||
"""A function-type dict with no name still falls back to required"""
|
||||
result = LiteLLMCompletionResponsesConfig._transform_tool_choice(
|
||||
{"type": "function"}
|
||||
)
|
||||
assert result == "required"
|
||||
|
||||
def test_transform_tool_choice_function_empty_name_falls_back_to_required(self):
|
||||
"""An empty top-level name is falsy and must not produce an empty function name"""
|
||||
result = LiteLLMCompletionResponsesConfig._transform_tool_choice(
|
||||
{"type": "function", "name": ""}
|
||||
)
|
||||
assert result == "required"
|
||||
|
||||
|
||||
class TestContentTypeTransformation:
|
||||
"""Test content type transformation from Responses API to Chat Completion format"""
|
||||
|
|
|
|||
|
|
@ -1,13 +1,9 @@
|
|||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
from typing import Optional
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
import json
|
||||
|
||||
from litellm.types.utils import HiddenParams
|
||||
|
||||
|
|
@ -75,6 +71,48 @@ def test_usage_dump():
|
|||
assert new_usage.prompt_tokens_details.web_search_requests == 1
|
||||
|
||||
|
||||
def test_usage_server_tool_use_dict_is_coerced_and_round_trips():
|
||||
from litellm.types.utils import ServerToolUse, Usage
|
||||
|
||||
current_usage = Usage(
|
||||
completion_tokens=1,
|
||||
prompt_tokens=1,
|
||||
total_tokens=2,
|
||||
server_tool_use={"web_search_requests": 1},
|
||||
)
|
||||
|
||||
assert isinstance(current_usage.server_tool_use, ServerToolUse)
|
||||
assert current_usage.server_tool_use.web_search_requests == 1
|
||||
|
||||
new_usage = Usage(**current_usage.model_dump())
|
||||
assert isinstance(new_usage.server_tool_use, ServerToolUse)
|
||||
assert new_usage.server_tool_use.web_search_requests == 1
|
||||
|
||||
|
||||
def test_usage_converts_server_tool_use_dict():
|
||||
from litellm.types.utils import ServerToolUse, Usage
|
||||
|
||||
usage = Usage(
|
||||
completion_tokens=2,
|
||||
prompt_tokens=1,
|
||||
total_tokens=3,
|
||||
server_tool_use={"web_search_requests": 4, "tool_search_requests": 1},
|
||||
)
|
||||
|
||||
assert isinstance(usage.server_tool_use, ServerToolUse)
|
||||
assert usage.server_tool_use.web_search_requests == 4
|
||||
assert usage.server_tool_use["web_search_requests"] == 4
|
||||
assert usage.server_tool_use.tool_search_requests == 1
|
||||
with pytest.raises(KeyError):
|
||||
usage.server_tool_use["unknown_metric"]
|
||||
|
||||
round_trip = Usage(**usage.model_dump())
|
||||
assert isinstance(round_trip.server_tool_use, ServerToolUse)
|
||||
assert round_trip.server_tool_use.web_search_requests == 4
|
||||
assert round_trip.server_tool_use["web_search_requests"] == 4
|
||||
assert round_trip.server_tool_use.tool_search_requests == 1
|
||||
|
||||
|
||||
def test_usage_completion_tokens_details_text_tokens():
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
|
|
|
|||
|
|
@ -1351,11 +1351,6 @@
|
|||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/mcp_tools/mcp_server_edit.test.tsx": {
|
||||
"unused-imports/no-unused-imports": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/mcp_tools/mcp_server_edit.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
|
|
@ -1517,11 +1512,6 @@
|
|||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/organisms/RegenerateKeyModal.tsx": {
|
||||
"react-hooks/set-state-in-effect": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/organisms/create_key_button.test.tsx": {
|
||||
"@typescript-eslint/no-require-imports": {
|
||||
"count": 2
|
||||
|
|
|
|||
6
ui/litellm-dashboard/package-lock.json
generated
6
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -13770,9 +13770,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/ws": {
|
||||
"version": "8.19.0",
|
||||
"resolved": "https://registry.npmjs.org/ws/-/ws-8.19.0.tgz",
|
||||
"integrity": "sha512-blAT2mjOEIi0ZzruJfIhb3nps74PRWTCz1IjglWEEpQl5XS/UNama6u2/rjFkDDouqr4L67ry+1aGIALViWjDg==",
|
||||
"version": "8.20.1",
|
||||
"resolved": "https://registry.npmjs.org/ws/-/ws-8.20.1.tgz",
|
||||
"integrity": "sha512-It4dO0K5v//JtTXuPkfEOaI3uUN87iYPnqo/ZzqCoG3g8uhA66QUMs/SrM0YK7/NAu+r4LMh/9dq2A7k+rHs+w==",
|
||||
"devOptional": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
|
|
|
|||
|
|
@ -90,7 +90,7 @@
|
|||
"glob": "13.0.0",
|
||||
"minimatch": "10.2.4",
|
||||
"lodash": "4.18.1",
|
||||
"ws": "8.19.0",
|
||||
"ws": "8.20.1",
|
||||
"braces": "3.0.3",
|
||||
"axios": "1.13.6",
|
||||
"postcss": "8.5.13"
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue