Merge branch 'BerriAI:litellm_internal_staging' into litellm_internal_staging

This commit is contained in:
bhumikadangayach 2026-06-09 13:44:07 +05:30 • committed by GitHub
commit 00915f15fe
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
128 changed files with 7531 additions and 904 deletions

View file

@ -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:

View file

@ -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)

View file

@ -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",

View file

@ -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}

View file

@ -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",

View file

@ -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"
)

View 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

View file

@ -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",

View file

@ -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")

View file

@ -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(

View file

@ -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:
"""

View file

@ -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:")
):

View file

@ -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

View file

@ -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

View file

@ -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."

View file

@ -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,

View file

@ -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,

View file

@ -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()

View 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",
)

View file

@ -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:

View file

@ -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)}"
)

View 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

View file

@ -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]:

View file

@ -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

View file

@ -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():

View file

@ -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)

View file

@ -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,

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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.

View file

@ -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]:
"""

View file

@ -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

View file

@ -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
}
}
}

View file

@ -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]]:

View file

@ -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,

View file

@ -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,

View file

@ -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:

View file

@ -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(

View file

@ -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,

View file

@ -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.

View file

@ -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:

View file

@ -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.",

View file

@ -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

View file

@ -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,

View file

@ -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
)

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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,
}

View file

@ -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(

View file

@ -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
):

View file

@ -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

View file

@ -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

View file

@ -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",
)

View file

@ -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"

View file

@ -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:

View file

@ -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,

View file

@ -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",

View file

@ -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",

View file

@ -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},
}

View 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)

View file

@ -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"),
)

View file

@ -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

View file

@ -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"

View 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"

View file

@ -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)

View file

@ -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"]
)

View file

@ -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

View file

@ -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

View file

@ -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"]}
]

View file

@ -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'."""

View file

@ -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.

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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")

View file

@ -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

View file

@ -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):

View file

@ -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__])

View 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)

View 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"

View file

@ -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.

View file

@ -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

View file

@ -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."""

View file

@ -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.

View file

@ -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()

View file

@ -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,

View file

@ -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():
"""

View file

@ -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

View file

@ -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

View file

@ -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"""

View file

@ -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(

View file

@ -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
# ---------------------------------------------------------------------------

View file

@ -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"""

View file

@ -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

View file

@ -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

View file

@ -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": {

View file

@ -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