+ ⚠️ Note: Your API requests will continue to work, but you should monitor your usage closely.
+ If you reach your maximum budget, requests will be rejected.
+
+
+ You can view your usage and manage your budget in the LiteLLM Dashboard.
+
+ If you have any questions, please send an email to {email_support_contact}
+
+ Best,
+ The LiteLLM team
+"""
+
MAX_BUDGET_ALERT_EMAIL_TEMPLATE = """
diff --git a/litellm/integrations/langfuse/langfuse_otel.py b/litellm/integrations/langfuse/langfuse_otel.py
index 08493a0e8ec..8955d3619f7 100644
--- a/litellm/integrations/langfuse/langfuse_otel.py
+++ b/litellm/integrations/langfuse/langfuse_otel.py
@@ -8,9 +8,8 @@ from litellm.integrations.arize import _utils
from litellm.integrations.langfuse.langfuse_otel_attributes import (
LangfuseLLMObsOTELAttributes,
)
-from litellm.integrations.opentelemetry import OpenTelemetry
+from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig
from litellm.types.integrations.langfuse_otel import (
- LangfuseOtelConfig,
LangfuseSpanAttributes,
)
from litellm.types.utils import StandardCallbackDynamicParams
@@ -18,17 +17,8 @@ from litellm.types.utils import StandardCallbackDynamicParams
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
- from litellm.integrations.opentelemetry import (
- OpenTelemetryConfig as _OpenTelemetryConfig,
- )
- from litellm.types.integrations.arize import Protocol as _Protocol
-
- Protocol = _Protocol
- OpenTelemetryConfig = _OpenTelemetryConfig
Span = Union[_Span, Any]
else:
- Protocol = Any
- OpenTelemetryConfig = Any
Span = Any
@@ -37,8 +27,12 @@ LANGFUSE_CLOUD_US_ENDPOINT = "https://us.cloud.langfuse.com/api/public/otel"
class LangfuseOtelLogger(OpenTelemetry):
- def __init__(self, *args, **kwargs):
- super().__init__(*args, **kwargs)
+ def __init__(self, config=None, *args, **kwargs):
+ # Prevent LangfuseOtelLogger from modifying global environment variables by constructing config manually
+ # and passing it to the parent OpenTelemetry class
+ if config is None:
+ config = self._create_open_telemetry_config_from_langfuse_env()
+ super().__init__(config=config, *args, **kwargs)
@staticmethod
def set_langfuse_otel_attributes(span: Span, kwargs, response_obj):
@@ -114,6 +108,10 @@ class LangfuseOtelLogger(OpenTelemetry):
for key, enum_attr in mapping.items():
if key in metadata and metadata[key] is not None:
value = metadata[key]
+ if key == "trace_id" and isinstance(value, str):
+ # trace_id must be 32 hex char no dashes for langfuse : Litellm sends uuid with dashes (might be breaking at some point)
+ value = value.replace("-", "")
+
if isinstance(value, (list, dict)):
try:
value = json.dumps(value)
@@ -265,8 +263,47 @@ class LangfuseOtelLogger(OpenTelemetry):
"""
return os.environ.get("LANGFUSE_OTEL_HOST") or os.environ.get("LANGFUSE_HOST")
+ def _create_open_telemetry_config_from_langfuse_env(self) -> OpenTelemetryConfig:
+ """
+ Creates OpenTelemetryConfig from Langfuse environment variables.
+ Does NOT modify global environment variables.
+ """
+ from litellm.integrations.opentelemetry import OpenTelemetryConfig
+
+ public_key = os.environ.get("LANGFUSE_PUBLIC_KEY", None)
+ secret_key = os.environ.get("LANGFUSE_SECRET_KEY", None)
+
+ if not public_key or not secret_key:
+ # If no keys, return default from env (likely logging to console or something else)
+ return OpenTelemetryConfig.from_env()
+
+ # Determine endpoint - default to US cloud
+ langfuse_host = LangfuseOtelLogger._get_langfuse_otel_host()
+
+ if langfuse_host:
+ # If LANGFUSE_HOST is provided, construct OTEL endpoint from it
+ if not langfuse_host.startswith("http"):
+ langfuse_host = "https://" + langfuse_host
+ endpoint = f"{langfuse_host.rstrip('/')}/api/public/otel"
+ verbose_logger.debug(f"Using Langfuse OTEL endpoint from host: {endpoint}")
+ else:
+ # Default to US cloud endpoint
+ endpoint = LANGFUSE_CLOUD_US_ENDPOINT
+ verbose_logger.debug(f"Using Langfuse US cloud endpoint: {endpoint}")
+
+ auth_header = LangfuseOtelLogger._get_langfuse_authorization_header(
+ public_key=public_key, secret_key=secret_key
+ )
+ otlp_auth_headers = f"Authorization={auth_header}"
+
+ return OpenTelemetryConfig(
+ exporter="otlp_http",
+ endpoint=endpoint,
+ headers=otlp_auth_headers,
+ )
+
@staticmethod
- def get_langfuse_otel_config() -> LangfuseOtelConfig:
+ def get_langfuse_otel_config() -> "OpenTelemetryConfig":
"""
Retrieves the Langfuse OpenTelemetry configuration based on environment variables.
@@ -276,7 +313,7 @@ class LangfuseOtelLogger(OpenTelemetry):
LANGFUSE_HOST: Optional. Custom Langfuse host URL. Defaults to US cloud.
Returns:
- LangfuseOtelConfig: A Pydantic model containing Langfuse OTEL configuration.
+ OpenTelemetryConfig: A Pydantic model containing Langfuse OTEL configuration.
Raises:
ValueError: If required keys are missing.
@@ -308,12 +345,14 @@ class LangfuseOtelLogger(OpenTelemetry):
)
otlp_auth_headers = f"Authorization={auth_header}"
- # Set standard OTEL environment variables
- os.environ["OTEL_EXPORTER_OTLP_ENDPOINT"] = endpoint
- os.environ["OTEL_EXPORTER_OTLP_HEADERS"] = otlp_auth_headers
+ # Prevent modification of global env vars which causes leakage
+ # os.environ["OTEL_EXPORTER_OTLP_ENDPOINT"] = endpoint
+ # os.environ["OTEL_EXPORTER_OTLP_HEADERS"] = otlp_auth_headers
- return LangfuseOtelConfig(
- otlp_auth_headers=otlp_auth_headers, protocol="otlp_http"
+ return OpenTelemetryConfig(
+ exporter="otlp_http",
+ endpoint=endpoint,
+ headers=otlp_auth_headers,
)
@staticmethod
diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py
index 18898be7dce..296a88f9a0b 100644
--- a/litellm/integrations/opentelemetry.py
+++ b/litellm/integrations/opentelemetry.py
@@ -599,9 +599,9 @@ class OpenTelemetry(CustomLogger):
def _get_dynamic_otel_headers_from_kwargs(self, kwargs) -> Optional[dict]:
"""Extract dynamic headers from kwargs if available."""
- standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = (
- kwargs.get("standard_callback_dynamic_params")
- )
+ standard_callback_dynamic_params: Optional[
+ StandardCallbackDynamicParams
+ ] = kwargs.get("standard_callback_dynamic_params")
if not standard_callback_dynamic_params:
return None
@@ -619,7 +619,9 @@ class OpenTelemetry(CustomLogger):
# Prevents thread exhaustion by reusing providers for the same credential sets (e.g. per-team keys)
cache_key = str(sorted(dynamic_headers.items()))
if cache_key in self._tracer_provider_cache:
- return self._tracer_provider_cache[cache_key].get_tracer(LITELLM_TRACER_NAME)
+ return self._tracer_provider_cache[cache_key].get_tracer(
+ LITELLM_TRACER_NAME
+ )
# Create a temporary tracer provider with dynamic headers
temp_provider = TracerProvider(resource=self._get_litellm_resource(self.config))
@@ -674,7 +676,10 @@ class OpenTelemetry(CustomLogger):
kwargs, response_obj, start_time, end_time, span
)
# Ensure proxy-request parent span is annotated with the actual operation kind
- if parent_span is not None and parent_span.name == LITELLM_PROXY_REQUEST_SPAN_NAME:
+ if (
+ parent_span is not None
+ and parent_span.name == LITELLM_PROXY_REQUEST_SPAN_NAME
+ ):
self.set_attributes(parent_span, kwargs, response_obj)
else:
# Do not create primary span (keep hierarchy shallow when parent exists)
@@ -1003,14 +1008,11 @@ class OpenTelemetry(CustomLogger):
# TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords
from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider
+
try:
- from opentelemetry.sdk._logs import (
- LogRecord as SdkLogRecord, # type: ignore[attr-defined] # OTEL < 1.39.0
- )
+ from opentelemetry.sdk._logs import LogRecord as SdkLogRecord # type: ignore[attr-defined] # OTEL < 1.39.0
except ImportError:
- from opentelemetry.sdk._logs._internal import (
- LogRecord as SdkLogRecord, # OTEL >= 1.39.0
- )
+ from opentelemetry.sdk._logs._internal import LogRecord as SdkLogRecord # type: ignore[attr-defined, no-redef] # OTEL >= 1.39.0
otel_logger = get_logger(LITELLM_LOGGER_NAME)
@@ -1618,7 +1620,6 @@ class OpenTelemetry(CustomLogger):
for idx, choice in enumerate(response_obj.get("choices")):
if choice.get("finish_reason"):
-
message = choice.get("message")
tool_calls = message.get("tool_calls")
if tool_calls:
@@ -1631,7 +1632,9 @@ class OpenTelemetry(CustomLogger):
)
except Exception as e:
- self.handle_callback_failure(callback_name=self.callback_name or "opentelemetry")
+ self.handle_callback_failure(
+ callback_name=self.callback_name or "opentelemetry"
+ )
verbose_logger.exception(
"OpenTelemetry logging error in set_attributes %s", str(e)
)
@@ -1722,6 +1725,7 @@ class OpenTelemetry(CustomLogger):
def set_raw_request_attributes(self, span: Span, kwargs, response_obj):
try:
+ self.set_attributes(span, kwargs, response_obj)
kwargs.get("optional_params", {})
litellm_params = kwargs.get("litellm_params", {}) or {}
custom_llm_provider = litellm_params.get("custom_llm_provider", "Unknown")
diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py
index 2c897cb0692..0a61dab0680 100644
--- a/litellm/integrations/prometheus.py
+++ b/litellm/integrations/prometheus.py
@@ -1,6 +1,7 @@
# used for /metrics endpoint on LiteLLM Proxy
#### What this does ####
# On success, log events to Prometheus
+import asyncio
import os
import sys
from datetime import datetime, timedelta
@@ -1188,28 +1189,34 @@ class PrometheusLogger(CustomLogger):
_user_spend = _metadata.get("user_api_key_user_spend", None)
_user_max_budget = _metadata.get("user_api_key_user_max_budget", None)
- await self._set_api_key_budget_metrics_after_api_request(
- user_api_key=user_api_key,
- user_api_key_alias=user_api_key_alias,
- response_cost=response_cost,
- key_max_budget=_api_key_max_budget,
- key_spend=_api_key_spend,
- )
-
- await self._set_team_budget_metrics_after_api_request(
- user_api_team=user_api_team,
- user_api_team_alias=user_api_team_alias,
- team_spend=_team_spend,
- team_max_budget=_team_max_budget,
- response_cost=response_cost,
- )
-
- await self._set_user_budget_metrics_after_api_request(
- user_id=user_id,
- user_spend=_user_spend,
- user_max_budget=_user_max_budget,
- response_cost=response_cost,
+ results = await asyncio.gather(
+ self._set_api_key_budget_metrics_after_api_request(
+ user_api_key=user_api_key,
+ user_api_key_alias=user_api_key_alias,
+ response_cost=response_cost,
+ key_max_budget=_api_key_max_budget,
+ key_spend=_api_key_spend,
+ ),
+ self._set_team_budget_metrics_after_api_request(
+ user_api_team=user_api_team,
+ user_api_team_alias=user_api_team_alias,
+ team_spend=_team_spend,
+ team_max_budget=_team_max_budget,
+ response_cost=response_cost,
+ ),
+ self._set_user_budget_metrics_after_api_request(
+ user_id=user_id,
+ user_spend=_user_spend,
+ user_max_budget=_user_max_budget,
+ response_cost=response_cost,
+ ),
+ return_exceptions=True,
)
+ for i, r in enumerate(results):
+ if isinstance(r, Exception):
+ verbose_logger.debug(
+ f"[Non-Blocking] Prometheus: Budget metric lookup {['key', 'team', 'user'][i]} failed: {r}"
+ )
def _increment_top_level_request_and_spend_metrics(
self,
@@ -1683,6 +1690,108 @@ class PrometheusLogger(CustomLogger):
)
pass
+ def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any:
+ """Get value from dict or Pydantic model."""
+ if obj is None:
+ return default
+ if isinstance(obj, dict):
+ return obj.get(key, default)
+ return getattr(obj, key, default)
+
+ def _extract_deployment_failure_label_values(
+ self, request_kwargs: dict
+ ) -> Dict[str, Optional[str]]:
+ """
+ Extract label values for deployment failure metrics from all available
+ sources in request_kwargs. Falls back to litellm_params metadata and
+ user_api_key_auth when standard_logging_payload has None values.
+ """
+ standard_logging_payload = (
+ request_kwargs.get("standard_logging_object", {}) or {}
+ )
+ _litellm_params = request_kwargs.get("litellm_params", {}) or {}
+ _metadata_raw = self._safe_get(standard_logging_payload, "metadata") or {}
+ if isinstance(_metadata_raw, dict):
+ _metadata = _metadata_raw
+ else:
+ _metadata = {
+ "user_api_key_alias": getattr(
+ _metadata_raw, "user_api_key_alias", None
+ ),
+ "user_api_key_team_id": getattr(
+ _metadata_raw, "user_api_key_team_id", None
+ ),
+ "user_api_key_team_alias": getattr(
+ _metadata_raw, "user_api_key_team_alias", None
+ ),
+ "user_api_key_hash": getattr(_metadata_raw, "user_api_key_hash", None),
+ "requester_ip_address": getattr(
+ _metadata_raw, "requester_ip_address", None
+ ),
+ "user_agent": getattr(_metadata_raw, "user_agent", None),
+ }
+ _litellm_params_metadata = _litellm_params.get("metadata", {}) or {}
+
+ # Extract user_api_key_auth if present (proxy injects this, skipped in merge)
+ user_api_key_auth = _litellm_params_metadata.get("user_api_key_auth")
+
+ def _get_api_key_alias() -> Optional[str]:
+ val = _metadata.get("user_api_key_alias")
+ if val is not None:
+ return val
+ val = _litellm_params_metadata.get("user_api_key_alias")
+ if val is not None:
+ return val
+ if user_api_key_auth is not None:
+ return getattr(user_api_key_auth, "key_alias", None)
+ return None
+
+ def _get_team_id() -> Optional[str]:
+ val = _metadata.get("user_api_key_team_id")
+ if val is not None:
+ return val
+ val = _litellm_params_metadata.get("user_api_key_team_id")
+ if val is not None:
+ return val
+ if user_api_key_auth is not None:
+ return getattr(user_api_key_auth, "team_id", None)
+ return None
+
+ def _get_team_alias() -> Optional[str]:
+ val = _metadata.get("user_api_key_team_alias")
+ if val is not None:
+ return val
+ val = _litellm_params_metadata.get("user_api_key_team_alias")
+ if val is not None:
+ return val
+ if user_api_key_auth is not None:
+ return getattr(user_api_key_auth, "team_alias", None)
+ return None
+
+ def _get_hashed_api_key() -> Optional[str]:
+ val = _metadata.get("user_api_key_hash")
+ if val is not None:
+ return val
+ val = _litellm_params_metadata.get("user_api_key_hash")
+ if val is not None:
+ return val
+ if user_api_key_auth is not None:
+ return getattr(user_api_key_auth, "api_key", None) or getattr(
+ user_api_key_auth, "api_key_hash", None
+ )
+ return None
+
+ return {
+ "api_key_alias": _get_api_key_alias(),
+ "team": _get_team_id(),
+ "team_alias": _get_team_alias(),
+ "hashed_api_key": _get_hashed_api_key(),
+ "client_ip": _metadata.get("requester_ip_address")
+ or _litellm_params_metadata.get("requester_ip_address"),
+ "user_agent": _metadata.get("user_agent")
+ or _litellm_params_metadata.get("user_agent"),
+ }
+
def set_llm_deployment_failure_metrics(self, request_kwargs: dict):
"""
Sets Failure metrics when an LLM API call fails
@@ -1707,6 +1816,21 @@ class PrometheusLogger(CustomLogger):
model_id = standard_logging_payload.get("model_id", None)
exception = request_kwargs.get("exception", None)
+ # Fallback: model_id from litellm_metadata.model_info
+ if model_id is None:
+ _model_info = (
+ (_litellm_params.get("litellm_metadata") or {}).get("model_info")
+ or (_litellm_params.get("metadata") or {}).get("model_info")
+ or {}
+ )
+ model_id = _model_info.get("id")
+
+ # Fallback: model_group from litellm_metadata
+ if model_group is None:
+ model_group = (_litellm_params.get("litellm_metadata") or {}).get(
+ "model_group"
+ ) or (_litellm_params.get("metadata") or {}).get("model_group")
+
llm_provider = _litellm_params.get("custom_llm_provider", None)
if self._should_skip_metrics_for_invalid_key(
@@ -1714,9 +1838,37 @@ class PrometheusLogger(CustomLogger):
standard_logging_payload=standard_logging_payload,
):
return
- hashed_api_key = standard_logging_payload.get("metadata", {}).get(
+
+ # Extract context labels from all available sources (fix for None labels)
+ fallback_values = self._extract_deployment_failure_label_values(
+ request_kwargs
+ )
+ _metadata = standard_logging_payload.get("metadata", {}) or {}
+ hashed_api_key = fallback_values.get("hashed_api_key") or _metadata.get(
"user_api_key_hash"
)
+ api_key_alias = fallback_values.get("api_key_alias") or _metadata.get(
+ "user_api_key_alias"
+ )
+ team = fallback_values.get("team") or _metadata.get("user_api_key_team_id")
+ team_alias = fallback_values.get("team_alias") or _metadata.get(
+ "user_api_key_team_alias"
+ )
+ client_ip = fallback_values.get("client_ip") or _metadata.get(
+ "requester_ip_address"
+ )
+ user_agent = fallback_values.get("user_agent") or _metadata.get(
+ "user_agent"
+ )
+
+ # exception_status: prefer status_code, fallback to exception class for known types
+ exception_status = None
+ if exception is not None:
+ exception_status = str(getattr(exception, "status_code", None))
+ if exception_status == "None" or not exception_status:
+ code = getattr(exception, "code", None)
+ if code is not None:
+ exception_status = str(code)
# Create enum_values for the label factory (always create for use in different metrics)
enum_values = UserAPIKeyLabelValues(
@@ -1724,26 +1876,18 @@ class PrometheusLogger(CustomLogger):
model_id=model_id,
api_base=api_base,
api_provider=llm_provider,
- exception_status=(
- str(getattr(exception, "status_code", None)) if exception else None
- ),
+ exception_status=exception_status,
exception_class=(
self._get_exception_class_name(exception) if exception else None
),
- requested_model=model_group,
+ requested_model=model_group or litellm_model_name,
hashed_api_key=hashed_api_key,
- api_key_alias=standard_logging_payload["metadata"][
- "user_api_key_alias"
- ],
- team=standard_logging_payload["metadata"]["user_api_key_team_id"],
- team_alias=standard_logging_payload["metadata"][
- "user_api_key_team_alias"
- ],
+ api_key_alias=api_key_alias,
+ team=team,
+ team_alias=team_alias,
tags=standard_logging_payload.get("request_tags", []),
- client_ip=standard_logging_payload["metadata"].get(
- "requester_ip_address"
- ),
- user_agent=standard_logging_payload["metadata"].get("user_agent"),
+ client_ip=client_ip,
+ user_agent=user_agent,
)
"""
@@ -2761,12 +2905,14 @@ class PrometheusLogger(CustomLogger):
max_budget=max_budget,
)
try:
+ # Note: Setting check_db_only=True bypasses cache and hits DB on every request,
+ # causing huge latency increase and CPU spikes. Keep check_db_only=False.
user_info = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
- check_db_only=True,
+ check_db_only=False,
)
except Exception as e:
verbose_logger.debug(
diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py
index 00695cbfb5b..7c8e2ebeaff 100644
--- a/litellm/litellm_core_utils/core_helpers.py
+++ b/litellm/litellm_core_utils/core_helpers.py
@@ -94,8 +94,8 @@ def map_finish_reason(
return "length"
elif finish_reason == "tool_use": # anthropic
return "tool_calls"
- elif finish_reason == "content_filtered":
- return "content_filter"
+ elif finish_reason == "compaction":
+ return "length"
return finish_reason
diff --git a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py
index c425319b4d4..ff521d47804 100644
--- a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py
+++ b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py
@@ -1,8 +1,35 @@
from typing import Dict, Optional
-
from litellm.secret_managers.main import get_secret_str
from litellm.types.utils import StandardCallbackDynamicParams
+# Hardcoded list of supported callback params to avoid runtime inspection issues with TypedDict
+_supported_callback_params = [
+ "langfuse_public_key",
+ "langfuse_secret",
+ "langfuse_secret_key",
+ "langfuse_host",
+ "langfuse_prompt_version",
+ "gcs_bucket_name",
+ "gcs_path_service_account",
+ "langsmith_api_key",
+ "langsmith_project",
+ "langsmith_base_url",
+ "langsmith_sampling_rate",
+ "langsmith_tenant_id",
+ "humanloop_api_key",
+ "arize_api_key",
+ "arize_space_key",
+ "arize_space_id",
+ "posthog_api_key",
+ "posthog_host",
+ "braintrust_api_key",
+ "braintrust_project",
+ "braintrust_host",
+ "slack_webhook_url",
+ "lunary_public_key",
+ "turn_off_message_logging",
+]
+
def initialize_standard_callback_dynamic_params(
kwargs: Optional[Dict] = None,
@@ -15,13 +42,10 @@ def initialize_standard_callback_dynamic_params(
standard_callback_dynamic_params = StandardCallbackDynamicParams()
if kwargs:
- _supported_callback_params = (
- StandardCallbackDynamicParams.__annotations__.keys()
- )
-
+ # 1. Check top-level kwargs
for param in _supported_callback_params:
if param in kwargs:
- _param_value = kwargs.pop(param)
+ _param_value = kwargs.get(param)
if (
_param_value is not None
and isinstance(_param_value, str)
@@ -30,4 +54,22 @@ def initialize_standard_callback_dynamic_params(
_param_value = get_secret_str(secret_name=_param_value)
standard_callback_dynamic_params[param] = _param_value # type: ignore
+ # 2. Fallback: check "metadata" or "litellm_params" -> "metadata"
+ metadata = (kwargs.get("metadata") or {}).copy()
+ litellm_params = kwargs.get("litellm_params") or {}
+ if isinstance(litellm_params, dict):
+ metadata.update(litellm_params.get("metadata") or {})
+
+ if isinstance(metadata, dict):
+ for param in _supported_callback_params:
+ if param not in standard_callback_dynamic_params and param in metadata:
+ _param_value = metadata.get(param)
+ if (
+ _param_value is not None
+ and isinstance(_param_value, str)
+ and "os.environ/" in _param_value
+ ):
+ _param_value = get_secret_str(secret_name=_param_value)
+ standard_callback_dynamic_params[param] = _param_value # type: ignore
+
return standard_callback_dynamic_params
diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py
index 4ad2d1002bc..14015225f38 100644
--- a/litellm/litellm_core_utils/litellm_logging.py
+++ b/litellm/litellm_core_utils/litellm_logging.py
@@ -2435,6 +2435,36 @@ class Logging(LiteLLMLoggingBaseClass):
standard_built_in_tools_params=self.standard_built_in_tools_params,
)
+ # print standard logging payload
+ if (
+ standard_logging_payload := self.model_call_details.get(
+ "standard_logging_object"
+ )
+ ) is not None:
+ emit_standard_logging_payload(standard_logging_payload)
+ elif self.call_type == "pass_through_endpoint":
+ print_verbose(
+ "Async success callbacks: Got a pass-through endpoint response"
+ )
+
+ self.model_call_details["async_complete_streaming_response"] = result
+
+ # cost calculation not possible for pass-through
+ self.model_call_details["response_cost"] = None
+
+ ## STANDARDIZED LOGGING PAYLOAD
+ self.model_call_details[
+ "standard_logging_object"
+ ] = get_standard_logging_object_payload(
+ kwargs=self.model_call_details,
+ init_response_obj=result,
+ start_time=start_time,
+ end_time=end_time,
+ logging_obj=self,
+ status="success",
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
+ )
+
# print standard logging payload
if (
standard_logging_payload := self.model_call_details.get(
@@ -3887,18 +3917,6 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
return langfuse_logger # type: ignore
elif logging_integration == "langfuse_otel":
from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger
- from litellm.integrations.opentelemetry import (
- OpenTelemetry,
- OpenTelemetryConfig,
- )
-
- langfuse_otel_config = LangfuseOtelLogger.get_langfuse_otel_config()
-
- # The endpoint and headers are now set as environment variables by get_langfuse_otel_config()
- otel_config = OpenTelemetryConfig(
- exporter=langfuse_otel_config.protocol,
- headers=langfuse_otel_config.otlp_auth_headers,
- )
for callback in _in_memory_loggers:
if (
@@ -3906,8 +3924,10 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
and callback.callback_name == "langfuse_otel"
):
return callback # type: ignore
+ # Allow LangfuseOtelLogger to initialize its own config safely
+ # This prevents startup crashes if LANGFUSE keys are not in env (e.g. for dynamic usage)
_otel_logger = LangfuseOtelLogger(
- config=otel_config, callback_name="langfuse_otel"
+ config=None, callback_name="langfuse_otel"
)
_in_memory_loggers.append(_otel_logger)
return _otel_logger # type: ignore
diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py
index fe06641a389..2308dc7beca 100644
--- a/litellm/litellm_core_utils/llm_cost_calc/utils.py
+++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py
@@ -215,6 +215,9 @@ def _get_token_base_cost(
cache_creation_tiered_key = (
f"cache_creation_input_token_cost_above_{threshold_str}_tokens"
)
+ cache_creation_1hr_tiered_key = (
+ f"cache_creation_input_token_cost_above_1hr_above_{threshold_str}_tokens"
+ )
cache_read_tiered_key = (
f"cache_read_input_token_cost_above_{threshold_str}_tokens"
)
@@ -229,6 +232,16 @@ def _get_token_base_cost(
),
)
+ if cache_creation_1hr_tiered_key in model_info:
+ cache_creation_cost_above_1hr = cast(
+ float,
+ _get_cost_per_unit(
+ model_info,
+ cache_creation_1hr_tiered_key,
+ cache_creation_cost_above_1hr,
+ ),
+ )
+
if cache_read_tiered_key in model_info:
cache_read_cost = cast(
float,
diff --git a/litellm/litellm_core_utils/logging_callback_manager.py b/litellm/litellm_core_utils/logging_callback_manager.py
index 4f76a5bad03..435ae078a65 100644
--- a/litellm/litellm_core_utils/logging_callback_manager.py
+++ b/litellm/litellm_core_utils/logging_callback_manager.py
@@ -114,6 +114,27 @@ class LoggingCallbackManager:
for c in remove_list:
callback_list.remove(c)
+ def remove_callbacks_by_type(self, callback_list, callback_type):
+ """
+ Remove all callbacks of a specific type from a callback list.
+
+ Args:
+ callback_list: The list to remove callbacks from (e.g., litellm.callbacks)
+ callback_type: The class type to match (e.g., SemanticToolFilterHook)
+
+ Example:
+ litellm.logging_callback_manager.remove_callbacks_by_type(
+ litellm.callbacks, SemanticToolFilterHook
+ )
+ """
+ if not isinstance(callback_list, list):
+ return
+
+ remove_list = [c for c in callback_list if isinstance(c, callback_type)]
+
+ for c in remove_list:
+ callback_list.remove(c)
+
def _add_string_callback_to_list(
self, callback: str, parent_list: List[Union[CustomLogger, Callable, str]]
):
diff --git a/litellm/litellm_core_utils/model_param_helper.py b/litellm/litellm_core_utils/model_param_helper.py
index 91f2f1341cf..4d45c47c224 100644
--- a/litellm/litellm_core_utils/model_param_helper.py
+++ b/litellm/litellm_core_utils/model_param_helper.py
@@ -17,15 +17,16 @@ from litellm.types.rerank import RerankRequest
class ModelParamHelper:
+ # Cached at class level — deterministic set built from static OpenAI type annotations
+ _relevant_logging_args: frozenset = frozenset()
+
@staticmethod
def get_standard_logging_model_parameters(
model_parameters: dict,
) -> dict:
""" """
standard_logging_model_parameters: dict = {}
- supported_model_parameters = (
- ModelParamHelper._get_relevant_args_to_use_for_logging()
- )
+ supported_model_parameters = ModelParamHelper._relevant_logging_args
for key, value in model_parameters.items():
if key in supported_model_parameters:
@@ -172,3 +173,8 @@ class ModelParamHelper:
Get the kwargs to exclude from the cache key
"""
return set(["metadata"])
+
+
+ModelParamHelper._relevant_logging_args = frozenset(
+ ModelParamHelper._get_relevant_args_to_use_for_logging()
+)
diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py
index 7790fb83361..b1c2d0a52f5 100644
--- a/litellm/litellm_core_utils/prompt_templates/common_utils.py
+++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py
@@ -443,13 +443,21 @@ def update_messages_with_model_file_ids(
def update_responses_input_with_model_file_ids(
input: Any,
+ model_id: Optional[str] = None,
+ model_file_id_mapping: Optional[Dict[str, Dict[str, str]]] = None,
) -> Union[str, List[Dict[str, Any]]]:
"""
Updates responses API input with provider-specific file IDs.
File IDs are always inside the content array, not as direct input_file items.
- For managed files (unified file IDs), decodes the base64-encoded unified file ID
- and extracts the llm_output_file_id directly.
+ For managed files (unified file IDs), uses model_file_id_mapping if provided,
+ otherwise decodes the base64-encoded unified file ID and extracts the llm_output_file_id directly.
+
+ Args:
+ input: The responses API input parameter
+ model_id: The model ID to use for looking up provider-specific file IDs
+ model_file_id_mapping: Dictionary mapping litellm file IDs to provider file IDs
+ Format: {"litellm_file_id": {"model_id": "provider_file_id"}}
"""
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
@@ -479,22 +487,35 @@ def update_responses_input_with_model_file_ids(
):
file_id = content_item.get("file_id")
if file_id:
- # Check if this is a managed file ID (base64-encoded unified file ID)
- is_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
- if is_unified_file_id:
- unified_file_id = convert_b64_uid_to_unified_uid(file_id)
- if "llm_output_file_id," in unified_file_id:
- provider_file_id = unified_file_id.split(
- "llm_output_file_id,"
- )[1].split(";")[0]
- else:
- # Fallback: keep original if we can't extract
- provider_file_id = file_id
+ provider_file_id = file_id # Default to original
+
+ # Check if we have a mapping for this file ID
+ if model_file_id_mapping and model_id and file_id in model_file_id_mapping:
+ # Use the model-specific file ID from mapping
+ provider_file_id = (
+ model_file_id_mapping.get(file_id, {}).get(model_id)
+ or file_id
+ )
updated_content_item = content_item.copy()
updated_content_item["file_id"] = provider_file_id
updated_content.append(updated_content_item)
else:
- updated_content.append(content_item)
+ # Check if this is a base64-encoded unified file ID without mapping
+ is_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
+ if is_unified_file_id:
+ # Fallback: decode unified file ID
+ unified_file_id = convert_b64_uid_to_unified_uid(file_id)
+ if "llm_output_file_id," in unified_file_id:
+ provider_file_id = unified_file_id.split(
+ "llm_output_file_id,"
+ )[1].split(";")[0]
+
+ updated_content_item = content_item.copy()
+ updated_content_item["file_id"] = provider_file_id
+ updated_content.append(updated_content_item)
+ else:
+ # Not a managed file, keep as-is
+ updated_content.append(content_item)
else:
updated_content.append(content_item)
else:
@@ -506,6 +527,68 @@ def update_responses_input_with_model_file_ids(
return updated_input
+def update_responses_tools_with_model_file_ids(
+ tools: Optional[List[Dict[str, Any]]],
+ model_id: Optional[str] = None,
+ model_file_id_mapping: Optional[Dict[str, Dict[str, str]]] = None,
+) -> Optional[List[Dict[str, Any]]]:
+ """
+ Updates responses API tools with provider-specific file IDs.
+
+ Handles code_interpreter tools with container.file_ids.
+
+ Args:
+ tools: The responses API tools parameter
+ model_id: The model ID to use for looking up provider-specific file IDs
+ model_file_id_mapping: Dictionary mapping litellm file IDs to provider file IDs
+ Format: {"litellm_file_id": {"model_id": "provider_file_id"}}
+ """
+ if not tools or not isinstance(tools, list):
+ return tools
+
+ if not model_file_id_mapping or not model_id:
+ return tools
+
+ updated_tools = []
+ for tool in tools:
+ if not isinstance(tool, dict):
+ updated_tools.append(tool)
+ continue
+
+ updated_tool = tool.copy()
+
+ # Handle code_interpreter with container file_ids
+ if tool.get("type") == "code_interpreter":
+ container = tool.get("container")
+ if isinstance(container, dict):
+ container_file_ids = container.get("file_ids")
+ if isinstance(container_file_ids, list):
+ updated_file_ids = []
+ for file_id in container_file_ids:
+ if isinstance(file_id, str):
+ # Check if we have a mapping for this file ID
+ if file_id in model_file_id_mapping:
+ # Map to provider-specific file ID
+ provider_file_id = (
+ model_file_id_mapping.get(file_id, {}).get(model_id)
+ or file_id
+ )
+ updated_file_ids.append(provider_file_id)
+ else:
+ updated_file_ids.append(file_id)
+ else:
+ updated_file_ids.append(file_id)
+
+ # Update the tool with new file IDs
+ updated_container = container.copy()
+ updated_container["file_ids"] = updated_file_ids
+ updated_tool["container"] = updated_container
+
+ updated_tools.append(updated_tool)
+
+ return updated_tools
+
+
def extract_file_data(file_data: FileTypes) -> ExtractedFileData:
"""
Extracts and processes file data from various input formats.
diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py
index 0e1637a65ba..f9ecd78ff1c 100644
--- a/litellm/litellm_core_utils/prompt_templates/factory.py
+++ b/litellm/litellm_core_utils/prompt_templates/factory.py
@@ -2190,6 +2190,16 @@ def anthropic_messages_pt( # noqa: PLR0915
while msg_i < len(messages) and messages[msg_i]["role"] == "assistant":
assistant_content_block: ChatCompletionAssistantMessage = messages[msg_i] # type: ignore
+ # Extract compaction_blocks from provider_specific_fields and add them first
+ _provider_specific_fields_raw = assistant_content_block.get(
+ "provider_specific_fields"
+ )
+ if isinstance(_provider_specific_fields_raw, dict):
+ _compaction_blocks = _provider_specific_fields_raw.get("compaction_blocks")
+ if _compaction_blocks and isinstance(_compaction_blocks, list):
+ # Add compaction blocks at the beginning of assistant content : https://platform.claude.com/docs/en/build-with-claude/compaction
+ assistant_content.extend(_compaction_blocks) # type: ignore
+
thinking_blocks = assistant_content_block.get("thinking_blocks", None)
if (
thinking_blocks is not None
@@ -3399,6 +3409,59 @@ def _convert_to_bedrock_tool_call_result(
return content_block
+def _deduplicate_bedrock_content_blocks(
+ blocks: List[BedrockContentBlock],
+ block_key: str,
+ id_key: str = "toolUseId",
+) -> List[BedrockContentBlock]:
+ """
+ Remove duplicate content blocks that share the same ID under ``block_key``.
+
+ Bedrock requires all toolResult and toolUse IDs within a single message to
+ be unique. When merging consecutive messages, duplicates can occur if the
+ same tool_call_id appears multiple times in conversation history.
+
+ When duplicates exist, the first occurrence is retained and subsequent ones
+ are discarded. A warning is logged for every dropped block so that
+ upstream duplication bugs remain visible.
+
+ Blocks that do not contain ``block_key`` (e.g., cachePoint, text) are
+ always preserved.
+
+ Args:
+ blocks: The list of Bedrock content blocks to deduplicate.
+ block_key: The dict key to inspect (e.g. ``"toolResult"`` or ``"toolUse"``).
+ id_key: The nested key that holds the unique ID (default ``"toolUseId"``).
+ """
+ seen_ids: Set[str] = set()
+ deduplicated: List[BedrockContentBlock] = []
+ for block in blocks:
+ keyed = block.get(block_key)
+ if keyed is not None and isinstance(keyed, dict):
+ block_id = keyed.get(id_key)
+ if block_id:
+ if block_id in seen_ids:
+ verbose_logger.warning(
+ "Bedrock Converse: dropping duplicate %s block with "
+ "%s=%s. This may indicate duplicate tool messages in "
+ "conversation history.",
+ block_key,
+ id_key,
+ block_id,
+ )
+ continue
+ seen_ids.add(block_id)
+ deduplicated.append(block)
+ return deduplicated
+
+
+def _deduplicate_bedrock_tool_content(
+ tool_content: List[BedrockContentBlock],
+) -> List[BedrockContentBlock]:
+ """Convenience wrapper: deduplicate ``toolResult`` blocks by ``toolUseId``."""
+ return _deduplicate_bedrock_content_blocks(tool_content, "toolResult")
+
+
def _insert_assistant_continue_message(
messages: List[BedrockMessageBlock],
assistant_continue_message: Optional[
@@ -3867,6 +3930,8 @@ class BedrockConverseMessagesProcessor:
tool_content.append(cache_point_block)
msg_i += 1
+ # Deduplicate toolResult blocks with the same toolUseId
+ tool_content = _deduplicate_bedrock_tool_content(tool_content)
if tool_content:
# if last message was a 'user' message, then add a blank assistant message (bedrock requires alternating roles)
if len(contents) > 0 and contents[-1]["role"] == "user":
@@ -3932,10 +3997,12 @@ class BedrockConverseMessagesProcessor:
assistant_parts=assistants_parts,
)
elif element["type"] == "text":
- assistants_part = BedrockContentBlock(
- text=element["text"]
- )
- assistants_parts.append(assistants_part)
+ # Skip completely empty strings to avoid blank content blocks
+ if element.get("text", "").strip():
+ assistants_part = BedrockContentBlock(
+ text=element["text"]
+ )
+ assistants_parts.append(assistants_part)
elif element["type"] == "image_url":
if isinstance(element["image_url"], dict):
image_url = element["image_url"]["url"]
@@ -3960,9 +4027,12 @@ class BedrockConverseMessagesProcessor:
elif _assistant_content is not None and isinstance(
_assistant_content, str
):
- assistant_content.append(
- BedrockContentBlock(text=_assistant_content)
- )
+ # Skip completely empty strings to avoid blank content blocks
+ if _assistant_content.strip():
+ assistant_content.append(
+ BedrockContentBlock(text=_assistant_content)
+ )
+ # If content is empty/whitespace, skip it (don't add a placeholder)
# Add cache point block for assistant string content
_cache_point_block = (
litellm.AmazonConverseConfig()._get_cache_point_block(
@@ -3980,6 +4050,8 @@ class BedrockConverseMessagesProcessor:
msg_i += 1
+ assistant_content = _deduplicate_bedrock_content_blocks(assistant_content, "toolUse")
+
if assistant_content:
contents.append(
BedrockMessageBlock(role="assistant", content=assistant_content)
@@ -4230,6 +4302,8 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
tool_content.append(cache_point_block)
msg_i += 1
+ # Deduplicate toolResult blocks with the same toolUseId
+ tool_content = _deduplicate_bedrock_tool_content(tool_content)
if tool_content:
# if last message was a 'user' message, then add a blank assistant message (bedrock requires alternating roles)
if len(contents) > 0 and contents[-1]["role"] == "user":
@@ -4289,12 +4363,11 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
assistant_parts=assistants_parts,
)
elif element["type"] == "text":
- # AWS Bedrock doesn't allow empty or whitespace-only text content, so use placeholder for empty strings
- text_content = (
- element["text"] if element["text"].strip() else "."
- )
- assistants_part = BedrockContentBlock(text=text_content)
- assistants_parts.append(assistants_part)
+ # AWS Bedrock doesn't allow empty or whitespace-only text content
+ # Skip completely empty strings to avoid blank content blocks
+ if element.get("text", "").strip():
+ assistants_part = BedrockContentBlock(text=element["text"])
+ assistants_parts.append(assistants_part)
elif element["type"] == "image_url":
if isinstance(element["image_url"], dict):
image_url = element["image_url"]["url"]
@@ -4317,9 +4390,9 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
assistants_parts.append(_cache_point_block)
assistant_content.extend(assistants_parts)
elif _assistant_content is not None and isinstance(_assistant_content, str):
- # AWS Bedrock doesn't allow empty or whitespace-only text content, so use placeholder for empty strings
- text_content = _assistant_content if _assistant_content.strip() else "."
- assistant_content.append(BedrockContentBlock(text=text_content))
+ # Skip completely empty strings to avoid blank content blocks
+ if _assistant_content.strip():
+ assistant_content.append(BedrockContentBlock(text=_assistant_content))
# Add cache point block for assistant string content
_cache_point_block = (
litellm.AmazonConverseConfig()._get_cache_point_block(
@@ -4336,6 +4409,8 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
msg_i += 1
+ assistant_content = _deduplicate_bedrock_content_blocks(assistant_content, "toolUse")
+
if assistant_content:
contents.append(
BedrockMessageBlock(role="assistant", content=assistant_content)
diff --git a/litellm/litellm_core_utils/redact_messages.py b/litellm/litellm_core_utils/redact_messages.py
index 0effed3db70..aa763dc9899 100644
--- a/litellm/litellm_core_utils/redact_messages.py
+++ b/litellm/litellm_core_utils/redact_messages.py
@@ -130,6 +130,11 @@ def perform_redaction(model_call_details: dict, result):
def should_redact_message_logging(model_call_details: dict) -> bool:
"""
Determine if message logging should be redacted.
+
+ Priority order:
+ 1. Dynamic parameter (turn_off_message_logging in request)
+ 2. Headers (litellm-disable-message-redaction / litellm-enable-message-redaction)
+ 3. Global setting (litellm.turn_off_message_logging)
"""
litellm_params = model_call_details.get("litellm_params", {})
@@ -139,36 +144,36 @@ def should_redact_message_logging(model_call_details: dict) -> bool:
# Get headers from the metadata
request_headers = metadata.get("headers", {}) if isinstance(metadata, dict) else {}
- possible_request_headers = [
+ # Check for headers that explicitly control redaction
+ if request_headers and bool(
+ request_headers.get("litellm-disable-message-redaction", False)
+ ):
+ # User explicitly disabled redaction via header
+ return False
+
+ possible_enable_headers = [
"litellm-enable-message-redaction", # old header. maintain backwards compatibility
"x-litellm-enable-message-redaction", # new header
]
is_redaction_enabled_via_header = False
- for header in possible_request_headers:
+ for header in possible_enable_headers:
if bool(request_headers.get(header, False)):
is_redaction_enabled_via_header = True
break
- # check if user opted out of logging message/response to callbacks
- if (
- litellm.turn_off_message_logging is not True
- and is_redaction_enabled_via_header is not True
- and _get_turn_off_message_logging_from_dynamic_params(model_call_details)
- is not True
- ):
- return False
-
- if request_headers and bool(
- request_headers.get("litellm-disable-message-redaction", False)
- ):
- return False
-
- # user has OPTED OUT of message redaction
- if _get_turn_off_message_logging_from_dynamic_params(model_call_details) is False:
- return False
-
- return True
+ # Priority 1: Check dynamic parameter first (if explicitly set)
+ dynamic_turn_off = _get_turn_off_message_logging_from_dynamic_params(model_call_details)
+ if dynamic_turn_off is not None:
+ # Dynamic parameter is explicitly set, use it
+ return dynamic_turn_off
+
+ # Priority 2: Check if header explicitly enables redaction
+ if is_redaction_enabled_via_header:
+ return True
+
+ # Priority 3: Fall back to global setting
+ return litellm.turn_off_message_logging is True
def redact_message_input_output_from_logging(
diff --git a/litellm/llms/a2a/__init__.py b/litellm/llms/a2a/__init__.py
new file mode 100644
index 00000000000..043efa5e8bf
--- /dev/null
+++ b/litellm/llms/a2a/__init__.py
@@ -0,0 +1,6 @@
+"""
+A2A (Agent-to-Agent) Protocol Provider for LiteLLM
+"""
+from .chat.transformation import A2AConfig
+
+__all__ = ["A2AConfig"]
diff --git a/litellm/llms/a2a/chat/__init__.py b/litellm/llms/a2a/chat/__init__.py
new file mode 100644
index 00000000000..76bf4dd71d9
--- /dev/null
+++ b/litellm/llms/a2a/chat/__init__.py
@@ -0,0 +1,6 @@
+"""
+A2A Chat Completion Implementation
+"""
+from .transformation import A2AConfig
+
+__all__ = ["A2AConfig"]
diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py
new file mode 100644
index 00000000000..4b689414ddd
--- /dev/null
+++ b/litellm/llms/a2a/chat/streaming_iterator.py
@@ -0,0 +1,103 @@
+"""
+A2A Streaming Response Iterator
+"""
+from typing import Optional, Union
+
+from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
+from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
+
+from ..common_utils import extract_text_from_a2a_response
+
+
+class A2AModelResponseIterator(BaseModelResponseIterator):
+ """
+ Iterator for parsing A2A streaming responses.
+
+ Converts A2A JSON-RPC streaming chunks to OpenAI-compatible format.
+ """
+
+ def __init__(
+ self,
+ streaming_response,
+ sync_stream: bool,
+ json_mode: Optional[bool] = False,
+ model: str = "a2a/agent",
+ ):
+ super().__init__(
+ streaming_response=streaming_response,
+ sync_stream=sync_stream,
+ json_mode=json_mode,
+ )
+ self.model = model
+
+ def chunk_parser(self, chunk: dict) -> Union[GenericStreamingChunk, ModelResponseStream]:
+ """
+ Parse A2A streaming chunk to OpenAI format.
+
+ A2A chunk format:
+ {
+ "jsonrpc": "2.0",
+ "id": "request-id",
+ "result": {
+ "message": {
+ "parts": [{"kind": "text", "text": "content"}]
+ }
+ }
+ }
+
+ Or for tasks:
+ {
+ "jsonrpc": "2.0",
+ "result": {
+ "kind": "task",
+ "status": {"state": "running"},
+ "artifacts": [{"parts": [{"kind": "text", "text": "content"}]}]
+ }
+ }
+ """
+ try:
+ # Extract text from A2A response
+ text = extract_text_from_a2a_response(chunk)
+
+ # Determine finish reason
+ finish_reason = self._get_finish_reason(chunk)
+
+ # Return generic streaming chunk
+ return GenericStreamingChunk(
+ text=text,
+ is_finished=bool(finish_reason),
+ finish_reason=finish_reason or "",
+ usage=None,
+ index=0,
+ tool_use=None,
+ )
+ except Exception:
+ # Return empty chunk on parse error
+ return GenericStreamingChunk(
+ text="",
+ is_finished=False,
+ finish_reason="",
+ usage=None,
+ index=0,
+ tool_use=None,
+ )
+
+ def _get_finish_reason(self, chunk: dict) -> Optional[str]:
+ """Extract finish reason from A2A chunk"""
+ result = chunk.get("result", {})
+
+ # Check for task completion
+ if isinstance(result, dict):
+ status = result.get("status", {})
+ if isinstance(status, dict):
+ state = status.get("state")
+ if state == "completed":
+ return "stop"
+ elif state == "failed":
+ return "stop" # Map failed state to 'stop' (valid finish_reason)
+
+ # Check for [DONE] marker
+ if chunk.get("done") is True:
+ return "stop"
+
+ return None
diff --git a/litellm/llms/a2a/chat/transformation.py b/litellm/llms/a2a/chat/transformation.py
new file mode 100644
index 00000000000..163cd5ab22e
--- /dev/null
+++ b/litellm/llms/a2a/chat/transformation.py
@@ -0,0 +1,370 @@
+"""
+A2A Protocol Transformation for LiteLLM
+"""
+import uuid
+from typing import Any, Dict, Iterator, List, Optional, Union
+
+import httpx
+
+from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
+from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
+from litellm.types.llms.openai import AllMessageValues
+from litellm.types.utils import Choices, Message, ModelResponse
+
+from ..common_utils import (
+ A2AError,
+ convert_messages_to_prompt,
+ extract_text_from_a2a_response,
+)
+from .streaming_iterator import A2AModelResponseIterator
+
+
+class A2AConfig(BaseConfig):
+ """
+ Configuration for A2A (Agent-to-Agent) Protocol.
+
+ Handles transformation between OpenAI and A2A JSON-RPC 2.0 formats.
+ """
+
+ @staticmethod
+ def resolve_agent_config_from_registry(
+ model: str,
+ api_base: Optional[str],
+ api_key: Optional[str],
+ headers: Optional[Dict[str, Any]],
+ optional_params: Dict[str, Any],
+ ) -> tuple[Optional[str], Optional[str], Optional[Dict[str, Any]]]:
+ """
+ Resolve agent configuration from registry if model format is "a2a/".
+
+ Extracts agent name from model string and looks up configuration in the
+ agent registry (if available in proxy context).
+
+ Args:
+ model: Model string (e.g., "a2a/my-agent")
+ api_base: Explicit api_base (takes precedence over registry)
+ api_key: Explicit api_key (takes precedence over registry)
+ headers: Explicit headers (takes precedence over registry)
+ optional_params: Dict to merge additional litellm_params into
+
+ Returns:
+ Tuple of (api_base, api_key, headers) with registry values filled in
+ """
+ # Extract agent name from model (e.g., "a2a/my-agent" -> "my-agent")
+ agent_name = model.split("/", 1)[1] if "/" in model else None
+
+ # Only lookup if agent name exists and some config is missing
+ if not agent_name or (api_base is not None and api_key is not None and headers is not None):
+ return api_base, api_key, headers
+
+ # Try registry lookup (only available in proxy context)
+ try:
+ from litellm.proxy.agent_endpoints.agent_registry import (
+ global_agent_registry,
+ )
+
+ agent = global_agent_registry.get_agent_by_name(agent_name)
+ if agent:
+ # Get api_base from agent card URL
+ if api_base is None and agent.agent_card_params:
+ api_base = agent.agent_card_params.get("url")
+
+ # Get api_key, headers, and other params from litellm_params
+ if agent.litellm_params:
+ if api_key is None:
+ api_key = agent.litellm_params.get("api_key")
+
+ if headers is None:
+ agent_headers = agent.litellm_params.get("headers")
+ if agent_headers:
+ headers = agent_headers
+
+ # Merge other litellm_params (timeout, max_retries, etc.)
+ for key, value in agent.litellm_params.items():
+ if key not in ["api_key", "api_base", "headers", "model"] and key not in optional_params:
+ optional_params[key] = value
+ except ImportError:
+ pass # Registry not available (not running in proxy context)
+
+ return api_base, api_key, headers
+
+ def get_supported_openai_params(self, model: str) -> List[str]:
+ """Return list of supported OpenAI parameters"""
+ return [
+ "stream",
+ "temperature",
+ "max_tokens",
+ "top_p",
+ ]
+
+ def map_openai_params(
+ self,
+ non_default_params: dict,
+ optional_params: dict,
+ model: str,
+ drop_params: bool,
+ ) -> dict:
+ """
+ Map OpenAI parameters to A2A parameters.
+
+ For A2A protocol, we need to map the stream parameter so
+ transform_request can determine which JSON-RPC method to use.
+ """
+ # Map stream parameter
+ for param, value in non_default_params.items():
+ if param == "stream" and value is True:
+ optional_params["stream"] = value
+
+ return optional_params
+
+ def validate_environment(
+ self,
+ headers: dict,
+ model: str,
+ messages: List[AllMessageValues],
+ optional_params: dict,
+ litellm_params: dict,
+ api_key: Optional[str] = None,
+ api_base: Optional[str] = None,
+ ) -> dict:
+ """
+ Validate environment and set headers for A2A requests.
+
+ Args:
+ headers: Request headers dict
+ model: Model name
+ messages: Messages list
+ optional_params: Optional parameters
+ litellm_params: LiteLLM parameters
+ api_key: API key (optional for A2A)
+ api_base: API base URL
+
+ Returns:
+ Updated headers dict
+ """
+ # Ensure Content-Type is set to application/json for JSON-RPC 2.0
+ if "content-type" not in headers and "Content-Type" not in headers:
+ headers["Content-Type"] = "application/json"
+
+ # Add Authorization header if API key is provided
+ if api_key is not None:
+ headers["Authorization"] = f"Bearer {api_key}"
+
+ return headers
+
+ def get_complete_url(
+ self,
+ api_base: Optional[str],
+ api_key: Optional[str],
+ model: str,
+ optional_params: dict,
+ litellm_params: dict,
+ stream: Optional[bool] = None,
+ ) -> str:
+ """
+ Get the complete A2A agent endpoint URL.
+
+ A2A agents use JSON-RPC 2.0 at the base URL, not specific paths.
+ The method (message/send or message/stream) is specified in the
+ JSON-RPC request body, not in the URL.
+
+ Args:
+ api_base: Base URL of the A2A agent (e.g., "http://0.0.0.0:9999")
+ api_key: API key (not used for URL construction)
+ model: Model name (not used for A2A, agent determined by api_base)
+ optional_params: Optional parameters
+ litellm_params: LiteLLM parameters
+ stream: Whether this is a streaming request (affects JSON-RPC method)
+
+ Returns:
+ Complete URL for the A2A endpoint (base URL)
+ """
+ if api_base is None:
+ raise ValueError("api_base is required for A2A provider")
+
+ # A2A uses JSON-RPC 2.0 at the base URL
+ # Remove trailing slash for consistency
+ return api_base.rstrip("/")
+
+ def transform_request(
+ self,
+ model: str,
+ messages: List[AllMessageValues],
+ optional_params: dict,
+ litellm_params: dict,
+ headers: dict,
+ ) -> dict:
+ """
+ Transform OpenAI request to A2A JSON-RPC 2.0 format.
+
+ Args:
+ model: Model name
+ messages: List of OpenAI messages
+ optional_params: Optional parameters
+ litellm_params: LiteLLM parameters
+ headers: Request headers
+
+ Returns:
+ A2A JSON-RPC 2.0 request dict
+ """
+ # Generate request ID
+ request_id = str(uuid.uuid4())
+
+ if not messages:
+ raise ValueError("At least one message is required for A2A completion")
+
+ # Convert all messages to maintain conversation history
+ # Use helper to format conversation with role prefixes
+ full_context = convert_messages_to_prompt(messages)
+
+ # Create single A2A message with full conversation context
+ a2a_message = {
+ "role": "user",
+ "parts": [{"kind": "text", "text": full_context}],
+ "messageId": str(uuid.uuid4()),
+ }
+
+ # Build JSON-RPC 2.0 request
+ # For A2A protocol, the method is "message/send" for non-streaming
+ # and "message/stream" for streaming
+ stream = optional_params.get("stream", False)
+ method = "message/stream" if stream else "message/send"
+
+ request_data = {
+ "jsonrpc": "2.0",
+ "id": request_id,
+ "method": method,
+ "params": {
+ "message": a2a_message
+ }
+ }
+
+ return request_data
+
+ def transform_response(
+ self,
+ model: str,
+ raw_response: httpx.Response,
+ model_response: ModelResponse,
+ logging_obj: Any,
+ request_data: dict,
+ messages: List[AllMessageValues],
+ optional_params: dict,
+ litellm_params: dict,
+ encoding: Any,
+ api_key: Optional[str] = None,
+ json_mode: Optional[bool] = None,
+ ) -> ModelResponse:
+ """
+ Transform A2A JSON-RPC 2.0 response to OpenAI format.
+
+ Args:
+ model: Model name
+ raw_response: HTTP response from A2A agent
+ model_response: Model response object to populate
+ logging_obj: Logging object
+ request_data: Original request data
+ messages: Original messages
+ optional_params: Optional parameters
+ litellm_params: LiteLLM parameters
+ encoding: Encoding object
+ api_key: API key
+ json_mode: JSON mode flag
+
+ Returns:
+ Populated ModelResponse object
+ """
+ try:
+ response_json = raw_response.json()
+ except Exception as e:
+ raise A2AError(
+ status_code=raw_response.status_code,
+ message=f"Failed to parse A2A response: {str(e)}",
+ headers=dict(raw_response.headers),
+ )
+
+ # Check for JSON-RPC error
+ if "error" in response_json:
+ error = response_json["error"]
+ raise A2AError(
+ status_code=raw_response.status_code,
+ message=f"A2A error: {error.get('message', 'Unknown error')}",
+ headers=dict(raw_response.headers),
+ )
+
+ # Extract text from A2A response
+ text = extract_text_from_a2a_response(response_json)
+
+ # Populate model response
+ model_response.choices = [
+ Choices(
+ finish_reason="stop",
+ index=0,
+ message=Message(
+ content=text,
+ role="assistant",
+ ),
+ )
+ ]
+
+ # Set model
+ model_response.model = model
+
+ # Set ID from response
+ model_response.id = response_json.get("id", str(uuid.uuid4()))
+
+ return model_response
+
+ def get_model_response_iterator(
+ self,
+ streaming_response: Union[Iterator, Any],
+ sync_stream: bool,
+ json_mode: Optional[bool] = False,
+ ) -> BaseModelResponseIterator:
+ """
+ Get streaming iterator for A2A responses.
+
+ Args:
+ streaming_response: Streaming response iterator
+ sync_stream: Whether this is a sync stream
+ json_mode: JSON mode flag
+
+ Returns:
+ A2A streaming iterator
+ """
+ return A2AModelResponseIterator(
+ streaming_response=streaming_response,
+ sync_stream=sync_stream,
+ json_mode=json_mode,
+ )
+
+ def _openai_message_to_a2a_message(self, message: Dict[str, Any]) -> Dict[str, Any]:
+ """
+ Convert OpenAI message to A2A message format.
+
+ Args:
+ message: OpenAI message dict
+
+ Returns:
+ A2A message dict
+ """
+ content = message.get("content", "")
+ role = message.get("role", "user")
+
+ return {
+ "role": role,
+ "parts": [{"kind": "text", "text": str(content)}],
+ "messageId": str(uuid.uuid4()),
+ }
+
+ def get_error_class(
+ self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
+ ) -> BaseLLMException:
+ """Return appropriate error class for A2A errors"""
+ # Convert headers to dict if needed
+ headers_dict = dict(headers) if isinstance(headers, httpx.Headers) else headers
+ return A2AError(
+ status_code=status_code,
+ message=error_message,
+ headers=headers_dict,
+ )
diff --git a/litellm/llms/a2a/common_utils.py b/litellm/llms/a2a/common_utils.py
new file mode 100644
index 00000000000..116e1205409
--- /dev/null
+++ b/litellm/llms/a2a/common_utils.py
@@ -0,0 +1,152 @@
+"""
+Common utilities for A2A (Agent-to-Agent) Protocol
+"""
+from typing import Any, Dict, List
+
+from pydantic import BaseModel
+
+from litellm.litellm_core_utils.prompt_templates.common_utils import (
+ convert_content_list_to_str,
+)
+from litellm.llms.base_llm.chat.transformation import BaseLLMException
+from litellm.types.llms.openai import AllMessageValues
+
+
+class A2AError(BaseLLMException):
+ """Base exception for A2A protocol errors"""
+
+ def __init__(
+ self,
+ status_code: int,
+ message: str,
+ headers: Dict[str, Any] = {},
+ ):
+ super().__init__(
+ status_code=status_code,
+ message=message,
+ headers=headers,
+ )
+
+
+def convert_messages_to_prompt(messages: List[AllMessageValues]) -> str:
+ """
+ Convert OpenAI messages to a single prompt string for A2A agent.
+
+ Formats each message as "{role}: {content}" and joins with newlines
+ to preserve conversation history. Handles both string and list content.
+
+ Args:
+ messages: List of OpenAI-format messages
+
+ Returns:
+ Formatted prompt string with full conversation context
+ """
+ conversation_parts = []
+ for msg in messages:
+ # Use LiteLLM's helper to extract text from content (handles both str and list)
+ content_text = convert_content_list_to_str(message=msg)
+
+ # Get role
+ if isinstance(msg, BaseModel):
+ role = msg.model_dump().get("role", "user")
+ elif isinstance(msg, dict):
+ role = msg.get("role", "user")
+ else:
+ role = dict(msg).get("role", "user") # type: ignore
+
+ if content_text:
+ conversation_parts.append(f"{role}: {content_text}")
+
+ return "\n".join(conversation_parts)
+
+
+def extract_text_from_a2a_message(
+ message: Dict[str, Any], depth: int = 0, max_depth: int = 10
+) -> str:
+ """
+ Extract text content from A2A message parts.
+
+ Args:
+ message: A2A message dict with 'parts' containing text parts
+ depth: Current recursion depth (internal use)
+ max_depth: Maximum recursion depth to prevent infinite loops
+
+ Returns:
+ Concatenated text from all text parts
+ """
+ if message is None or depth >= max_depth:
+ return ""
+
+ parts = message.get("parts", [])
+ text_parts: List[str] = []
+
+ for part in parts:
+ if part.get("kind") == "text":
+ text_parts.append(part.get("text", ""))
+ # Handle nested parts if they exist
+ elif "parts" in part:
+ nested_text = extract_text_from_a2a_message(part, depth + 1, max_depth)
+ if nested_text:
+ text_parts.append(nested_text)
+
+ return " ".join(text_parts)
+
+
+def extract_text_from_a2a_response(
+ response_dict: Dict[str, Any], max_depth: int = 10
+) -> str:
+ """
+ Extract text content from A2A response result.
+
+ Args:
+ response_dict: A2A response dict with 'result' containing message
+ max_depth: Maximum recursion depth to prevent infinite loops
+
+ Returns:
+ Text from response message parts
+ """
+ result = response_dict.get("result", {})
+ if not isinstance(result, dict):
+ return ""
+
+ # A2A response can have different formats:
+ # 1. Direct message: {"result": {"kind": "message", "parts": [...]}}
+ # 2. Nested message: {"result": {"message": {"parts": [...]}}}
+ # 3. Task with artifacts: {"result": {"kind": "task", "artifacts": [{"parts": [...]}]}}
+ # 4. Task with status message: {"result": {"kind": "task", "status": {"message": {"parts": [...]}}}}
+ # 5. Streaming artifact-update: {"result": {"kind": "artifact-update", "artifact": {"parts": [...]}}}
+
+ # Check if result itself has parts (direct message)
+ if "parts" in result:
+ return extract_text_from_a2a_message(result, depth=0, max_depth=max_depth)
+
+ # Check for nested message
+ message = result.get("message")
+ if message:
+ return extract_text_from_a2a_message(message, depth=0, max_depth=max_depth)
+
+ # Check for streaming artifact-update (singular artifact)
+ artifact = result.get("artifact")
+ if artifact and isinstance(artifact, dict):
+ return extract_text_from_a2a_message(
+ artifact, depth=0, max_depth=max_depth
+ )
+
+ # Check for task status message (common in Gemini A2A agents)
+ status = result.get("status", {})
+ if isinstance(status, dict):
+ status_message = status.get("message")
+ if status_message:
+ return extract_text_from_a2a_message(
+ status_message, depth=0, max_depth=max_depth
+ )
+
+ # Handle task result with artifacts (plural, array)
+ artifacts = result.get("artifacts", [])
+ if artifacts and len(artifacts) > 0:
+ first_artifact = artifacts[0]
+ return extract_text_from_a2a_message(
+ first_artifact, depth=0, max_depth=max_depth
+ )
+
+ return ""
diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py
index 71d74121a30..a14e7d118e8 100644
--- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py
+++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py
@@ -34,6 +34,7 @@ from litellm.types.llms.openai import (
)
from litellm.types.utils import (
ChatCompletionMessageToolCall,
+ Choices,
GenericGuardrailAPIInputs,
ModelResponse,
)
@@ -74,9 +75,10 @@ class AnthropicMessagesHandler(BaseTranslation):
if messages is None:
return data
- chat_completion_compatible_request = (
+ chat_completion_compatible_request, tool_name_mapping = (
LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
- anthropic_message_request=cast(AnthropicMessagesRequest, data)
+ # Use a shallow copy to avoid mutating request data (pop on litellm_metadata).
+ anthropic_message_request=cast(AnthropicMessagesRequest, data.copy())
)
)
@@ -84,9 +86,9 @@ class AnthropicMessagesHandler(BaseTranslation):
texts_to_check: List[str] = []
images_to_check: List[str] = []
- tools_to_check: List[ChatCompletionToolParam] = (
- chat_completion_compatible_request.get("tools", [])
- )
+ tools_to_check: List[
+ ChatCompletionToolParam
+ ] = chat_completion_compatible_request.get("tools", [])
task_mappings: List[Tuple[int, Optional[int]]] = []
# Track (message_index, content_index) for each text
# content_index is None for string content, int for list content
@@ -282,7 +284,10 @@ class AnthropicMessagesHandler(BaseTranslation):
if hasattr(content_block, "model_dump"):
block_dict = content_block.model_dump()
else:
- block_dict = {"type": block_type, "text": getattr(content_block, "text", None)}
+ block_dict = {
+ "type": block_type,
+ "text": getattr(content_block, "text", None),
+ }
else:
continue
@@ -358,30 +363,40 @@ class AnthropicMessagesHandler(BaseTranslation):
"""
has_ended = self._check_streaming_has_ended(responses_so_far)
if has_ended:
-
# build the model response from the responses_so_far
- model_response = cast(
- ModelResponse,
- AnthropicPassthroughLoggingHandler._build_complete_streaming_response(
- all_chunks=responses_so_far,
- litellm_logging_obj=cast("LiteLLMLoggingObj", litellm_logging_obj),
- model="",
- ),
+ built_response = AnthropicPassthroughLoggingHandler._build_complete_streaming_response(
+ all_chunks=responses_so_far,
+ litellm_logging_obj=cast("LiteLLMLoggingObj", litellm_logging_obj),
+ model="",
)
- tool_calls_list = cast(Optional[List[ChatCompletionMessageToolCall]], model_response.choices[0].message.tool_calls) # type: ignore
- string_so_far = model_response.choices[0].message.content # type: ignore
- guardrail_inputs = GenericGuardrailAPIInputs()
- if string_so_far:
- guardrail_inputs["texts"] = [string_so_far]
- if tool_calls_list:
- guardrail_inputs["tool_calls"] = tool_calls_list
- _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid
- inputs=guardrail_inputs,
- request_data={},
- input_type="response",
- logging_obj=litellm_logging_obj,
- )
+ # Check if model_response is valid and has choices before accessing
+ if (
+ built_response is not None
+ and hasattr(built_response, "choices")
+ and built_response.choices
+ ):
+ model_response = cast(ModelResponse, built_response)
+ first_choice = cast(Choices, model_response.choices[0])
+ tool_calls_list = cast(
+ Optional[List[ChatCompletionMessageToolCall]],
+ first_choice.message.tool_calls,
+ )
+ string_so_far = first_choice.message.content
+ guardrail_inputs = GenericGuardrailAPIInputs()
+ if string_so_far:
+ guardrail_inputs["texts"] = [string_so_far]
+ if tool_calls_list:
+ guardrail_inputs["tool_calls"] = tool_calls_list
+
+ _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid
+ inputs=guardrail_inputs,
+ request_data={},
+ input_type="response",
+ logging_obj=litellm_logging_obj,
+ )
+ else:
+ verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices")
return responses_so_far
string_so_far = self.get_streaming_string_so_far(responses_so_far)
@@ -648,7 +663,10 @@ class AnthropicMessagesHandler(BaseTranslation):
if isinstance(content_block, dict):
if content_block.get("type") == "text":
cast(Dict[str, Any], content_block)["text"] = guardrail_response
- elif hasattr(content_block, "type") and getattr(content_block, "type", None) == "text":
+ elif (
+ hasattr(content_block, "type")
+ and getattr(content_block, "type", None) == "text"
+ ):
# Update Pydantic object's text attribute
if hasattr(content_block, "text"):
content_block.text = guardrail_response
diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py
index 6a9aafd076b..485e95d6489 100644
--- a/litellm/llms/anthropic/chat/handler.py
+++ b/litellm/llms/anthropic/chat/handler.py
@@ -512,6 +512,9 @@ class ModelResponseIterator:
# Accumulate web_search_tool_result blocks for multi-turn reconstruction
# See: https://github.com/BerriAI/litellm/issues/17737
self.web_search_results: List[Dict[str, Any]] = []
+
+ # Accumulate compaction blocks for multi-turn reconstruction
+ self.compaction_blocks: List[Dict[str, Any]] = []
def check_empty_tool_call_args(self) -> bool:
"""
@@ -592,6 +595,12 @@ class ModelResponseIterator:
)
]
provider_specific_fields["thinking_blocks"] = thinking_blocks
+ elif "content" in content_block["delta"] and content_block["delta"].get("type") == "compaction_delta":
+ # Handle compaction delta
+ provider_specific_fields["compaction_delta"] = {
+ "type": "compaction_delta",
+ "content": content_block["delta"]["content"]
+ }
return text, tool_use, thinking_blocks, provider_specific_fields
@@ -721,6 +730,20 @@ class ModelResponseIterator:
provider_specific_fields=provider_specific_fields,
)
+ elif content_block_start["content_block"]["type"] == "compaction":
+ # Handle compaction blocks
+ # The full content comes in content_block_start
+ self.compaction_blocks.append(
+ content_block_start["content_block"]
+ )
+ provider_specific_fields["compaction_blocks"] = (
+ self.compaction_blocks
+ )
+ provider_specific_fields["compaction_start"] = {
+ "type": "compaction",
+ "content": content_block_start["content_block"].get("content", "")
+ }
+
elif content_block_start["content_block"]["type"].endswith("_tool_result"):
# Handle all tool result types (web_search, bash_code_execution, text_editor, etc.)
content_type = content_block_start["content_block"]["type"]
diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py
index 1b61b533275..02b8d952445 100644
--- a/litellm/llms/anthropic/chat/transformation.py
+++ b/litellm/llms/anthropic/chat/transformation.py
@@ -170,9 +170,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
tool_call["caller"] = cast(Dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item]
return tool_call
- def _is_claude_opus_4_5(self, model: str) -> bool:
+ @staticmethod
+ def _is_claude_opus_4_6(model: str) -> bool:
"""Check if the model is Claude Opus 4.5."""
- return "opus-4-5" in model.lower() or "opus_4_5" in model.lower()
+ return "opus-4-6" in model.lower() or "opus_4_6" in model.lower()
def get_supported_openai_params(self, model: str):
params = [
@@ -659,32 +660,38 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def _map_reasoning_effort(
- reasoning_effort: Optional[Union[REASONING_EFFORT, str]],
+ reasoning_effort: Optional[Union[REASONING_EFFORT, str]],
+ model: str,
) -> Optional[AnthropicThinkingParam]:
- if reasoning_effort is None:
- return None
- elif reasoning_effort == "low":
+ if AnthropicConfig._is_claude_opus_4_6(model):
return AnthropicThinkingParam(
- type="enabled",
- budget_tokens=DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
- )
- elif reasoning_effort == "medium":
- return AnthropicThinkingParam(
- type="enabled",
- budget_tokens=DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
- )
- elif reasoning_effort == "high":
- return AnthropicThinkingParam(
- type="enabled",
- budget_tokens=DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
- )
- elif reasoning_effort == "minimal":
- return AnthropicThinkingParam(
- type="enabled",
- budget_tokens=DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET,
+ type="adaptive",
)
else:
- raise ValueError(f"Unmapped reasoning effort: {reasoning_effort}")
+ if reasoning_effort is None:
+ return None
+ elif reasoning_effort == "low":
+ return AnthropicThinkingParam(
+ type="enabled",
+ budget_tokens=DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
+ )
+ elif reasoning_effort == "medium":
+ return AnthropicThinkingParam(
+ type="enabled",
+ budget_tokens=DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
+ )
+ elif reasoning_effort == "high":
+ return AnthropicThinkingParam(
+ type="enabled",
+ budget_tokens=DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
+ )
+ elif reasoning_effort == "minimal":
+ return AnthropicThinkingParam(
+ type="enabled",
+ budget_tokens=DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET,
+ )
+ else:
+ raise ValueError(f"Unmapped reasoning effort: {reasoning_effort}")
def _extract_json_schema_from_response_format(
self, value: Optional[dict]
@@ -860,13 +867,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if param == "thinking":
optional_params["thinking"] = value
elif param == "reasoning_effort" and isinstance(value, str):
- # For Claude Opus 4.5, map reasoning_effort to output_config
- if self._is_claude_opus_4_5(model):
- optional_params["output_config"] = {"effort": value}
-
- # For other models, map to thinking parameter
optional_params["thinking"] = AnthropicConfig._map_reasoning_effort(
- value
+ reasoning_effort=value, model=model
)
elif param == "web_search_options" and isinstance(value, dict):
hosted_web_search_tool = self.map_web_search_tool(
@@ -877,6 +879,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
elif param == "extra_headers":
optional_params["extra_headers"] = value
+ elif param == "context_management" and isinstance(value, dict):
+ # Pass through Anthropic-specific context_management parameter
+ optional_params["context_management"] = value
## handle thinking tokens
self.update_optional_params_with_thinking_tokens(
@@ -1026,9 +1031,37 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if beta_value not in existing_values:
headers["anthropic-beta"] = f"{existing_beta}, {beta_value}"
- def _ensure_context_management_beta_header(self, headers: dict) -> None:
- beta_value = ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value
- self._ensure_beta_header(headers, beta_value)
+ def _ensure_context_management_beta_header(
+ self, headers: dict, context_management: dict
+ ) -> None:
+ """
+ Add appropriate beta headers based on context_management edits.
+ - If any edit has type "compact_20260112", add compact-2026-01-12 header
+ - For all other edits, add context-management-2025-06-27 header
+ """
+ edits = context_management.get("edits", [])
+
+ has_compact = False
+ has_other = False
+
+ for edit in edits:
+ edit_type = edit.get("type", "")
+ if edit_type == "compact_20260112":
+ has_compact = True
+ else:
+ has_other = True
+
+ # Add compact header if any compact edits exist
+ if has_compact:
+ self._ensure_beta_header(
+ headers, ANTHROPIC_BETA_HEADER_VALUES.COMPACT_2026_01_12.value
+ )
+
+ # Add context management header if any other edits exist
+ if has_other:
+ self._ensure_beta_header(
+ headers, ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value
+ )
def update_headers_with_optional_anthropic_beta(
self, headers: dict, optional_params: dict
@@ -1056,7 +1089,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
headers, ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value
)
if optional_params.get("context_management") is not None:
- self._ensure_context_management_beta_header(headers)
+ self._ensure_context_management_beta_header(
+ headers, optional_params["context_management"]
+ )
if optional_params.get("output_format") is not None:
self._ensure_beta_header(
headers, ANTHROPIC_BETA_HEADER_VALUES.STRUCTURED_OUTPUT_2025_09_25.value
@@ -1225,6 +1260,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
List[ChatCompletionToolCallChunk],
Optional[List[Any]],
Optional[List[Any]],
+ Optional[List[Any]],
]:
text_content = ""
citations: Optional[List[Any]] = None
@@ -1237,6 +1273,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
tool_calls: List[ChatCompletionToolCallChunk] = []
web_search_results: Optional[List[Any]] = None
tool_results: Optional[List[Any]] = None
+ compaction_blocks: Optional[List[Any]] = None
for idx, content in enumerate(completion_response["content"]):
if content["type"] == "text":
text_content += content["text"]
@@ -1278,6 +1315,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
thinking_blocks.append(
cast(ChatCompletionRedactedThinkingBlock, content)
)
+
+ ## COMPACTION
+ elif content["type"] == "compaction":
+ if compaction_blocks is None:
+ compaction_blocks = []
+ compaction_blocks.append(content)
## CITATIONS
if content.get("citations") is not None:
@@ -1299,7 +1342,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if thinking_content is not None:
reasoning_content += thinking_content
- return text_content, citations, thinking_blocks, reasoning_content, tool_calls, web_search_results, tool_results
+ return text_content, citations, thinking_blocks, reasoning_content, tool_calls, web_search_results, tool_results, compaction_blocks
def calculate_usage(
self,
@@ -1316,6 +1359,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
web_search_requests: Optional[int] = None
tool_search_requests: Optional[int] = None
+ inference_geo: Optional[str] = None
+ if "inference_geo" in _usage and _usage["inference_geo"] is not None:
+ inference_geo = _usage["inference_geo"]
+
if (
"cache_creation_input_tokens" in _usage
and _usage["cache_creation_input_tokens"] is not None
@@ -1399,6 +1446,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if (web_search_requests is not None or tool_search_requests is not None)
else None
),
+ inference_geo=inference_geo,
)
return usage
@@ -1442,6 +1490,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
tool_calls,
web_search_results,
tool_results,
+ compaction_blocks,
) = self.extract_response_content(completion_response=completion_response)
if (
@@ -1469,6 +1518,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
provider_specific_fields["tool_results"] = tool_results
if container is not None:
provider_specific_fields["container"] = container
+ if compaction_blocks is not None:
+ provider_specific_fields["compaction_blocks"] = compaction_blocks
_message = litellm.Message(
tool_calls=tool_calls,
@@ -1477,6 +1528,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
thinking_blocks=thinking_blocks,
reasoning_content=reasoning_content,
)
+ _message.provider_specific_fields = provider_specific_fields
## HANDLE JSON MODE - anthropic returns single function call
json_mode_message = self._transform_response_for_json_mode(
@@ -1507,18 +1559,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
model_response.created = int(time.time())
model_response.model = completion_response["model"]
- context_management_response = completion_response.get("context_management")
- if context_management_response is not None:
- _hidden_params["context_management"] = context_management_response
- try:
- model_response.__dict__["context_management"] = (
- context_management_response
- )
- except Exception:
- pass
-
model_response._hidden_params = _hidden_params
-
return model_response
def get_prefix_prompt(self, messages: List[AllMessageValues]) -> Optional[str]:
diff --git a/litellm/llms/anthropic/cost_calculation.py b/litellm/llms/anthropic/cost_calculation.py
index 8f34eb00ce5..11b61cc92f0 100644
--- a/litellm/llms/anthropic/cost_calculation.py
+++ b/litellm/llms/anthropic/cost_calculation.py
@@ -22,10 +22,17 @@ def cost_per_token(model: str, usage: "Usage") -> Tuple[float, float]:
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
"""
- return generic_cost_per_token(
- model=model, usage=usage, custom_llm_provider="anthropic"
+ # If usage has inference_geo, prepend it as prefix to model name
+ if hasattr(usage, "inference_geo") and usage.inference_geo and usage.inference_geo.lower() not in ["global", "not_available"]:
+ model_with_geo_prefix = f"{usage.inference_geo}/{model}"
+ else:
+ model_with_geo_prefix = model
+ prompt_cost, completion_cost = generic_cost_per_token(
+ model=model_with_geo_prefix, usage=usage, custom_llm_provider="anthropic"
)
+ return prompt_cost, completion_cost
+
def get_cost_for_anthropic_web_search(
model_info: Optional["ModelInfo"] = None,
diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py b/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py
index 8fa7bb7e65e..a17eba75b3b 100644
--- a/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py
+++ b/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py
@@ -6,6 +6,7 @@ from typing import (
Dict,
List,
Optional,
+ Tuple,
Union,
cast,
)
@@ -47,8 +48,14 @@ class LiteLLMMessagesToCompletionTransformationHandler:
top_p: Optional[float] = None,
output_format: Optional[Dict] = None,
extra_kwargs: Optional[Dict[str, Any]] = None,
- ) -> Dict[str, Any]:
- """Prepare kwargs for litellm.completion/acompletion"""
+ ) -> Tuple[Dict[str, Any], Dict[str, str]]:
+ """Prepare kwargs for litellm.completion/acompletion.
+
+ Returns:
+ Tuple of (completion_kwargs, tool_name_mapping)
+ - tool_name_mapping maps truncated tool names back to original names
+ for tools that exceeded OpenAI's 64-char limit
+ """
from litellm.litellm_core_utils.litellm_logging import (
Logging as LiteLLMLoggingObject,
)
@@ -80,7 +87,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
if output_format:
request_data["output_format"] = output_format
- openai_request = ANTHROPIC_ADAPTER.translate_completion_input_params(
+ openai_request, tool_name_mapping = ANTHROPIC_ADAPTER.translate_completion_input_params_with_tool_mapping(
request_data
)
@@ -116,7 +123,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
):
completion_kwargs[key] = value
- return completion_kwargs
+ return completion_kwargs, tool_name_mapping
@staticmethod
async def async_anthropic_messages_handler(
@@ -137,7 +144,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
**kwargs,
) -> Union[AnthropicMessagesResponse, AsyncIterator]:
"""Handle non-Anthropic models asynchronously using the adapter"""
- completion_kwargs = (
+ completion_kwargs, tool_name_mapping = (
LiteLLMMessagesToCompletionTransformationHandler._prepare_completion_kwargs(
max_tokens=max_tokens,
messages=messages,
@@ -164,6 +171,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
ANTHROPIC_ADAPTER.translate_completion_output_params_streaming(
completion_response,
model=model,
+ tool_name_mapping=tool_name_mapping,
)
)
if transformed_stream is not None:
@@ -172,7 +180,8 @@ class LiteLLMMessagesToCompletionTransformationHandler:
else:
anthropic_response = (
ANTHROPIC_ADAPTER.translate_completion_output_params(
- cast(ModelResponse, completion_response)
+ cast(ModelResponse, completion_response),
+ tool_name_mapping=tool_name_mapping,
)
)
if anthropic_response is not None:
@@ -222,7 +231,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
**kwargs,
)
- completion_kwargs = (
+ completion_kwargs, tool_name_mapping = (
LiteLLMMessagesToCompletionTransformationHandler._prepare_completion_kwargs(
max_tokens=max_tokens,
messages=messages,
@@ -249,6 +258,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
ANTHROPIC_ADAPTER.translate_completion_output_params_streaming(
completion_response,
model=model,
+ tool_name_mapping=tool_name_mapping,
)
)
if transformed_stream is not None:
@@ -257,7 +267,8 @@ class LiteLLMMessagesToCompletionTransformationHandler:
else:
anthropic_response = (
ANTHROPIC_ADAPTER.translate_completion_output_params(
- cast(ModelResponse, completion_response)
+ cast(ModelResponse, completion_response),
+ tool_name_mapping=tool_name_mapping,
)
)
if anthropic_response is not None:
diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py
index 24524233ddf..a86820f82e8 100644
--- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py
+++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py
@@ -3,7 +3,7 @@
import json
import traceback
from collections import deque
-from typing import TYPE_CHECKING, Any, AsyncIterator, Iterator, Literal, Optional
+from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, Iterator, Literal, Optional
from litellm import verbose_logger
from litellm._uuid import uuid
@@ -44,9 +44,16 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
pending_new_content_block: bool = False
chunk_queue: deque = deque() # Queue for buffering multiple chunks
- def __init__(self, completion_stream: Any, model: str):
+ def __init__(
+ self,
+ completion_stream: Any,
+ model: str,
+ tool_name_mapping: Optional[Dict[str, str]] = None,
+ ):
super().__init__(completion_stream)
self.model = model
+ # Mapping of truncated tool names to original names (for OpenAI's 64-char limit)
+ self.tool_name_mapping = tool_name_mapping or {}
def _create_initial_usage_delta(self) -> UsageDelta:
"""
@@ -401,6 +408,19 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
choices=chunk.choices # type: ignore
)
+ # Restore original tool name if it was truncated for OpenAI's 64-char limit
+ if block_type == "tool_use":
+ # Type narrowing: content_block_start is ToolUseBlock when block_type is "tool_use"
+ from typing import cast
+ from litellm.types.llms.anthropic import ToolUseBlock
+
+ tool_block = cast(ToolUseBlock, content_block_start)
+
+ if tool_block.get("name"):
+ truncated_name = tool_block["name"]
+ original_name = self.tool_name_mapping.get(truncated_name, truncated_name)
+ tool_block["name"] = original_name
+
if block_type != self.current_content_block_type:
self.current_content_block_type = block_type
self.current_content_block_start = content_block_start
@@ -408,9 +428,14 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
# For parallel tool calls, we'll necessarily have a new content block
# if we get a function name since it signals a new tool call
- if block_type == "tool_use" and content_block_start.get("name"):
- self.current_content_block_type = block_type
- self.current_content_block_start = content_block_start
- return True
+ if block_type == "tool_use":
+ from typing import cast
+ from litellm.types.llms.anthropic import ToolUseBlock
+
+ tool_block = cast(ToolUseBlock, content_block_start)
+ if tool_block.get("name"):
+ self.current_content_block_type = block_type
+ self.current_content_block_start = content_block_start
+ return True
return False
diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py
index 5ba0754b744..169b138a5f7 100644
--- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py
+++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py
@@ -1,3 +1,4 @@
+import hashlib
import json
from typing import (
TYPE_CHECKING,
@@ -12,6 +13,54 @@ from typing import (
cast,
)
+# OpenAI has a 64-character limit for function/tool names
+# Anthropic does not have this limit, so we need to truncate long names
+OPENAI_MAX_TOOL_NAME_LENGTH = 64
+TOOL_NAME_HASH_LENGTH = 8
+TOOL_NAME_PREFIX_LENGTH = OPENAI_MAX_TOOL_NAME_LENGTH - TOOL_NAME_HASH_LENGTH - 1 # 55
+
+
+def truncate_tool_name(name: str) -> str:
+ """
+ Truncate tool names that exceed OpenAI's 64-character limit.
+
+ Uses format: {55-char-prefix}_{8-char-hash} to avoid collisions
+ when multiple tools have similar long names.
+
+ Args:
+ name: The original tool name
+
+ Returns:
+ The original name if <= 64 chars, otherwise truncated with hash
+ """
+ if len(name) <= OPENAI_MAX_TOOL_NAME_LENGTH:
+ return name
+
+ # Create deterministic hash from full name to avoid collisions
+ name_hash = hashlib.sha256(name.encode()).hexdigest()[:TOOL_NAME_HASH_LENGTH]
+ return f"{name[:TOOL_NAME_PREFIX_LENGTH]}_{name_hash}"
+
+
+def create_tool_name_mapping(
+ tools: List[Dict[str, Any]],
+) -> Dict[str, str]:
+ """
+ Create a mapping of truncated tool names to original names.
+
+ Args:
+ tools: List of tool definitions with 'name' field
+
+ Returns:
+ Dict mapping truncated names to original names (only for truncated tools)
+ """
+ mapping: Dict[str, str] = {}
+ for tool in tools:
+ original_name = tool.get("name", "")
+ truncated_name = truncate_tool_name(original_name)
+ if truncated_name != original_name:
+ mapping[truncated_name] = original_name
+ return mapping
+
from openai.types.chat.chat_completion_chunk import Choice as OpenAIStreamingChoice
from litellm.litellm_core_utils.prompt_templates.common_utils import (
@@ -77,8 +126,29 @@ class AnthropicAdapter:
self, kwargs
) -> Optional[ChatCompletionRequest]:
"""
+ Translate Anthropic request params to OpenAI format.
+
- translate params, where needed
- pass rest, as is
+
+ Note: Use translate_completion_input_params_with_tool_mapping() if you need
+ the tool name mapping for restoring original names in responses.
+ """
+ result, _ = self.translate_completion_input_params_with_tool_mapping(kwargs)
+ return result
+
+ def translate_completion_input_params_with_tool_mapping(
+ self, kwargs
+ ) -> Tuple[Optional[ChatCompletionRequest], Dict[str, str]]:
+ """
+ Translate Anthropic request params to OpenAI format, returning tool name mapping.
+
+ This method handles truncation of tool names that exceed OpenAI's 64-character
+ limit. The mapping allows restoring original names when translating responses.
+
+ Returns:
+ Tuple of (openai_request, tool_name_mapping)
+ - tool_name_mapping maps truncated tool names back to original names
"""
#########################################################
@@ -102,26 +172,51 @@ class AnthropicAdapter:
model=model, messages=messages, **kwargs
)
- translated_body = (
+ translated_body, tool_name_mapping = (
LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
anthropic_message_request=request_body
)
)
- return translated_body
+ return translated_body, tool_name_mapping
def translate_completion_output_params(
- self, response: ModelResponse
+ self,
+ response: ModelResponse,
+ tool_name_mapping: Optional[Dict[str, str]] = None,
) -> Optional[AnthropicMessagesResponse]:
+ """
+ Translate OpenAI response to Anthropic format.
+
+ Args:
+ response: The OpenAI ModelResponse
+ tool_name_mapping: Optional mapping of truncated tool names to original names.
+ Used to restore original names for tools that exceeded
+ OpenAI's 64-char limit.
+ """
return LiteLLMAnthropicMessagesAdapter().translate_openai_response_to_anthropic(
- response=response
+ response=response,
+ tool_name_mapping=tool_name_mapping,
)
def translate_completion_output_params_streaming(
- self, completion_stream: Any, model: str
+ self,
+ completion_stream: Any,
+ model: str,
+ tool_name_mapping: Optional[Dict[str, str]] = None,
) -> Union[AsyncIterator[bytes], None]:
+ """
+ Translate OpenAI streaming response to Anthropic format.
+
+ Args:
+ completion_stream: The OpenAI streaming response
+ model: The model name
+ tool_name_mapping: Optional mapping of truncated tool names to original names.
+ """
anthropic_wrapper = AnthropicStreamWrapper(
- completion_stream=completion_stream, model=model
+ completion_stream=completion_stream,
+ model=model,
+ tool_name_mapping=tool_name_mapping,
)
# Return the SSE-wrapped version for proper event formatting
return anthropic_wrapper.async_anthropic_sse_wrapper()
@@ -417,8 +512,10 @@ class LiteLLMAnthropicMessagesAdapter:
has_cache_control_in_text = True
assistant_content_list.append(text_block)
elif content.get("type") == "tool_use":
+ # Truncate tool name for OpenAI's 64-char limit
+ tool_name = truncate_tool_name(content.get("name", ""))
function_chunk: ChatCompletionToolCallFunctionChunk = {
- "name": content.get("name", ""),
+ "name": tool_name,
"arguments": json.dumps(content.get("input", {})),
}
signature = (
@@ -587,8 +684,11 @@ class LiteLLMAnthropicMessagesAdapter:
elif tool_choice["type"] == "auto":
return "auto"
elif tool_choice["type"] == "tool":
+ # Truncate tool name if it exceeds OpenAI's 64-char limit
+ original_name = tool_choice.get("name", "")
+ truncated_name = truncate_tool_name(original_name)
tc_function_param = ChatCompletionToolChoiceFunctionParam(
- name=tool_choice.get("name", "")
+ name=truncated_name
)
return ChatCompletionToolChoiceObjectParam(
type="function", function=tc_function_param
@@ -600,12 +700,28 @@ class LiteLLMAnthropicMessagesAdapter:
def translate_anthropic_tools_to_openai(
self, tools: List[AllAnthropicToolsValues], model: Optional[str] = None
- ) -> List[ChatCompletionToolParam]:
+ ) -> Tuple[List[ChatCompletionToolParam], Dict[str, str]]:
+ """
+ Translate Anthropic tools to OpenAI format.
+
+ Returns:
+ Tuple of (translated_tools, tool_name_mapping)
+ - tool_name_mapping maps truncated names back to original names
+ for tools that exceeded OpenAI's 64-char limit
+ """
new_tools: List[ChatCompletionToolParam] = []
+ tool_name_mapping: Dict[str, str] = {}
mapped_tool_params = ["name", "input_schema", "description", "cache_control"]
for tool in tools:
+ original_name = tool["name"]
+ truncated_name = truncate_tool_name(original_name)
+
+ # Store mapping if name was truncated
+ if truncated_name != original_name:
+ tool_name_mapping[truncated_name] = original_name
+
function_chunk = ChatCompletionToolParamFunctionChunk(
- name=tool["name"],
+ name=truncated_name,
)
if "input_schema" in tool:
function_chunk["parameters"] = tool["input_schema"] # type: ignore
@@ -619,7 +735,7 @@ class LiteLLMAnthropicMessagesAdapter:
self._add_cache_control_if_applicable(tool, tool_param, model)
new_tools.append(tool_param) # type: ignore[arg-type]
- return new_tools # type: ignore[return-value]
+ return new_tools, tool_name_mapping # type: ignore[return-value]
def translate_anthropic_output_format_to_openai(
self, output_format: Any
@@ -694,12 +810,18 @@ class LiteLLMAnthropicMessagesAdapter:
def translate_anthropic_to_openai(
self, anthropic_message_request: AnthropicMessagesRequest
- ) -> ChatCompletionRequest:
+ ) -> Tuple[ChatCompletionRequest, Dict[str, str]]:
"""
This is used by the beta Anthropic Adapter, for translating anthropic `/v1/messages` requests to the openai format.
+
+ Returns:
+ Tuple of (openai_request, tool_name_mapping)
+ - tool_name_mapping maps truncated tool names back to original names
+ for tools that exceeded OpenAI's 64-char limit
"""
# Debug: Processing Anthropic message request
new_messages: List[AllMessageValues] = []
+ tool_name_mapping: Dict[str, str] = {}
## CONVERT ANTHROPIC MESSAGES TO OPENAI
messages_list: List[
@@ -750,7 +872,7 @@ class LiteLLMAnthropicMessagesAdapter:
if "tools" in anthropic_message_request:
tools = anthropic_message_request["tools"]
if tools:
- new_kwargs["tools"] = self.translate_anthropic_tools_to_openai(
+ new_kwargs["tools"], tool_name_mapping = self.translate_anthropic_tools_to_openai(
tools=cast(List[AllAnthropicToolsValues], tools),
model=new_kwargs.get("model"),
)
@@ -784,7 +906,7 @@ class LiteLLMAnthropicMessagesAdapter:
if k not in translatable_params: # pass remaining params as is
new_kwargs[k] = v # type: ignore
- return new_kwargs
+ return new_kwargs, tool_name_mapping
def _translate_anthropic_image_to_openai(self, image_source: dict) -> Optional[str]:
"""
@@ -813,22 +935,12 @@ class LiteLLMAnthropicMessagesAdapter:
return None
- def _translate_openai_content_to_anthropic(self, choices: List[Choices]) -> List[
- Union[
- AnthropicResponseContentBlockText,
- AnthropicResponseContentBlockToolUse,
- AnthropicResponseContentBlockThinking,
- AnthropicResponseContentBlockRedactedThinking,
- ]
- ]:
- new_content: List[
- Union[
- AnthropicResponseContentBlockText,
- AnthropicResponseContentBlockToolUse,
- AnthropicResponseContentBlockThinking,
- AnthropicResponseContentBlockRedactedThinking,
- ]
- ] = []
+ def _translate_openai_content_to_anthropic(
+ self,
+ choices: List[Choices],
+ tool_name_mapping: Optional[Dict[str, str]] = None,
+ ) -> List[Dict[str, Any]]:
+ new_content: List[Dict[str, Any]] = []
for choice in choices:
# Handle thinking blocks first
if (
@@ -852,7 +964,7 @@ class LiteLLMAnthropicMessagesAdapter:
if signature_value is not None
else None
),
- )
+ ).model_dump()
)
elif thinking_block.get("type") == "redacted_thinking":
data_value = thinking_block.get("data", "")
@@ -860,15 +972,27 @@ class LiteLLMAnthropicMessagesAdapter:
AnthropicResponseContentBlockRedactedThinking(
type="redacted_thinking",
data=str(data_value) if data_value is not None else "",
- )
+ ).model_dump()
)
+ # Handle reasoning_content when thinking_blocks is not present
+ elif (
+ hasattr(choice.message, "reasoning_content")
+ and choice.message.reasoning_content
+ ):
+ new_content.append(
+ AnthropicResponseContentBlockThinking(
+ type="thinking",
+ thinking=str(choice.message.reasoning_content),
+ signature=None,
+ ).model_dump()
+ )
# Handle text content
if choice.message.content is not None:
new_content.append(
AnthropicResponseContentBlockText(
type="text", text=choice.message.content
- )
+ ).model_dump()
)
# Handle tool calls (in parallel to text content)
if (
@@ -883,13 +1007,21 @@ class LiteLLMAnthropicMessagesAdapter:
if signature:
provider_specific_fields["signature"] = signature
+ # Restore original tool name if it was truncated
+ truncated_name = tool_call.function.name or ""
+ original_name = (
+ tool_name_mapping.get(truncated_name, truncated_name)
+ if tool_name_mapping
+ else truncated_name
+ )
+
tool_use_block = AnthropicResponseContentBlockToolUse(
type="tool_use",
id=tool_call.id,
- name=tool_call.function.name or "",
+ name=original_name,
input=parse_tool_call_arguments(
tool_call.function.arguments,
- tool_name=tool_call.function.name,
+ tool_name=original_name,
context="Anthropic pass-through adapter",
),
)
@@ -898,7 +1030,7 @@ class LiteLLMAnthropicMessagesAdapter:
tool_use_block.provider_specific_fields = (
provider_specific_fields
)
- new_content.append(tool_use_block)
+ new_content.append(tool_use_block.model_dump())
return new_content
@@ -914,10 +1046,24 @@ class LiteLLMAnthropicMessagesAdapter:
return "end_turn"
def translate_openai_response_to_anthropic(
- self, response: ModelResponse
+ self,
+ response: ModelResponse,
+ tool_name_mapping: Optional[Dict[str, str]] = None,
) -> AnthropicMessagesResponse:
+ """
+ Translate OpenAI response to Anthropic format.
+
+ Args:
+ response: The OpenAI ModelResponse
+ tool_name_mapping: Optional mapping of truncated tool names to original names.
+ Used to restore original names for tools that exceeded
+ OpenAI's 64-char limit.
+ """
## translate content block
- anthropic_content = self._translate_openai_content_to_anthropic(choices=response.choices) # type: ignore
+ anthropic_content = self._translate_openai_content_to_anthropic(
+ choices=response.choices, # type: ignore
+ tool_name_mapping=tool_name_mapping,
+ )
## extract finish reason
anthropic_finish_reason = self._translate_openai_finish_reason_to_anthropic(
openai_finish_reason=response.choices[0].finish_reason # type: ignore
@@ -1036,6 +1182,13 @@ class LiteLLMAnthropicMessagesAdapter:
reasoning_content += thinking
reasoning_signature += signature
+ # Handle reasoning_content when thinking_blocks is not present
+ # This handles providers like OpenRouter that return reasoning_content
+ elif isinstance(choice, StreamingChoices) and hasattr(
+ choice.delta, "reasoning_content"
+ ):
+ if choice.delta.reasoning_content is not None:
+ reasoning_content += choice.delta.reasoning_content
if reasoning_content and reasoning_signature:
raise ValueError(
diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py
index 308bf367d06..bb40f9df266 100644
--- a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py
+++ b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py
@@ -2,6 +2,9 @@ from typing import Any, AsyncIterator, Dict, List, Optional, Tuple
import httpx
+from litellm.anthropic_beta_headers_manager import (
+ update_headers_with_filtered_beta,
+)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import verbose_logger
from litellm.llms.base_llm.anthropic_messages.transformation import (
@@ -90,6 +93,11 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
optional_params=optional_params,
)
+ headers = update_headers_with_filtered_beta(
+ headers=headers,
+ provider="anthropic",
+ )
+
return headers, api_base
def transform_anthropic_messages_request(
@@ -189,8 +197,27 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
beta_values.update(b.strip() for b in existing_beta.split(","))
# Check for context management
- if optional_params.get("context_management") is not None:
- beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value)
+ context_management_param = optional_params.get("context_management")
+ if context_management_param is not None:
+ # Check edits array for compact_20260112 type
+ edits = context_management_param.get("edits", [])
+ has_compact = False
+ has_other = False
+
+ for edit in edits:
+ edit_type = edit.get("type", "")
+ if edit_type == "compact_20260112":
+ has_compact = True
+ else:
+ has_other = True
+
+ # Add compact header if any compact edits exist
+ if has_compact:
+ beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.COMPACT_2026_01_12.value)
+
+ # Add context management header if any other edits exist
+ if has_other:
+ beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value)
# Check for structured outputs
if optional_params.get("output_format") is not None:
diff --git a/litellm/llms/azure_ai/anthropic/count_tokens/transformation.py b/litellm/llms/azure_ai/anthropic/count_tokens/transformation.py
index e284595cc8a..09b83b7c971 100644
--- a/litellm/llms/azure_ai/anthropic/count_tokens/transformation.py
+++ b/litellm/llms/azure_ai/anthropic/count_tokens/transformation.py
@@ -30,30 +30,32 @@ class AzureAIAnthropicCountTokensConfig(AnthropicCountTokensConfig):
"""
Get the required headers for the Azure AI Anthropic CountTokens API.
- Uses Azure authentication (api-key header) instead of Anthropic's x-api-key.
+ Azure AI Anthropic uses Anthropic's native API format, which requires the
+ x-api-key header for authentication (in addition to Azure's api-key header).
Args:
api_key: The Azure AI API key
litellm_params: Optional LiteLLM parameters for additional auth config
Returns:
- Dictionary of required headers with Azure authentication
+ Dictionary of required headers with both x-api-key and Azure authentication
"""
- # Start with base headers
+ # Start with base headers including x-api-key for Anthropic API compatibility
headers = {
"Content-Type": "application/json",
"anthropic-version": "2023-06-01",
"anthropic-beta": ANTHROPIC_TOKEN_COUNTING_BETA_VERSION,
+ "x-api-key": api_key, # Azure AI Anthropic requires this header
}
- # Use Azure authentication
+ # Also set up Azure auth headers for flexibility
litellm_params = litellm_params or {}
if "api_key" not in litellm_params:
litellm_params["api_key"] = api_key
litellm_params_obj = GenericLiteLLMParams(**litellm_params)
- # Get Azure auth headers
+ # Get Azure auth headers (api-key or Authorization)
azure_headers = BaseAzureLLM._base_validate_azure_environment(
headers={}, litellm_params=litellm_params_obj
)
@@ -68,7 +70,7 @@ class AzureAIAnthropicCountTokensConfig(AnthropicCountTokensConfig):
Get the Azure AI Anthropic CountTokens API endpoint.
Args:
- api_base: The Azure AI API base URL
+ api_base: The Azure AI API base URL
(e.g., https://my-resource.services.ai.azure.com or
https://my-resource.services.ai.azure.com/anthropic)
diff --git a/litellm/llms/azure_ai/anthropic/transformation.py b/litellm/llms/azure_ai/anthropic/transformation.py
index 2d8d3b987c7..753bc9c08eb 100644
--- a/litellm/llms/azure_ai/anthropic/transformation.py
+++ b/litellm/llms/azure_ai/anthropic/transformation.py
@@ -3,6 +3,9 @@ Azure Anthropic transformation config - extends AnthropicConfig with Azure authe
"""
from typing import TYPE_CHECKING, Dict, List, Optional, Union
+from litellm.anthropic_beta_headers_manager import (
+ update_headers_with_filtered_beta,
+)
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.types.llms.openai import AllMessageValues
@@ -87,6 +90,12 @@ class AzureAnthropicConfig(AnthropicConfig):
if "anthropic-version" not in headers:
headers["anthropic-version"] = "2023-06-01"
+ # Filter out unsupported beta headers for Azure AI
+ headers = update_headers_with_filtered_beta(
+ headers=headers,
+ provider="azure_ai",
+ )
+
return headers
def transform_request(
diff --git a/litellm/llms/azure_ai/rerank/transformation.py b/litellm/llms/azure_ai/rerank/transformation.py
index a47b6082c37..f577a42ed58 100644
--- a/litellm/llms/azure_ai/rerank/transformation.py
+++ b/litellm/llms/azure_ai/rerank/transformation.py
@@ -11,6 +11,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.llms.cohere.rerank.transformation import CohereRerankConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.utils import RerankResponse
+from litellm.utils import _add_path_to_api_base
class AzureAIRerankConfig(CohereRerankConfig):
@@ -28,9 +29,34 @@ class AzureAIRerankConfig(CohereRerankConfig):
raise ValueError(
"Azure AI API Base is required. api_base=None. Set in call or via `AZURE_AI_API_BASE` env var."
)
- if not api_base.endswith("/v1/rerank"):
- api_base = f"{api_base}/v1/rerank"
- return api_base
+ original_url = httpx.URL(api_base)
+ if not original_url.is_absolute_url:
+ raise ValueError(
+ "Azure AI API Base must be an absolute URL including scheme (e.g. "
+ "'https://.services.ai.azure.com'). "
+ f"Got api_base={api_base!r}."
+ )
+ normalized_path = original_url.path.rstrip("/")
+
+ # Allow callers to pass either full v1/v2 rerank endpoints:
+ # - https://.services.ai.azure.com/v1/rerank
+ # - https://.services.ai.azure.com/providers/cohere/v2/rerank
+ if normalized_path.endswith("/v1/rerank") or normalized_path.endswith("/v2/rerank"):
+ return str(original_url.copy_with(path=normalized_path or "/"))
+
+ # If callers pass just the version path (e.g. ".../v2" or ".../providers/cohere/v2"), append "/rerank"
+ if (
+ normalized_path.endswith("/v1")
+ or normalized_path.endswith("/v2")
+ or normalized_path.endswith("/providers/cohere/v2")
+ ):
+ return _add_path_to_api_base(
+ api_base=str(original_url.copy_with(path=normalized_path or "/")),
+ ending_path="/rerank",
+ )
+
+ # Backwards compatible default: Azure AI rerank was originally exposed under /v1/rerank
+ return _add_path_to_api_base(api_base=api_base, ending_path="/v1/rerank")
def validate_environment(
self,
diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py
index d4e4d3591ba..7fc51263ebb 100644
--- a/litellm/llms/bedrock/chat/converse_transformation.py
+++ b/litellm/llms/bedrock/chat/converse_transformation.py
@@ -11,6 +11,9 @@ import httpx
import litellm
from litellm._logging import verbose_logger
+from litellm.anthropic_beta_headers_manager import (
+ filter_and_transform_beta_headers,
+)
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.litellm_core_utils.core_helpers import (
filter_exceptions_from_params,
@@ -66,6 +69,7 @@ from ..common_utils import (
BedrockModelInfo,
get_anthropic_beta_from_headers,
get_bedrock_tool_name,
+ is_claude_4_5_on_bedrock,
)
# Computer use tool prefixes supported by Bedrock
@@ -81,6 +85,7 @@ BEDROCK_COMPUTER_USE_TOOLS = [
UNSUPPORTED_BEDROCK_CONVERSE_BETA_PATTERNS = [
"advanced-tool-use", # Bedrock Converse doesn't support advanced-tool-use beta headers
"prompt-caching", # Prompt caching not supported in Converse API
+ "compact-2026-01-12", # The compact beta feature is not currently supported on the Converse and ConverseStream APIs
]
@@ -306,9 +311,7 @@ class AmazonConverseConfig(BaseConfig):
return "nova-2-lite" in model_without_region
def _map_web_search_options(
- self,
- web_search_options: dict,
- model: str
+ self, web_search_options: dict, model: str
) -> Optional[BedrockToolBlock]:
"""
Map web_search_options to Nova grounding systemTool.
@@ -431,7 +434,7 @@ class AmazonConverseConfig(BaseConfig):
else:
# Anthropic and other models: convert to thinking parameter
optional_params["thinking"] = AnthropicConfig._map_reasoning_effort(
- reasoning_effort
+ reasoning_effort=reasoning_effort, model=model
)
def get_supported_openai_params(self, model: str) -> List[str]:
@@ -617,37 +620,6 @@ class AmazonConverseConfig(BaseConfig):
return transformed_tools
- def _filter_unsupported_beta_headers_for_bedrock(
- self, model: str, beta_list: list
- ) -> list:
- """
- Remove beta headers that are not supported on Bedrock Converse API for the given model.
-
- Extended thinking beta headers are only supported on specific Claude 4+ models.
- Some beta headers are universally unsupported on Bedrock Converse API.
-
- Args:
- model: The model name
- beta_list: The list of beta headers to filter
-
- Returns:
- Filtered list of beta headers
- """
- filtered_betas = []
-
- # 1. Filter out beta headers that are universally unsupported on Bedrock Converse
- for beta in beta_list:
- should_keep = True
- for unsupported_pattern in UNSUPPORTED_BEDROCK_CONVERSE_BETA_PATTERNS:
- if unsupported_pattern in beta.lower():
- should_keep = False
- break
-
- if should_keep:
- filtered_betas.append(beta)
-
- return filtered_betas
-
def _separate_computer_use_tools(
self, tools: List[OpenAIChatCompletionToolParam], model: str
) -> Tuple[
@@ -808,11 +780,11 @@ class AmazonConverseConfig(BaseConfig):
if param == "web_search_options" and isinstance(value, dict):
# Note: we use `isinstance(value, dict)` instead of `value and isinstance(value, dict)`
# because empty dict {} is falsy but is a valid way to enable Nova grounding
- grounding_tool = self._map_web_search_options(value, model)
- if grounding_tool is not None:
- optional_params = self._add_tools_to_optional_params(
- optional_params=optional_params, tools=[grounding_tool]
- )
+ grounding_tool = self._map_web_search_options(value, model)
+ if grounding_tool is not None:
+ optional_params = self._add_tools_to_optional_params(
+ optional_params=optional_params, tools=[grounding_tool]
+ )
# Only update thinking tokens for non-GPT-OSS models and non-Nova-Lite-2 models
# Nova Lite 2 handles token budgeting differently through reasoningConfig
@@ -926,6 +898,7 @@ class AmazonConverseConfig(BaseConfig):
ChatCompletionAssistantMessage,
],
block_type: Literal["system"],
+ model: Optional[str] = None,
) -> Optional[SystemContentBlock]:
pass
@@ -939,6 +912,7 @@ class AmazonConverseConfig(BaseConfig):
ChatCompletionAssistantMessage,
],
block_type: Literal["content_block"],
+ model: Optional[str] = None,
) -> Optional[ContentBlock]:
pass
@@ -951,16 +925,26 @@ class AmazonConverseConfig(BaseConfig):
ChatCompletionAssistantMessage,
],
block_type: Literal["system", "content_block"],
+ model: Optional[str] = None,
) -> Optional[Union[SystemContentBlock, ContentBlock]]:
- if message_block.get("cache_control", None) is None:
+ cache_control = message_block.get("cache_control", None)
+ if cache_control is None:
return None
+
+ cache_point = CachePointBlock(type="default")
+ if isinstance(cache_control, dict) and "ttl" in cache_control:
+ ttl = cache_control["ttl"]
+ if ttl in ["5m", "1h"] and model is not None:
+ if is_claude_4_5_on_bedrock(model):
+ cache_point["ttl"] = ttl
+
if block_type == "system":
- return SystemContentBlock(cachePoint=CachePointBlock(type="default"))
+ return SystemContentBlock(cachePoint=cache_point)
else:
- return ContentBlock(cachePoint=CachePointBlock(type="default"))
+ return ContentBlock(cachePoint=cache_point)
def _transform_system_message(
- self, messages: List[AllMessageValues]
+ self, messages: List[AllMessageValues], model: Optional[str] = None
) -> Tuple[List[AllMessageValues], List[SystemContentBlock]]:
system_prompt_indices = []
system_content_blocks: List[SystemContentBlock] = []
@@ -972,7 +956,7 @@ class AmazonConverseConfig(BaseConfig):
SystemContentBlock(text=message["content"])
)
cache_block = self._get_cache_point_block(
- message, block_type="system"
+ message, block_type="system", model=model
)
if cache_block:
system_content_blocks.append(cache_block)
@@ -983,7 +967,7 @@ class AmazonConverseConfig(BaseConfig):
SystemContentBlock(text=m["text"])
)
cache_block = self._get_cache_point_block(
- m, block_type="system"
+ m, block_type="system", model=model
)
if cache_block:
system_content_blocks.append(cache_block)
@@ -1081,10 +1065,16 @@ class AmazonConverseConfig(BaseConfig):
user_betas = get_anthropic_beta_from_headers(headers)
anthropic_beta_list.extend(user_betas)
- # Filter out tool search tools - Bedrock Converse API doesn't support them
+ # Separate pre-formatted Bedrock tools (e.g. systemTool from web_search_options)
+ # from OpenAI-format tools that need transformation via _bedrock_tools_pt
filtered_tools = []
+ pre_formatted_tools: List[ToolBlock] = []
if original_tools:
for tool in original_tools:
+ # Already-formatted Bedrock tools (e.g. systemTool for Nova grounding)
+ if "systemTool" in tool:
+ pre_formatted_tools.append(tool)
+ continue
tool_type = tool.get("type", "")
if tool_type in (
"tool_search_tool_regex_20251119",
@@ -1106,7 +1096,28 @@ class AmazonConverseConfig(BaseConfig):
# Add computer use tools and anthropic_beta if needed (only when computer use tools are present)
if computer_use_tools:
- anthropic_beta_list.append("computer-use-2024-10-22")
+ # Determine the correct computer-use beta header based on model
+ # "computer-use-2025-11-24" for Claude Opus 4.6, Claude Opus 4.5
+ # "computer-use-2025-01-24" for Claude Sonnet 4.5, Haiku 4.5, Opus 4.1, Sonnet 4, Opus 4, and Sonnet 3.7
+ # "computer-use-2024-10-22" for older models
+ model_lower = model.lower()
+ if "opus-4.6" in model_lower or "opus_4.6" in model_lower or "opus-4-6" in model_lower or "opus_4_6" in model_lower:
+ computer_use_header = "computer-use-2025-11-24"
+ elif "opus-4.5" in model_lower or "opus_4.5" in model_lower or "opus-4-5" in model_lower or "opus_4_5" in model_lower:
+ computer_use_header = "computer-use-2025-11-24"
+ elif any(pattern in model_lower for pattern in [
+ "sonnet-4.5", "sonnet_4.5", "sonnet-4-5", "sonnet_4_5",
+ "haiku-4.5", "haiku_4.5", "haiku-4-5", "haiku_4_5",
+ "opus-4.1", "opus_4.1", "opus-4-1", "opus_4_1",
+ "sonnet-4", "sonnet_4",
+ "opus-4", "opus_4",
+ "sonnet-3.7", "sonnet_3.7", "sonnet-3-7", "sonnet_3_7"
+ ]):
+ computer_use_header = "computer-use-2025-01-24"
+ else:
+ computer_use_header = "computer-use-2024-10-22"
+
+ anthropic_beta_list.append(computer_use_header)
# Transform computer use tools to proper Bedrock format
transformed_computer_tools = self._transform_computer_use_tools(
computer_use_tools
@@ -1116,6 +1127,9 @@ class AmazonConverseConfig(BaseConfig):
# No computer use tools, process all tools as regular tools
bedrock_tools = _bedrock_tools_pt(filtered_tools)
+ # Append pre-formatted tools (systemTool etc.) after transformation
+ bedrock_tools.extend(pre_formatted_tools)
+
# Set anthropic_beta in additional_request_params if we have any beta features
# ONLY apply to Anthropic/Claude models - other models (e.g., Qwen, Llama) don't support this field
# and will error with "unknown variant anthropic_beta" if included
@@ -1128,14 +1142,14 @@ class AmazonConverseConfig(BaseConfig):
if beta not in seen:
unique_betas.append(beta)
seen.add(beta)
-
- # Filter out unsupported beta headers for Bedrock Converse API
- filtered_betas = self._filter_unsupported_beta_headers_for_bedrock(
- model=model,
- beta_list=unique_betas,
+
+ filtered_betas = filter_and_transform_beta_headers(
+ beta_headers=unique_betas,
+ provider="bedrock_converse",
)
- additional_request_params["anthropic_beta"] = filtered_betas
+ if filtered_betas:
+ additional_request_params["anthropic_beta"] = filtered_betas
return bedrock_tools, anthropic_beta_list
@@ -1187,9 +1201,11 @@ class AmazonConverseConfig(BaseConfig):
)
# Prepare and separate parameters
- inference_params, additional_request_params, request_metadata = self._prepare_request_params(
- optional_params, model
- )
+ (
+ inference_params,
+ additional_request_params,
+ request_metadata,
+ ) = self._prepare_request_params(optional_params, model)
original_tools = inference_params.pop("tools", [])
@@ -1241,7 +1257,9 @@ class AmazonConverseConfig(BaseConfig):
litellm_params: dict,
headers: Optional[dict] = None,
) -> RequestObject:
- messages, system_content_blocks = self._transform_system_message(messages)
+ messages, system_content_blocks = self._transform_system_message(
+ messages, model=model
+ )
# Convert last user message to guarded_text if guardrailConfig is present
messages = self._convert_consecutive_user_messages_to_guarded_text(
@@ -1297,7 +1315,9 @@ class AmazonConverseConfig(BaseConfig):
litellm_params: dict,
headers: Optional[dict] = None,
) -> RequestObject:
- messages, system_content_blocks = self._transform_system_message(messages)
+ messages, system_content_blocks = self._transform_system_message(
+ messages, model=model
+ )
# Convert last user message to guarded_text if guardrailConfig is present
messages = self._convert_consecutive_user_messages_to_guarded_text(
@@ -1475,7 +1495,9 @@ class AmazonConverseConfig(BaseConfig):
return message, returned_finish_reason
- def _translate_message_content(self, content_blocks: List[ContentBlock]) -> Tuple[
+ def _translate_message_content(
+ self, content_blocks: List[ContentBlock]
+ ) -> Tuple[
str,
List[ChatCompletionToolCallChunk],
Optional[List[BedrockConverseReasoningContentBlock]],
@@ -1492,9 +1514,9 @@ class AmazonConverseConfig(BaseConfig):
"""
content_str = ""
tools: List[ChatCompletionToolCallChunk] = []
- reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = (
- None
- )
+ reasoningContentBlocks: Optional[
+ List[BedrockConverseReasoningContentBlock]
+ ] = None
citationsContentBlocks: Optional[List[CitationsContentBlock]] = None
for idx, content in enumerate(content_blocks):
"""
@@ -1548,7 +1570,7 @@ class AmazonConverseConfig(BaseConfig):
return content_str, tools, reasoningContentBlocks, citationsContentBlocks
- def _transform_response( # noqa: PLR0915
+ def _transform_response( # noqa: PLR0915
self,
model: str,
response: httpx.Response,
@@ -1621,9 +1643,9 @@ class AmazonConverseConfig(BaseConfig):
chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"}
content_str = ""
tools: List[ChatCompletionToolCallChunk] = []
- reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = (
- None
- )
+ reasoningContentBlocks: Optional[
+ List[BedrockConverseReasoningContentBlock]
+ ] = None
citationsContentBlocks: Optional[List[CitationsContentBlock]] = None
if message is not None:
@@ -1642,15 +1664,17 @@ class AmazonConverseConfig(BaseConfig):
provider_specific_fields["citationsContent"] = citationsContentBlocks
if provider_specific_fields:
- chat_completion_message["provider_specific_fields"] = provider_specific_fields
+ chat_completion_message[
+ "provider_specific_fields"
+ ] = provider_specific_fields
if reasoningContentBlocks is not None:
- chat_completion_message["reasoning_content"] = (
- self._transform_reasoning_content(reasoningContentBlocks)
- )
- chat_completion_message["thinking_blocks"] = (
- self._transform_thinking_blocks(reasoningContentBlocks)
- )
+ chat_completion_message[
+ "reasoning_content"
+ ] = self._transform_reasoning_content(reasoningContentBlocks)
+ chat_completion_message[
+ "thinking_blocks"
+ ] = self._transform_thinking_blocks(reasoningContentBlocks)
chat_completion_message["content"] = content_str
if (
json_mode is True
diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py
index 65d237bdbdf..4c87f6fa994 100644
--- a/litellm/llms/bedrock/common_utils.py
+++ b/litellm/llms/bedrock/common_utils.py
@@ -446,6 +446,29 @@ def get_bedrock_base_model(model: str) -> str:
return model
+def is_claude_4_5_on_bedrock(model: str) -> bool:
+ """
+ Check if the model is a Claude 4.5 model on Bedrock.
+ Claude 4.5 models support prompt caching with '5m' and '1h' TTL on Bedrock.
+ """
+ model_lower = model.lower()
+ claude_4_5_patterns = [
+ "sonnet-4.5",
+ "sonnet_4.5",
+ "sonnet-4-5",
+ "sonnet_4_5",
+ "haiku-4.5",
+ "haiku_4.5",
+ "haiku-4-5",
+ "haiku_4_5",
+ "opus-4.5",
+ "opus_4.5",
+ "opus-4-5",
+ "opus_4_5",
+ ]
+ return any(pattern in model_lower for pattern in claude_4_5_patterns)
+
+
# Import after standalone functions to avoid circular imports
from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter
@@ -815,21 +838,23 @@ def get_anthropic_beta_from_headers(headers: dict) -> List[str]:
# If it's already a list, return it
if isinstance(anthropic_beta_header, list):
return anthropic_beta_header
-
+
# Try to parse as JSON array first (e.g., '["interleaved-thinking-2025-05-14", "claude-code-20250219"]')
if isinstance(anthropic_beta_header, str):
anthropic_beta_header = anthropic_beta_header.strip()
- if anthropic_beta_header.startswith("[") and anthropic_beta_header.endswith("]"):
+ if anthropic_beta_header.startswith("[") and anthropic_beta_header.endswith(
+ "]"
+ ):
try:
parsed = json.loads(anthropic_beta_header)
if isinstance(parsed, list):
return [str(beta).strip() for beta in parsed]
except json.JSONDecodeError:
pass # Fall through to comma-separated parsing
-
+
# Fall back to comma-separated values
return [beta.strip() for beta in anthropic_beta_header.split(",")]
-
+
return []
diff --git a/litellm/llms/bedrock/embed/cohere_transformation.py b/litellm/llms/bedrock/embed/cohere_transformation.py
index 490cd71b793..d00cb74aae0 100644
--- a/litellm/llms/bedrock/embed/cohere_transformation.py
+++ b/litellm/llms/bedrock/embed/cohere_transformation.py
@@ -15,7 +15,7 @@ class BedrockCohereEmbeddingConfig:
pass
def get_supported_openai_params(self) -> List[str]:
- return ["encoding_format"]
+ return ["encoding_format", "dimensions"]
def map_openai_params(
self, non_default_params: dict, optional_params: dict
@@ -23,6 +23,8 @@ class BedrockCohereEmbeddingConfig:
for k, v in non_default_params.items():
if k == "encoding_format":
optional_params["embedding_types"] = v
+ elif k == "dimensions":
+ optional_params["output_dimension"] = v
return optional_params
def _is_v3_model(self, model: str) -> bool:
diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py
index b1c45ea83a2..19fe7d8c140 100644
--- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py
+++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py
@@ -12,6 +12,9 @@ from typing import (
import httpx
+from litellm.anthropic_beta_headers_manager import (
+ filter_and_transform_beta_headers,
+)
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
@@ -23,7 +26,10 @@ from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
AmazonInvokeConfig,
)
-from litellm.llms.bedrock.common_utils import get_anthropic_beta_from_headers
+from litellm.llms.bedrock.common_utils import (
+ get_anthropic_beta_from_headers,
+ is_claude_4_5_on_bedrock,
+)
from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER
from litellm.types.llms.openai import AllMessageValues
from litellm.types.router import GenericLiteLLMParams
@@ -52,10 +58,6 @@ class AmazonAnthropicClaudeMessagesConfig(
# Beta header patterns that are not supported by Bedrock Invoke API
# These will be filtered out to prevent 400 "invalid beta flag" errors
- UNSUPPORTED_BEDROCK_INVOKE_BETA_PATTERNS = [
- "advanced-tool-use", # Bedrock Invoke doesn't support advanced-tool-use beta headers
- "prompt-caching-scope"
- ]
def __init__(self, **kwargs):
BaseAnthropicMessagesConfig.__init__(self, **kwargs)
@@ -116,15 +118,22 @@ class AmazonAnthropicClaudeMessagesConfig(
)
def _remove_ttl_from_cache_control(
- self, anthropic_messages_request: Dict
+ self, anthropic_messages_request: Dict, model: Optional[str] = None
) -> None:
"""
Remove `ttl` field from cache_control in messages.
Bedrock doesn't support the ttl field in cache_control.
+ Update: Bedock supports `5m` and `1h` for Claude 4.5 models.
+
Args:
anthropic_messages_request: The request dictionary to modify in-place
+ model: The model name to check if it supports ttl
"""
+ is_claude_4_5 = False
+ if model:
+ is_claude_4_5 = self._is_claude_4_5_on_bedrock(model)
+
if "messages" in anthropic_messages_request:
for message in anthropic_messages_request["messages"]:
if isinstance(message, dict) and "content" in message:
@@ -133,7 +142,14 @@ class AmazonAnthropicClaudeMessagesConfig(
for item in content:
if isinstance(item, dict) and "cache_control" in item:
cache_control = item["cache_control"]
- if isinstance(cache_control, dict) and "ttl" in cache_control:
+ if (
+ isinstance(cache_control, dict)
+ and "ttl" in cache_control
+ ):
+ ttl = cache_control["ttl"]
+ if is_claude_4_5 and ttl in ["5m", "1h"]:
+ continue
+
cache_control.pop("ttl", None)
def _supports_extended_thinking_on_bedrock(self, model: str) -> bool:
@@ -155,10 +171,18 @@ class AmazonAnthropicClaudeMessagesConfig(
# Supported models on Bedrock for extended thinking
supported_patterns = [
- "opus-4.5", "opus_4.5", "opus-4-5", "opus_4_5", # Opus 4.5
- "opus-4.1", "opus_4.1", "opus-4-1", "opus_4_1", # Opus 4.1
- "opus-4", "opus_4", # Opus 4
- "sonnet-4", "sonnet_4", # Sonnet 4
+ "opus-4.5",
+ "opus_4.5",
+ "opus-4-5",
+ "opus_4_5", # Opus 4.5
+ "opus-4.1",
+ "opus_4.1",
+ "opus-4-1",
+ "opus_4_1", # Opus 4.1
+ "opus-4",
+ "opus_4", # Opus 4
+ "sonnet-4",
+ "sonnet_4", # Sonnet 4
]
return any(pattern in model_lower for pattern in supported_patterns)
@@ -175,10 +199,27 @@ class AmazonAnthropicClaudeMessagesConfig(
"""
model_lower = model.lower()
opus_4_5_patterns = [
- "opus-4.5", "opus_4.5", "opus-4-5", "opus_4_5",
+ "opus-4.5",
+ "opus_4.5",
+ "opus-4-5",
+ "opus_4_5",
]
return any(pattern in model_lower for pattern in opus_4_5_patterns)
+ def _is_claude_4_5_on_bedrock(self, model: str) -> bool:
+ """
+ Check if the model is Claude 4.5 on Bedrock.
+
+ Claude Sonnet 4.5, Haiku 4.5, and Opus 4.5 support 1-hour prompt caching.
+
+ Args:
+ model: The model name
+
+ Returns:
+ True if the model is Claude 4.5
+ """
+ return is_claude_4_5_on_bedrock(model)
+
def _supports_tool_search_on_bedrock(self, model: str) -> bool:
"""
Check if the model supports tool search on Bedrock.
@@ -199,9 +240,15 @@ class AmazonAnthropicClaudeMessagesConfig(
# Supported models for tool search on Bedrock
supported_patterns = [
# Opus 4.5
- "opus-4.5", "opus_4.5", "opus-4-5", "opus_4_5",
+ "opus-4.5",
+ "opus_4.5",
+ "opus-4-5",
+ "opus_4_5",
# Sonnet 4.5
- "sonnet-4.5", "sonnet_4.5", "sonnet-4-5", "sonnet_4_5",
+ "sonnet-4.5",
+ "sonnet_4.5",
+ "sonnet-4-5",
+ "sonnet_4_5",
]
return any(pattern in model_lower for pattern in supported_patterns)
@@ -228,41 +275,48 @@ class AmazonAnthropicClaudeMessagesConfig(
model: The model name
beta_set: The set of beta headers to filter in-place
"""
- beta_headers_to_remove = set()
- has_advanced_tool_use = False
-
- # 1. Filter out beta headers that are universally unsupported on Bedrock Invoke and track if advanced-tool-use header is present
- for beta in beta_set:
- for unsupported_pattern in self.UNSUPPORTED_BEDROCK_INVOKE_BETA_PATTERNS:
- if unsupported_pattern in beta.lower():
- beta_headers_to_remove.add(beta)
- has_advanced_tool_use = True
- break
-
+ # 1. Handle header transformations BEFORE filtering
+ # (advanced-tool-use -> tool-search-tool)
+ # This must happen before filtering because advanced-tool-use is in the unsupported list
+ has_advanced_tool_use = "advanced-tool-use-2025-11-20" in beta_set
+ if has_advanced_tool_use and self._supports_tool_search_on_bedrock(model):
+ beta_set.discard("advanced-tool-use-2025-11-20")
+ beta_set.add("tool-search-tool-2025-10-19")
+ beta_set.add("tool-examples-2025-10-29")
- # 2. Filter out extended thinking headers for models that don't support them
+ # 2. Apply provider-level filtering using centralized JSON config
+ beta_list = list(beta_set)
+ filtered_list = filter_and_transform_beta_headers(
+ beta_headers=beta_list,
+ provider="bedrock",
+ )
+
+ # Update the set with filtered headers
+ beta_set.clear()
+ beta_set.update(filtered_list)
+
+ # 2.1. Handle model-specific exceptions: structured-outputs is only supported on Opus 4.6
+ # Re-add structured-outputs if it was in the original set and model is Opus 4.6
+ model_lower = model.lower()
+ is_opus_4_6 = any(pattern in model_lower for pattern in ["opus-4.6", "opus_4.6", "opus-4-6", "opus_4_6"])
+ if is_opus_4_6 and "structured-outputs-2025-11-13" in beta_list:
+ beta_set.add("structured-outputs-2025-11-13")
+
+ # 3. Filter out extended thinking headers for models that don't support them
extended_thinking_patterns = [
"extended-thinking",
"interleaved-thinking",
]
if not self._supports_extended_thinking_on_bedrock(model):
+ beta_headers_to_remove = set()
for beta in beta_set:
for pattern in extended_thinking_patterns:
if pattern in beta.lower():
beta_headers_to_remove.add(beta)
break
-
- # Remove all filtered headers
- for beta in beta_headers_to_remove:
- beta_set.discard(beta)
-
- # 3. Translate advanced-tool-use to Bedrock-specific headers for models that support tool search
- # Ref: https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-anthropic-claude-messages-request-response.html
- # Ref: https://platform.claude.com/docs/en/agents-and-tools/tool-use/tool-search-tool
- if has_advanced_tool_use and self._supports_tool_search_on_bedrock(model):
- beta_set.add("tool-search-tool-2025-10-19")
- beta_set.add("tool-examples-2025-10-29")
-
+
+ for beta in beta_headers_to_remove:
+ beta_set.discard(beta)
def _get_tool_search_beta_header_for_bedrock(
self,
@@ -290,7 +344,9 @@ class AmazonAnthropicClaudeMessagesConfig(
input_examples_used: Whether input examples are used
beta_set: The set of beta headers to modify in-place
"""
- if tool_search_used and not (programmatic_tool_calling_used or input_examples_used):
+ if tool_search_used and not (
+ programmatic_tool_calling_used or input_examples_used
+ ):
beta_set.discard(ANTHROPIC_TOOL_SEARCH_BETA_HEADER)
if "opus-4" in model.lower() or "opus_4" in model.lower():
beta_set.add("tool-search-tool-2025-10-19")
@@ -302,13 +358,13 @@ class AmazonAnthropicClaudeMessagesConfig(
) -> None:
"""
Convert Anthropic output_format to inline schema in message content.
-
+
Bedrock Invoke doesn't support the output_format parameter, so we embed
the schema directly into the user message content as text instructions.
-
+
This approach adds the schema to the last user message, instructing the model
to respond in the specified JSON format.
-
+
Args:
output_format: The output_format dict with 'type' and 'schema'
anthropic_messages_request: The request dict to modify in-place
@@ -321,35 +377,32 @@ class AmazonAnthropicClaudeMessagesConfig(
schema = output_format.get("schema")
if not schema:
return
-
+
# Get messages from the request
messages = anthropic_messages_request.get("messages", [])
if not messages:
return
-
+
# Find the last user message
last_user_message_idx = None
for idx in range(len(messages) - 1, -1, -1):
if messages[idx].get("role") == "user":
last_user_message_idx = idx
break
-
+
if last_user_message_idx is None:
return
-
+
last_user_message = messages[last_user_message_idx]
content = last_user_message.get("content", [])
-
+
# Ensure content is a list
if isinstance(content, str):
content = [{"type": "text", "text": content}]
last_user_message["content"] = content
-
+
# Add schema as text content to the message
- schema_text = {
- "type": "text",
- "text": json.dumps(schema)
- }
+ schema_text = {"type": "text", "text": json.dumps(schema)}
content.append(schema_text)
def transform_anthropic_messages_request(
@@ -374,9 +427,9 @@ class AmazonAnthropicClaudeMessagesConfig(
# 1. anthropic_version is required for all claude models
if "anthropic_version" not in anthropic_messages_request:
- anthropic_messages_request["anthropic_version"] = (
- self.DEFAULT_BEDROCK_ANTHROPIC_API_VERSION
- )
+ anthropic_messages_request[
+ "anthropic_version"
+ ] = self.DEFAULT_BEDROCK_ANTHROPIC_API_VERSION
# 2. `stream` is not allowed in request body for bedrock invoke
if "stream" in anthropic_messages_request:
@@ -386,8 +439,10 @@ class AmazonAnthropicClaudeMessagesConfig(
if "model" in anthropic_messages_request:
anthropic_messages_request.pop("model", None)
- # 4. Remove `ttl` field from cache_control in messages (Bedrock doesn't support it)
- self._remove_ttl_from_cache_control(anthropic_messages_request)
+ # 4. Remove `ttl` field from cache_control in messages (Bedrock doesn't support it for older models)
+ self._remove_ttl_from_cache_control(
+ anthropic_messages_request=anthropic_messages_request, model=model
+ )
# 5. Convert `output_format` to inline schema (Bedrock invoke doesn't support output_format)
output_format = anthropic_messages_request.pop("output_format", None)
@@ -396,14 +451,14 @@ class AmazonAnthropicClaudeMessagesConfig(
output_format=output_format,
anthropic_messages_request=anthropic_messages_request,
)
-
+
# 6. AUTO-INJECT beta headers based on features used
anthropic_model_info = AnthropicModelInfo()
tools = anthropic_messages_optional_request_params.get("tools")
messages_typed = cast(List[AllMessageValues], messages)
tool_search_used = anthropic_model_info.is_tool_search_used(tools)
- programmatic_tool_calling_used = anthropic_model_info.is_programmatic_tool_calling_used(
- tools
+ programmatic_tool_calling_used = (
+ anthropic_model_info.is_programmatic_tool_calling_used(tools)
)
input_examples_used = anthropic_model_info.is_input_examples_used(tools)
@@ -436,8 +491,7 @@ class AmazonAnthropicClaudeMessagesConfig(
if beta_set:
anthropic_messages_request["anthropic_beta"] = list(beta_set)
-
-
+
return anthropic_messages_request
def get_async_streaming_response_iterator(
@@ -455,7 +509,7 @@ class AmazonAnthropicClaudeMessagesConfig(
)
# Convert decoded Bedrock events to Server-Sent Events expected by Anthropic clients.
return self.bedrock_sse_wrapper(
- completion_stream=completion_stream,
+ completion_stream=completion_stream,
litellm_logging_obj=litellm_logging_obj,
request_body=request_body,
)
@@ -474,14 +528,14 @@ class AmazonAnthropicClaudeMessagesConfig(
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
BaseAnthropicMessagesStreamingIterator,
)
+
handler = BaseAnthropicMessagesStreamingIterator(
litellm_logging_obj=litellm_logging_obj,
request_body=request_body,
)
-
+
async for chunk in handler.async_sse_wrapper(completion_stream):
yield chunk
-
class AmazonAnthropicClaudeMessagesStreamDecoder(AWSEventStreamDecoder):
diff --git a/litellm/llms/bedrock/realtime/handler.py b/litellm/llms/bedrock/realtime/handler.py
new file mode 100644
index 00000000000..9b6a80f4a2f
--- /dev/null
+++ b/litellm/llms/bedrock/realtime/handler.py
@@ -0,0 +1,307 @@
+"""
+This file contains the handler for AWS Bedrock Nova Sonic realtime API.
+
+This uses aws_sdk_bedrock_runtime for bidirectional streaming with Nova Sonic.
+"""
+
+import asyncio
+import json
+from typing import Any, Optional
+
+from litellm._logging import verbose_proxy_logger
+from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
+
+from ..base_aws_llm import BaseAWSLLM
+from .transformation import BedrockRealtimeConfig
+
+
+class BedrockRealtime(BaseAWSLLM):
+ """Handler for Bedrock Nova Sonic realtime speech-to-speech API."""
+
+ def __init__(self):
+ super().__init__()
+
+ async def async_realtime(
+ self,
+ model: str,
+ websocket: Any,
+ logging_obj: LiteLLMLogging,
+ api_base: Optional[str] = None,
+ api_key: Optional[str] = None,
+ timeout: Optional[float] = None,
+ aws_region_name: Optional[str] = None,
+ aws_access_key_id: Optional[str] = None,
+ aws_secret_access_key: Optional[str] = None,
+ aws_session_token: Optional[str] = None,
+ aws_role_name: Optional[str] = None,
+ aws_session_name: Optional[str] = None,
+ aws_profile_name: Optional[str] = None,
+ aws_web_identity_token: Optional[str] = None,
+ aws_sts_endpoint: Optional[str] = None,
+ aws_bedrock_runtime_endpoint: Optional[str] = None,
+ aws_external_id: Optional[str] = None,
+ **kwargs,
+ ):
+ """
+ Establish bidirectional streaming connection with Bedrock Nova Sonic.
+
+ Args:
+ model: Model ID (e.g., 'amazon.nova-sonic-v1:0')
+ websocket: Client WebSocket connection
+ logging_obj: LiteLLM logging object
+ aws_region_name: AWS region
+ Various AWS authentication parameters
+ """
+ try:
+ from aws_sdk_bedrock_runtime.client import (
+ BedrockRuntimeClient,
+ InvokeModelWithBidirectionalStreamOperationInput,
+ )
+ from aws_sdk_bedrock_runtime.config import Config
+ from smithy_aws_core.identity.environment import (
+ EnvironmentCredentialsResolver,
+ )
+ except ImportError:
+ raise ImportError(
+ "Missing aws_sdk_bedrock_runtime. Install with: pip install aws-sdk-bedrock-runtime"
+ )
+
+ # Get AWS region
+ if aws_region_name is None:
+ optional_params = {
+ "aws_region_name": aws_region_name,
+ }
+ aws_region_name = self._get_aws_region_name(optional_params, model)
+
+ # Get endpoint URL
+ if api_base is not None:
+ endpoint_uri = api_base
+ elif aws_bedrock_runtime_endpoint is not None:
+ endpoint_uri = aws_bedrock_runtime_endpoint
+ else:
+ endpoint_uri = f"https://bedrock-runtime.{aws_region_name}.amazonaws.com"
+
+ verbose_proxy_logger.debug(
+ f"Bedrock Realtime: Connecting to {endpoint_uri} with model {model}"
+ )
+
+ # Initialize Bedrock client with aws_sdk_bedrock_runtime
+ config = Config(
+ endpoint_uri=endpoint_uri,
+ region=aws_region_name,
+ aws_credentials_identity_resolver=EnvironmentCredentialsResolver(),
+ )
+ bedrock_client = BedrockRuntimeClient(config=config)
+
+ transformation_config = BedrockRealtimeConfig()
+
+ try:
+ # Initialize the bidirectional stream
+ bedrock_stream = await bedrock_client.invoke_model_with_bidirectional_stream(
+ InvokeModelWithBidirectionalStreamOperationInput(model_id=model)
+ )
+
+ verbose_proxy_logger.debug(
+ "Bedrock Realtime: Bidirectional stream established"
+ )
+
+ # Track state for transformation
+ session_state = {
+ "current_output_item_id": None,
+ "current_response_id": None,
+ "current_conversation_id": None,
+ "current_delta_chunks": None,
+ "current_item_chunks": None,
+ "current_delta_type": None,
+ "session_configuration_request": None,
+ }
+
+ # Create tasks for bidirectional forwarding
+ client_to_bedrock_task = asyncio.create_task(
+ self._forward_client_to_bedrock(
+ websocket,
+ bedrock_stream,
+ transformation_config,
+ model,
+ session_state,
+ )
+ )
+
+ bedrock_to_client_task = asyncio.create_task(
+ self._forward_bedrock_to_client(
+ bedrock_stream,
+ websocket,
+ transformation_config,
+ model,
+ logging_obj,
+ session_state,
+ )
+ )
+
+ # Wait for both tasks to complete
+ await asyncio.gather(
+ client_to_bedrock_task,
+ bedrock_to_client_task,
+ return_exceptions=True,
+ )
+
+ except Exception as e:
+ verbose_proxy_logger.exception(
+ f"Error in BedrockRealtime.async_realtime: {e}"
+ )
+ try:
+ await websocket.close(code=1011, reason=f"Internal error: {str(e)}")
+ except Exception:
+ pass
+ raise
+
+ async def _forward_client_to_bedrock(
+ self,
+ client_ws: Any,
+ bedrock_stream: Any,
+ transformation_config: BedrockRealtimeConfig,
+ model: str,
+ session_state: dict,
+ ):
+ """Forward messages from client WebSocket to Bedrock stream."""
+ try:
+ from aws_sdk_bedrock_runtime.models import (
+ BidirectionalInputPayloadPart,
+ InvokeModelWithBidirectionalStreamInputChunk,
+ )
+
+ while True:
+ # Receive message from client
+ message = await client_ws.receive_text()
+ verbose_proxy_logger.debug(
+ f"Bedrock Realtime: Received from client: {message[:200]}"
+ )
+
+ # Transform OpenAI format to Bedrock format
+ transformed_messages = transformation_config.transform_realtime_request(
+ message=message,
+ model=model,
+ session_configuration_request=session_state.get(
+ "session_configuration_request"
+ ),
+ )
+
+ # Send transformed messages to Bedrock
+ for bedrock_message in transformed_messages:
+ event = InvokeModelWithBidirectionalStreamInputChunk(
+ value=BidirectionalInputPayloadPart(
+ bytes_=bedrock_message.encode("utf-8")
+ )
+ )
+ await bedrock_stream.input_stream.send(event)
+ verbose_proxy_logger.debug(
+ f"Bedrock Realtime: Sent to Bedrock: {bedrock_message[:200]}"
+ )
+
+ except Exception as e:
+ verbose_proxy_logger.debug(
+ f"Client to Bedrock forwarding ended: {e}", exc_info=True
+ )
+ # Close the Bedrock stream input
+ try:
+ await bedrock_stream.input_stream.close()
+ except Exception:
+ pass
+
+ async def _forward_bedrock_to_client(
+ self,
+ bedrock_stream: Any,
+ client_ws: Any,
+ transformation_config: BedrockRealtimeConfig,
+ model: str,
+ logging_obj: LiteLLMLogging,
+ session_state: dict,
+ ):
+ """Forward messages from Bedrock stream to client WebSocket."""
+ try:
+ while True:
+ # Receive from Bedrock
+ output = await bedrock_stream.await_output()
+ result = await output[1].receive()
+
+ if result.value and result.value.bytes_:
+ bedrock_response = result.value.bytes_.decode("utf-8")
+ verbose_proxy_logger.debug(
+ f"Bedrock Realtime: Received from Bedrock: {bedrock_response[:200]}"
+ )
+
+ # Transform Bedrock format to OpenAI format
+ from litellm.types.realtime import RealtimeResponseTransformInput
+
+ realtime_response_transform_input: RealtimeResponseTransformInput = {
+ "current_output_item_id": session_state.get(
+ "current_output_item_id"
+ ),
+ "current_response_id": session_state.get("current_response_id"),
+ "current_conversation_id": session_state.get(
+ "current_conversation_id"
+ ),
+ "current_delta_chunks": session_state.get(
+ "current_delta_chunks"
+ ),
+ "current_item_chunks": session_state.get("current_item_chunks"),
+ "current_delta_type": session_state.get("current_delta_type"),
+ "session_configuration_request": session_state.get(
+ "session_configuration_request"
+ ),
+ }
+
+ transformed_response = (
+ transformation_config.transform_realtime_response(
+ message=bedrock_response,
+ model=model,
+ logging_obj=logging_obj,
+ realtime_response_transform_input=realtime_response_transform_input,
+ )
+ )
+
+ # Update session state
+ session_state.update(
+ {
+ "current_output_item_id": transformed_response.get(
+ "current_output_item_id"
+ ),
+ "current_response_id": transformed_response.get(
+ "current_response_id"
+ ),
+ "current_conversation_id": transformed_response.get(
+ "current_conversation_id"
+ ),
+ "current_delta_chunks": transformed_response.get(
+ "current_delta_chunks"
+ ),
+ "current_item_chunks": transformed_response.get(
+ "current_item_chunks"
+ ),
+ "current_delta_type": transformed_response.get(
+ "current_delta_type"
+ ),
+ "session_configuration_request": transformed_response.get(
+ "session_configuration_request"
+ ),
+ }
+ )
+
+ # Send transformed messages to client
+ openai_messages = transformed_response.get("response", [])
+ for openai_message in openai_messages:
+ message_json = json.dumps(openai_message)
+ await client_ws.send_text(message_json)
+ verbose_proxy_logger.debug(
+ f"Bedrock Realtime: Sent to client: {message_json[:200]}"
+ )
+
+ except Exception as e:
+ verbose_proxy_logger.debug(
+ f"Bedrock to client forwarding ended: {e}", exc_info=True
+ )
+ # Close the client WebSocket
+ try:
+ await client_ws.close()
+ except Exception:
+ pass
diff --git a/litellm/llms/bedrock/realtime/transformation.py b/litellm/llms/bedrock/realtime/transformation.py
new file mode 100644
index 00000000000..1dde1b47fe3
--- /dev/null
+++ b/litellm/llms/bedrock/realtime/transformation.py
@@ -0,0 +1,1156 @@
+"""
+This file contains the transformation logic for Bedrock Nova Sonic realtime API.
+
+Transforms between OpenAI Realtime API format and Bedrock Nova Sonic format.
+"""
+
+import json
+import uuid as uuid_lib
+from typing import Any, List, Optional, Union
+
+from litellm._logging import verbose_logger
+from litellm._uuid import uuid
+from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
+from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
+from litellm.types.llms.openai import (
+ OpenAIRealtimeContentPartDone,
+ OpenAIRealtimeDoneEvent,
+ OpenAIRealtimeEvents,
+ OpenAIRealtimeOutputItemDone,
+ OpenAIRealtimeResponseAudioDone,
+ OpenAIRealtimeResponseContentPartAdded,
+ OpenAIRealtimeResponseDelta,
+ OpenAIRealtimeResponseDoneObject,
+ OpenAIRealtimeResponseTextDone,
+ OpenAIRealtimeStreamResponseBaseObject,
+ OpenAIRealtimeStreamResponseOutputItemAdded,
+ OpenAIRealtimeStreamSession,
+ OpenAIRealtimeStreamSessionEvents,
+)
+from litellm.types.realtime import (
+ ALL_DELTA_TYPES,
+ RealtimeResponseTransformInput,
+ RealtimeResponseTypedDict,
+)
+from litellm.utils import get_empty_usage
+
+
+class BedrockRealtimeConfig(BaseRealtimeConfig):
+ """Configuration for Bedrock Nova Sonic realtime transformations."""
+
+ def __init__(self):
+ # Track session state
+ self.prompt_name = str(uuid_lib.uuid4())
+ self.content_name = str(uuid_lib.uuid4())
+ self.audio_content_name = str(uuid_lib.uuid4())
+
+ # Default configuration values
+ # Inference configuration
+ self.max_tokens = 1024
+ self.top_p = 0.9
+ self.temperature = 0.7
+
+ # Audio output configuration
+ self.output_sample_rate_hertz = 24000
+ self.output_sample_size_bits = 16
+ self.output_channel_count = 1
+ self.voice_id = "matthew"
+ self.output_encoding = "base64"
+ self.output_audio_type = "SPEECH"
+ self.output_media_type = "audio/lpcm"
+
+ # Audio input configuration
+ self.input_sample_rate_hertz = 16000
+ self.input_sample_size_bits = 16
+ self.input_channel_count = 1
+ self.input_encoding = "base64"
+ self.input_audio_type = "SPEECH"
+ self.input_media_type = "audio/lpcm"
+
+ # Text configuration
+ self.text_media_type = "text/plain"
+
+ def validate_environment(
+ self, headers: dict, model: str, api_key: Optional[str] = None
+ ) -> dict:
+ """Validate environment - no special validation needed for Bedrock."""
+ return headers
+
+ def get_complete_url(
+ self, api_base: Optional[str], model: str, api_key: Optional[str] = None
+ ) -> str:
+ """Get complete URL - handled by aws_sdk_bedrock_runtime."""
+ return api_base or ""
+
+ def requires_session_configuration(self) -> bool:
+ """Bedrock requires session configuration."""
+ return True
+
+ def session_configuration_request(self, model: str, tools: Optional[List[dict]] = None) -> str:
+ """
+ Create initial session configuration for Bedrock Nova Sonic.
+
+ Args:
+ model: Model ID
+ tools: Optional list of tool definitions
+
+ Returns JSON string with session start and prompt start events.
+ """
+ session_start = {
+ "event": {
+ "sessionStart": {
+ "inferenceConfiguration": {
+ "maxTokens": self.max_tokens,
+ "topP": self.top_p,
+ "temperature": self.temperature,
+ }
+ }
+ }
+ }
+
+ prompt_start_config = {
+ "promptName": self.prompt_name,
+ "textOutputConfiguration": {"mediaType": self.text_media_type},
+ "audioOutputConfiguration": {
+ "mediaType": self.output_media_type,
+ "sampleRateHertz": self.output_sample_rate_hertz,
+ "sampleSizeBits": self.output_sample_size_bits,
+ "channelCount": self.output_channel_count,
+ "voiceId": self.voice_id,
+ "encoding": self.output_encoding,
+ "audioType": self.output_audio_type,
+ },
+ }
+
+ # Add tool configuration if tools are provided
+ if tools:
+ prompt_start_config["toolUseOutputConfiguration"] = {
+ "mediaType": "application/json"
+ }
+ prompt_start_config["toolConfiguration"] = {
+ "tools": self._transform_tools_to_bedrock_format(tools)
+ }
+
+ prompt_start = {"event": {"promptStart": prompt_start_config}}
+
+ # Return as a marker that we've sent the configuration
+ return json.dumps(
+ {"session_start": session_start, "prompt_start": prompt_start}
+ )
+
+ def _transform_tools_to_bedrock_format(self, tools: List[dict]) -> List[dict]:
+ """
+ Transform OpenAI tool format to Bedrock tool format.
+
+ Args:
+ tools: List of OpenAI format tools
+
+ Returns:
+ List of Bedrock format tools
+ """
+ bedrock_tools = []
+ for tool in tools:
+ if tool.get("type") == "function":
+ function = tool.get("function", {})
+ bedrock_tool = {
+ "toolSpec": {
+ "name": function.get("name", ""),
+ "description": function.get("description", ""),
+ "inputSchema": {
+ "json": json.dumps(function.get("parameters", {}))
+ }
+ }
+ }
+ bedrock_tools.append(bedrock_tool)
+ return bedrock_tools
+
+ def _map_audio_format_to_sample_rate(self, audio_format: str, is_output: bool = True) -> int:
+ """
+ Map OpenAI audio format to sample rate.
+
+ Args:
+ audio_format: OpenAI audio format (pcm16, g711_ulaw, g711_alaw)
+ is_output: Whether this is for output (True) or input (False)
+
+ Returns:
+ Sample rate in Hz
+ """
+ # OpenAI uses 24kHz for output and can vary for input
+ # Bedrock Nova Sonic uses 24kHz for output and 16kHz for input by default
+ if audio_format == "pcm16":
+ return 24000 if is_output else 16000
+ elif audio_format in ["g711_ulaw", "g711_alaw"]:
+ return 8000 # G.711 typically uses 8kHz
+ return 24000 if is_output else 16000
+
+ def transform_session_update_event(self, json_message: dict) -> List[str]:
+ """
+ Transform session.update event to Bedrock session configuration.
+
+ Args:
+ json_message: OpenAI session.update message
+
+ Returns:
+ List of Bedrock format messages (JSON strings)
+ """
+ verbose_logger.debug("Handling session.update")
+ messages: List[str] = []
+
+ session_config = json_message.get("session", {})
+
+ # Update inference configuration from session if provided
+ if "max_response_output_tokens" in session_config:
+ self.max_tokens = session_config["max_response_output_tokens"]
+ if "temperature" in session_config:
+ self.temperature = session_config["temperature"]
+
+ # Update audio output configuration from session if provided
+ if "voice" in session_config:
+ self.voice_id = session_config["voice"]
+ if "output_audio_format" in session_config:
+ output_format = session_config["output_audio_format"]
+ self.output_sample_rate_hertz = self._map_audio_format_to_sample_rate(
+ output_format, is_output=True
+ )
+
+ # Update audio input configuration from session if provided
+ if "input_audio_format" in session_config:
+ input_format = session_config["input_audio_format"]
+ self.input_sample_rate_hertz = self._map_audio_format_to_sample_rate(
+ input_format, is_output=False
+ )
+
+ # Allow direct override of sample rates if provided (custom extension)
+ if "output_sample_rate_hertz" in session_config:
+ self.output_sample_rate_hertz = session_config["output_sample_rate_hertz"]
+ if "input_sample_rate_hertz" in session_config:
+ self.input_sample_rate_hertz = session_config["input_sample_rate_hertz"]
+
+ # Send session start
+ session_start = {
+ "event": {
+ "sessionStart": {
+ "inferenceConfiguration": {
+ "maxTokens": self.max_tokens,
+ "topP": self.top_p,
+ "temperature": self.temperature,
+ }
+ }
+ }
+ }
+ messages.append(json.dumps(session_start))
+
+ # Send prompt start
+ prompt_start_config = {
+ "promptName": self.prompt_name,
+ "textOutputConfiguration": {"mediaType": self.text_media_type},
+ "audioOutputConfiguration": {
+ "mediaType": self.output_media_type,
+ "sampleRateHertz": self.output_sample_rate_hertz,
+ "sampleSizeBits": self.output_sample_size_bits,
+ "channelCount": self.output_channel_count,
+ "voiceId": self.voice_id,
+ "encoding": self.output_encoding,
+ "audioType": self.output_audio_type,
+ },
+ }
+
+ # Add tool configuration if tools are provided
+ tools = session_config.get("tools")
+ if tools:
+ prompt_start_config["toolUseOutputConfiguration"] = {
+ "mediaType": "application/json"
+ }
+ prompt_start_config["toolConfiguration"] = {
+ "tools": self._transform_tools_to_bedrock_format(tools)
+ }
+
+ prompt_start = {"event": {"promptStart": prompt_start_config}}
+ messages.append(json.dumps(prompt_start))
+
+ # Send system prompt if provided
+ instructions = session_config.get("instructions")
+ if instructions:
+ text_content_name = str(uuid_lib.uuid4())
+
+ # Content start
+ text_content_start = {
+ "event": {
+ "contentStart": {
+ "promptName": self.prompt_name,
+ "contentName": text_content_name,
+ "type": "TEXT",
+ "interactive": False,
+ "role": "SYSTEM",
+ "textInputConfiguration": {"mediaType": self.text_media_type},
+ }
+ }
+ }
+ messages.append(json.dumps(text_content_start))
+
+ # Text input
+ text_input = {
+ "event": {
+ "textInput": {
+ "promptName": self.prompt_name,
+ "contentName": text_content_name,
+ "content": instructions,
+ }
+ }
+ }
+ messages.append(json.dumps(text_input))
+
+ # Content end
+ text_content_end = {
+ "event": {
+ "contentEnd": {
+ "promptName": self.prompt_name,
+ "contentName": text_content_name,
+ }
+ }
+ }
+ messages.append(json.dumps(text_content_end))
+
+ return messages
+
+ def transform_input_audio_buffer_append_event(self, json_message: dict) -> List[str]:
+ """
+ Transform input_audio_buffer.append event to Bedrock audio input.
+
+ Args:
+ json_message: OpenAI input_audio_buffer.append message
+
+ Returns:
+ List of Bedrock format messages (JSON strings)
+ """
+ verbose_logger.debug("Handling input_audio_buffer.append")
+ messages: List[str] = []
+
+ # Check if we need to start audio content
+ if not hasattr(self, "_audio_content_started"):
+ audio_content_start = {
+ "event": {
+ "contentStart": {
+ "promptName": self.prompt_name,
+ "contentName": self.audio_content_name,
+ "type": "AUDIO",
+ "interactive": True,
+ "role": "USER",
+ "audioInputConfiguration": {
+ "mediaType": self.input_media_type,
+ "sampleRateHertz": self.input_sample_rate_hertz,
+ "sampleSizeBits": self.input_sample_size_bits,
+ "channelCount": self.input_channel_count,
+ "audioType": self.input_audio_type,
+ "encoding": self.input_encoding,
+ },
+ }
+ }
+ }
+ messages.append(json.dumps(audio_content_start))
+ self._audio_content_started = True
+
+ # Send audio chunk
+ audio_data = json_message.get("audio", "")
+ audio_event = {
+ "event": {
+ "audioInput": {
+ "promptName": self.prompt_name,
+ "contentName": self.audio_content_name,
+ "content": audio_data,
+ }
+ }
+ }
+ messages.append(json.dumps(audio_event))
+
+ return messages
+
+ def transform_input_audio_buffer_commit_event(self, json_message: dict) -> List[str]:
+ """
+ Transform input_audio_buffer.commit event to Bedrock audio content end.
+
+ Args:
+ json_message: OpenAI input_audio_buffer.commit message
+
+ Returns:
+ List of Bedrock format messages (JSON strings)
+ """
+ verbose_logger.debug("Handling input_audio_buffer.commit")
+ messages: List[str] = []
+
+ if hasattr(self, "_audio_content_started"):
+ audio_content_end = {
+ "event": {
+ "contentEnd": {
+ "promptName": self.prompt_name,
+ "contentName": self.audio_content_name,
+ }
+ }
+ }
+ messages.append(json.dumps(audio_content_end))
+ delattr(self, "_audio_content_started")
+
+ return messages
+
+ def transform_conversation_item_create_event(self, json_message: dict) -> List[str]:
+ """
+ Transform conversation.item.create event to Bedrock text input or tool result.
+
+ Args:
+ json_message: OpenAI conversation.item.create message
+
+ Returns:
+ List of Bedrock format messages (JSON strings)
+ """
+ verbose_logger.debug("Handling conversation.item.create")
+ messages: List[str] = []
+
+ item = json_message.get("item", {})
+ item_type = item.get("type")
+
+ # Handle tool result
+ if item_type == "function_call_output":
+ return self.transform_conversation_item_create_tool_result_event(json_message)
+
+ # Handle regular message
+ if item_type == "message":
+ content = item.get("content", [])
+ for content_part in content:
+ if content_part.get("type") == "input_text":
+ text_content_name = str(uuid_lib.uuid4())
+
+ # Content start
+ text_content_start = {
+ "event": {
+ "contentStart": {
+ "promptName": self.prompt_name,
+ "contentName": text_content_name,
+ "type": "TEXT",
+ "interactive": True,
+ "role": "USER",
+ "textInputConfiguration": {
+ "mediaType": self.text_media_type
+ },
+ }
+ }
+ }
+ messages.append(json.dumps(text_content_start))
+
+ # Text input
+ text_input = {
+ "event": {
+ "textInput": {
+ "promptName": self.prompt_name,
+ "contentName": text_content_name,
+ "content": content_part.get("text", ""),
+ }
+ }
+ }
+ messages.append(json.dumps(text_input))
+
+ # Content end
+ text_content_end = {
+ "event": {
+ "contentEnd": {
+ "promptName": self.prompt_name,
+ "contentName": text_content_name,
+ }
+ }
+ }
+ messages.append(json.dumps(text_content_end))
+
+ return messages
+
+ def transform_response_create_event(self, json_message: dict) -> List[str]:
+ """
+ Transform response.create event to Bedrock format.
+
+ Args:
+ json_message: OpenAI response.create message
+
+ Returns:
+ List of Bedrock format messages (JSON strings)
+ """
+ verbose_logger.debug("Handling response.create")
+ # Bedrock starts generating automatically, no explicit trigger needed
+ return []
+
+ def transform_response_cancel_event(self, json_message: dict) -> List[str]:
+ """
+ Transform response.cancel event to Bedrock format.
+
+ Args:
+ json_message: OpenAI response.cancel message
+
+ Returns:
+ List of Bedrock format messages (JSON strings)
+ """
+ verbose_logger.debug("Handling response.cancel")
+ # Send interrupt signal if needed
+ return []
+
+ def transform_realtime_request(
+ self,
+ message: str,
+ model: str,
+ session_configuration_request: Optional[str] = None,
+ ) -> List[str]:
+ """
+ Transform OpenAI realtime request to Bedrock Nova Sonic format.
+
+ Args:
+ message: OpenAI format message (JSON string)
+ model: Model ID
+ session_configuration_request: Previous session config
+
+ Returns:
+ List of Bedrock format messages (JSON strings)
+ """
+ try:
+ json_message = json.loads(message)
+ except json.JSONDecodeError:
+ verbose_logger.warning(f"Invalid JSON message: {message[:200]}")
+ return []
+
+ message_type = json_message.get("type")
+
+ # Route to appropriate transformation method
+ if message_type == "session.update":
+ return self.transform_session_update_event(json_message)
+ elif message_type == "input_audio_buffer.append":
+ return self.transform_input_audio_buffer_append_event(json_message)
+ elif message_type == "input_audio_buffer.commit":
+ return self.transform_input_audio_buffer_commit_event(json_message)
+ elif message_type == "conversation.item.create":
+ return self.transform_conversation_item_create_event(json_message)
+ elif message_type == "response.create":
+ return self.transform_response_create_event(json_message)
+ elif message_type == "response.cancel":
+ return self.transform_response_cancel_event(json_message)
+ else:
+ verbose_logger.warning(f"Unknown message type: {message_type}")
+ return []
+
+ def transform_session_start_event(
+ self,
+ event: dict,
+ model: str,
+ logging_obj: LiteLLMLoggingObj,
+ ) -> OpenAIRealtimeStreamSessionEvents:
+ """
+ Transform Bedrock sessionStart event to OpenAI session.created.
+
+ Args:
+ event: Bedrock sessionStart event
+ model: Model ID
+ logging_obj: Logging object
+
+ Returns:
+ OpenAI session.created event
+ """
+ verbose_logger.debug("Handling sessionStart")
+
+ session = OpenAIRealtimeStreamSession(
+ id=logging_obj.litellm_trace_id,
+ modalities=["text", "audio"],
+ )
+ if model is not None and isinstance(model, str):
+ session["model"] = model
+
+ return OpenAIRealtimeStreamSessionEvents(
+ type="session.created",
+ session=session,
+ event_id=str(uuid.uuid4()),
+ )
+
+ def transform_content_start_event(
+ self,
+ event: dict,
+ current_response_id: Optional[str],
+ current_output_item_id: Optional[str],
+ current_conversation_id: Optional[str],
+ ) -> tuple[
+ List[OpenAIRealtimeEvents],
+ Optional[str],
+ Optional[str],
+ Optional[str],
+ Optional[ALL_DELTA_TYPES],
+ ]:
+ """
+ Transform Bedrock contentStart event to OpenAI response events.
+
+ Args:
+ event: Bedrock contentStart event
+ current_response_id: Current response ID
+ current_output_item_id: Current output item ID
+ current_conversation_id: Current conversation ID
+
+ Returns:
+ Tuple of (events, response_id, output_item_id, conversation_id, delta_type)
+ """
+ content_start = event["contentStart"]
+ role = content_start.get("role")
+
+ if role != "ASSISTANT":
+ return [], current_response_id, current_output_item_id, current_conversation_id, None
+
+ verbose_logger.debug("Handling ASSISTANT contentStart")
+
+ # Initialize IDs if needed
+ if not current_response_id:
+ current_response_id = f"resp_{uuid.uuid4()}"
+ if not current_output_item_id:
+ current_output_item_id = f"item_{uuid.uuid4()}"
+ if not current_conversation_id:
+ current_conversation_id = f"conv_{uuid.uuid4()}"
+
+ # Determine content type
+ content_type = content_start.get("type", "TEXT")
+ current_delta_type: ALL_DELTA_TYPES = "text" if content_type == "TEXT" else "audio"
+
+ returned_messages: List[OpenAIRealtimeEvents] = []
+
+ # Send response.created
+ response_created = OpenAIRealtimeStreamResponseBaseObject(
+ type="response.created",
+ event_id=f"event_{uuid.uuid4()}",
+ response={
+ "object": "realtime.response",
+ "id": current_response_id,
+ "status": "in_progress",
+ "output": [],
+ "conversation_id": current_conversation_id,
+ },
+ )
+ returned_messages.append(response_created)
+
+ # Send response.output_item.added
+ output_item_added = OpenAIRealtimeStreamResponseOutputItemAdded(
+ type="response.output_item.added",
+ response_id=current_response_id,
+ output_index=0,
+ item={
+ "id": current_output_item_id,
+ "object": "realtime.item",
+ "type": "message",
+ "status": "in_progress",
+ "role": "assistant",
+ "content": [],
+ },
+ )
+ returned_messages.append(output_item_added)
+
+ # Send response.content_part.added
+ content_part_added = OpenAIRealtimeResponseContentPartAdded(
+ type="response.content_part.added",
+ content_index=0,
+ output_index=0,
+ event_id=f"event_{uuid.uuid4()}",
+ item_id=current_output_item_id,
+ part=(
+ {"type": "text", "text": ""}
+ if current_delta_type == "text"
+ else {"type": "audio", "transcript": ""}
+ ),
+ response_id=current_response_id,
+ )
+ returned_messages.append(content_part_added)
+
+ return (
+ returned_messages,
+ current_response_id,
+ current_output_item_id,
+ current_conversation_id,
+ current_delta_type,
+ )
+
+ def transform_text_output_event(
+ self,
+ event: dict,
+ current_output_item_id: Optional[str],
+ current_response_id: Optional[str],
+ current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]],
+ ) -> tuple[List[OpenAIRealtimeEvents], Optional[List[OpenAIRealtimeResponseDelta]]]:
+ """
+ Transform Bedrock textOutput event to OpenAI response.text.delta.
+
+ Args:
+ event: Bedrock textOutput event
+ current_output_item_id: Current output item ID
+ current_response_id: Current response ID
+ current_delta_chunks: Current delta chunks
+
+ Returns:
+ Tuple of (events, updated_delta_chunks)
+ """
+ verbose_logger.debug("Handling textOutput")
+ text_content = event["textOutput"].get("content", "")
+
+ if not current_output_item_id or not current_response_id:
+ return [], current_delta_chunks
+
+ text_delta = OpenAIRealtimeResponseDelta(
+ type="response.text.delta",
+ content_index=0,
+ event_id=f"event_{uuid.uuid4()}",
+ item_id=current_output_item_id,
+ output_index=0,
+ response_id=current_response_id,
+ delta=text_content,
+ )
+
+ # Track delta chunks
+ if current_delta_chunks is None:
+ current_delta_chunks = []
+ current_delta_chunks.append(text_delta)
+
+ return [text_delta], current_delta_chunks
+
+ def transform_audio_output_event(
+ self,
+ event: dict,
+ current_output_item_id: Optional[str],
+ current_response_id: Optional[str],
+ ) -> List[OpenAIRealtimeEvents]:
+ """
+ Transform Bedrock audioOutput event to OpenAI response.audio.delta.
+
+ Args:
+ event: Bedrock audioOutput event
+ current_output_item_id: Current output item ID
+ current_response_id: Current response ID
+
+ Returns:
+ List of OpenAI events
+ """
+ verbose_logger.debug("Handling audioOutput")
+ audio_content = event["audioOutput"].get("content", "")
+
+ if not current_output_item_id or not current_response_id:
+ return []
+
+ audio_delta = OpenAIRealtimeResponseDelta(
+ type="response.audio.delta",
+ content_index=0,
+ event_id=f"event_{uuid.uuid4()}",
+ item_id=current_output_item_id,
+ output_index=0,
+ response_id=current_response_id,
+ delta=audio_content,
+ )
+
+ return [audio_delta]
+
+ def transform_content_end_event(
+ self,
+ event: dict,
+ current_output_item_id: Optional[str],
+ current_response_id: Optional[str],
+ current_delta_type: Optional[str],
+ current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]],
+ ) -> tuple[List[OpenAIRealtimeEvents], Optional[List[OpenAIRealtimeResponseDelta]]]:
+ """
+ Transform Bedrock contentEnd event to OpenAI response done events.
+
+ Args:
+ event: Bedrock contentEnd event
+ current_output_item_id: Current output item ID
+ current_response_id: Current response ID
+ current_delta_type: Current delta type (text or audio)
+ current_delta_chunks: Current delta chunks
+
+ Returns:
+ Tuple of (events, reset_delta_chunks)
+ """
+ content_end = event["contentEnd"]
+ verbose_logger.debug(f"Handling contentEnd: {content_end}")
+
+ if not current_output_item_id or not current_response_id:
+ return [], current_delta_chunks
+
+ returned_messages: List[OpenAIRealtimeEvents] = []
+
+ # Send appropriate done event based on type
+ if current_delta_type == "text":
+ # Accumulate text
+ accumulated_text = ""
+ if current_delta_chunks:
+ accumulated_text = "".join(
+ [chunk.get("delta", "") for chunk in current_delta_chunks]
+ )
+
+ text_done = OpenAIRealtimeResponseTextDone(
+ type="response.text.done",
+ content_index=0,
+ event_id=f"event_{uuid.uuid4()}",
+ item_id=current_output_item_id,
+ output_index=0,
+ response_id=current_response_id,
+ text=accumulated_text,
+ )
+ returned_messages.append(text_done)
+
+ # Send content_part.done
+ content_part_done = OpenAIRealtimeContentPartDone(
+ type="response.content_part.done",
+ content_index=0,
+ event_id=f"event_{uuid.uuid4()}",
+ item_id=current_output_item_id,
+ output_index=0,
+ part={"type": "text", "text": accumulated_text},
+ response_id=current_response_id,
+ )
+ returned_messages.append(content_part_done)
+
+ elif current_delta_type == "audio":
+ audio_done = OpenAIRealtimeResponseAudioDone(
+ type="response.audio.done",
+ content_index=0,
+ event_id=f"event_{uuid.uuid4()}",
+ item_id=current_output_item_id,
+ output_index=0,
+ response_id=current_response_id,
+ )
+ returned_messages.append(audio_done)
+
+ # Send content_part.done
+ content_part_done = OpenAIRealtimeContentPartDone(
+ type="response.content_part.done",
+ content_index=0,
+ event_id=f"event_{uuid.uuid4()}",
+ item_id=current_output_item_id,
+ output_index=0,
+ part={"type": "audio", "transcript": ""},
+ response_id=current_response_id,
+ )
+ returned_messages.append(content_part_done)
+
+ # Send output_item.done
+ output_item_done = OpenAIRealtimeOutputItemDone(
+ type="response.output_item.done",
+ event_id=f"event_{uuid.uuid4()}",
+ output_index=0,
+ response_id=current_response_id,
+ item={
+ "id": current_output_item_id,
+ "object": "realtime.item",
+ "type": "message",
+ "status": "completed",
+ "role": "assistant",
+ "content": [],
+ },
+ )
+ returned_messages.append(output_item_done)
+
+ # Reset delta chunks
+ return returned_messages, None
+
+ def transform_prompt_end_event(
+ self,
+ event: dict,
+ current_response_id: Optional[str],
+ current_conversation_id: Optional[str],
+ ) -> tuple[List[OpenAIRealtimeEvents], Optional[str], Optional[str], Optional[ALL_DELTA_TYPES]]:
+ """
+ Transform Bedrock promptEnd event to OpenAI response.done.
+
+ Args:
+ event: Bedrock promptEnd event
+ current_response_id: Current response ID
+ current_conversation_id: Current conversation ID
+
+ Returns:
+ Tuple of (events, reset_output_item_id, reset_response_id, reset_delta_type)
+ """
+ verbose_logger.debug("Handling promptEnd")
+
+ if not current_response_id or not current_conversation_id:
+ return [], None, None, None
+
+ usage_obj = get_empty_usage()
+ response_done = OpenAIRealtimeDoneEvent(
+ type="response.done",
+ event_id=f"event_{uuid.uuid4()}",
+ response=OpenAIRealtimeResponseDoneObject(
+ object="realtime.response",
+ id=current_response_id,
+ status="completed",
+ output=[],
+ conversation_id=current_conversation_id,
+ usage={
+ "prompt_tokens": usage_obj.prompt_tokens,
+ "completion_tokens": usage_obj.completion_tokens,
+ "total_tokens": usage_obj.total_tokens,
+ },
+ ),
+ )
+
+ # Reset state for next response
+ return [response_done], None, None, None
+
+ def transform_tool_use_event(
+ self,
+ event: dict,
+ current_output_item_id: Optional[str],
+ current_response_id: Optional[str],
+ ) -> tuple[List[OpenAIRealtimeEvents], str, str]:
+ """
+ Transform Bedrock toolUse event to OpenAI format.
+
+ Args:
+ event: Bedrock toolUse event
+ current_output_item_id: Current output item ID
+ current_response_id: Current response ID
+
+ Returns:
+ Tuple of (events, tool_call_id, tool_name) for tracking
+ """
+ verbose_logger.debug("Handling toolUse")
+ tool_use = event["toolUse"]
+
+ if not current_output_item_id or not current_response_id:
+ return [], "", ""
+
+ # Parse the tool input
+ tool_input = {}
+ if "input" in tool_use:
+ try:
+ tool_input = json.loads(tool_use["input"]) if isinstance(tool_use["input"], str) else tool_use["input"]
+ except json.JSONDecodeError:
+ tool_input = {}
+
+ tool_call_id = tool_use.get("toolUseId", "")
+ tool_name = tool_use.get("toolName", "")
+
+ # Create a function call arguments done event
+ # This is a custom event format that matches what clients expect
+ from typing import cast
+ function_call_event: dict[str, Any] = {
+ "type": "response.function_call_arguments.done",
+ "event_id": f"event_{uuid.uuid4()}",
+ "response_id": current_response_id,
+ "item_id": current_output_item_id,
+ "output_index": 0,
+ "call_id": tool_call_id,
+ "name": tool_name,
+ "arguments": json.dumps(tool_input),
+ }
+
+ return [cast(OpenAIRealtimeEvents, function_call_event)], tool_call_id, tool_name
+
+ def transform_conversation_item_create_tool_result_event(self, json_message: dict) -> List[str]:
+ """
+ Transform conversation.item.create with tool result to Bedrock format.
+
+ Args:
+ json_message: OpenAI conversation.item.create message with tool result
+
+ Returns:
+ List of Bedrock format messages (JSON strings)
+ """
+ verbose_logger.debug("Handling conversation.item.create for tool result")
+ messages: List[str] = []
+
+ item = json_message.get("item", {})
+ if item.get("type") == "function_call_output":
+ tool_content_name = str(uuid_lib.uuid4())
+ call_id = item.get("call_id", "")
+ output = item.get("output", "")
+
+ # Content start for tool result
+ tool_content_start = {
+ "event": {
+ "contentStart": {
+ "promptName": self.prompt_name,
+ "contentName": tool_content_name,
+ "interactive": False,
+ "type": "TOOL",
+ "role": "TOOL",
+ "toolResultInputConfiguration": {
+ "toolUseId": call_id,
+ "type": "TEXT",
+ "textInputConfiguration": {
+ "mediaType": "text/plain"
+ }
+ }
+ }
+ }
+ }
+ messages.append(json.dumps(tool_content_start))
+
+ # Tool result
+ tool_result = {
+ "event": {
+ "toolResult": {
+ "promptName": self.prompt_name,
+ "contentName": tool_content_name,
+ "content": output if isinstance(output, str) else json.dumps(output)
+ }
+ }
+ }
+ messages.append(json.dumps(tool_result))
+
+ # Content end
+ tool_content_end = {
+ "event": {
+ "contentEnd": {
+ "promptName": self.prompt_name,
+ "contentName": tool_content_name,
+ }
+ }
+ }
+ messages.append(json.dumps(tool_content_end))
+
+ return messages
+
+ def transform_realtime_response(
+ self,
+ message: Union[str, bytes],
+ model: str,
+ logging_obj: LiteLLMLoggingObj,
+ realtime_response_transform_input: RealtimeResponseTransformInput,
+ ) -> RealtimeResponseTypedDict:
+ """
+ Transform Bedrock Nova Sonic response to OpenAI realtime format.
+
+ Args:
+ message: Bedrock format message (JSON string)
+ model: Model ID
+ logging_obj: Logging object
+ realtime_response_transform_input: Current state
+
+ Returns:
+ Transformed response with updated state
+ """
+ try:
+ json_message = json.loads(message)
+ except json.JSONDecodeError:
+ message_preview = message[:200].decode('utf-8', errors='replace') if isinstance(message, bytes) else message[:200]
+ verbose_logger.warning(f"Invalid JSON message: {message_preview}")
+ return {
+ "response": [],
+ "current_output_item_id": realtime_response_transform_input.get(
+ "current_output_item_id"
+ ),
+ "current_response_id": realtime_response_transform_input.get(
+ "current_response_id"
+ ),
+ "current_delta_chunks": realtime_response_transform_input.get(
+ "current_delta_chunks"
+ ),
+ "current_conversation_id": realtime_response_transform_input.get(
+ "current_conversation_id"
+ ),
+ "current_item_chunks": realtime_response_transform_input.get(
+ "current_item_chunks"
+ ),
+ "current_delta_type": realtime_response_transform_input.get(
+ "current_delta_type"
+ ),
+ "session_configuration_request": realtime_response_transform_input.get(
+ "session_configuration_request"
+ ),
+ }
+
+ # Extract state
+ current_output_item_id = realtime_response_transform_input.get(
+ "current_output_item_id"
+ )
+ current_response_id = realtime_response_transform_input.get(
+ "current_response_id"
+ )
+ current_conversation_id = realtime_response_transform_input.get(
+ "current_conversation_id"
+ )
+ current_delta_chunks = realtime_response_transform_input.get(
+ "current_delta_chunks"
+ )
+ current_delta_type = realtime_response_transform_input.get("current_delta_type")
+ session_configuration_request = realtime_response_transform_input.get(
+ "session_configuration_request"
+ )
+
+ returned_messages: List[OpenAIRealtimeEvents] = []
+
+ # Parse Bedrock event
+ event = json_message.get("event", {})
+
+ # Route to appropriate transformation method
+ if "sessionStart" in event:
+ session_created = self.transform_session_start_event(
+ event, model, logging_obj
+ )
+ returned_messages.append(session_created)
+ session_configuration_request = json.dumps({"configured": True})
+
+ elif "contentStart" in event:
+ (
+ events,
+ current_response_id,
+ current_output_item_id,
+ current_conversation_id,
+ current_delta_type,
+ ) = self.transform_content_start_event(
+ event,
+ current_response_id,
+ current_output_item_id,
+ current_conversation_id,
+ )
+ returned_messages.extend(events)
+
+ elif "textOutput" in event:
+ events, current_delta_chunks = self.transform_text_output_event(
+ event,
+ current_output_item_id,
+ current_response_id,
+ current_delta_chunks,
+ )
+ returned_messages.extend(events)
+
+ elif "audioOutput" in event:
+ events = self.transform_audio_output_event(
+ event, current_output_item_id, current_response_id
+ )
+ returned_messages.extend(events)
+
+ elif "contentEnd" in event:
+ events, current_delta_chunks = self.transform_content_end_event(
+ event,
+ current_output_item_id,
+ current_response_id,
+ current_delta_type,
+ current_delta_chunks,
+ )
+ returned_messages.extend(events)
+
+ elif "toolUse" in event:
+ events, tool_call_id, tool_name = self.transform_tool_use_event(
+ event, current_output_item_id, current_response_id
+ )
+ returned_messages.extend(events)
+ # Store tool call info for potential use
+ verbose_logger.debug(f"Tool use event: {tool_name} (ID: {tool_call_id})")
+
+ elif "promptEnd" in event:
+ (
+ events,
+ current_output_item_id,
+ current_response_id,
+ current_delta_type,
+ ) = self.transform_prompt_end_event(
+ event, current_response_id, current_conversation_id
+ )
+ returned_messages.extend(events)
+
+ return {
+ "response": returned_messages,
+ "current_output_item_id": current_output_item_id,
+ "current_response_id": current_response_id,
+ "current_delta_chunks": current_delta_chunks,
+ "current_conversation_id": current_conversation_id,
+ "current_item_chunks": realtime_response_transform_input.get(
+ "current_item_chunks"
+ ),
+ "current_delta_type": current_delta_type,
+ "session_configuration_request": session_configuration_request,
+ }
diff --git a/litellm/llms/cerebras/chat.py b/litellm/llms/cerebras/chat.py
index 4e9c6811a77..9929e2ab9a2 100644
--- a/litellm/llms/cerebras/chat.py
+++ b/litellm/llms/cerebras/chat.py
@@ -7,6 +7,7 @@ this is OpenAI compatible - no translation needed / occurs
from typing import Optional
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
+from litellm.utils import supports_reasoning
class CerebrasConfig(OpenAIGPTConfig):
@@ -24,6 +25,7 @@ class CerebrasConfig(OpenAIGPTConfig):
tool_choice: Optional[str] = None
tools: Optional[list] = None
user: Optional[str] = None
+ reasoning_effort: Optional[str] = None
def __init__(
self,
@@ -37,6 +39,7 @@ class CerebrasConfig(OpenAIGPTConfig):
tool_choice: Optional[str] = None,
tools: Optional[list] = None,
user: Optional[str] = None,
+ reasoning_effort: Optional[str] = None,
) -> None:
locals_ = locals().copy()
for key, value in locals_.items():
@@ -53,7 +56,7 @@ class CerebrasConfig(OpenAIGPTConfig):
"""
- return [
+ supported_params = [
"max_tokens",
"max_completion_tokens",
"response_format",
@@ -67,6 +70,12 @@ class CerebrasConfig(OpenAIGPTConfig):
"user",
]
+ # Only add reasoning_effort for models that support it
+ if supports_reasoning(model=model, custom_llm_provider="cerebras"):
+ supported_params.append("reasoning_effort")
+
+ return supported_params
+
def map_openai_params(
self,
non_default_params: dict,
diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py
index 4f86877a6c0..ac9dd5998e2 100644
--- a/litellm/llms/custom_httpx/http_handler.py
+++ b/litellm/llms/custom_httpx/http_handler.py
@@ -50,9 +50,21 @@ try:
except Exception:
version = "0.0.0"
-headers = {
- "User-Agent": f"litellm/{version}",
-}
+def get_default_headers() -> dict:
+ """
+ Get default headers for HTTP requests.
+
+ - Default: `User-Agent: litellm/{version}`
+ - Override: set `LITELLM_USER_AGENT` to fully override the header value.
+ """
+ user_agent = os.environ.get("LITELLM_USER_AGENT")
+ if user_agent is not None:
+ return {"User-Agent": user_agent}
+
+ return {"User-Agent": f"litellm/{version}"}
+
+# Initialize headers (User-Agent)
+headers = get_default_headers()
# https://www.python-httpx.org/advanced/timeouts
_DEFAULT_TIMEOUT = httpx.Timeout(timeout=5.0, connect=5.0)
@@ -371,13 +383,16 @@ class AsyncHTTPHandler:
shared_session=shared_session,
)
+ # Get default headers (User-Agent, overridable via LITELLM_USER_AGENT)
+ default_headers = get_default_headers()
+
return httpx.AsyncClient(
transport=transport,
event_hooks=event_hooks,
timeout=timeout,
verify=ssl_config,
cert=cert,
- headers=headers,
+ headers=default_headers,
follow_redirects=True,
)
@@ -899,6 +914,9 @@ class HTTPHandler:
# /path/to/client.pem
cert = os.getenv("SSL_CERTIFICATE", litellm.ssl_certificate)
+ # Get default headers (User-Agent, overridable via LITELLM_USER_AGENT)
+ default_headers = get_default_headers() if not disable_default_headers else None
+
if client is None:
transport = self._create_sync_transport()
@@ -908,7 +926,7 @@ class HTTPHandler:
timeout=timeout,
verify=ssl_config,
cert=cert,
- headers=headers if not disable_default_headers else None,
+ headers=default_headers,
follow_redirects=True,
)
else:
diff --git a/litellm/llms/custom_httpx/httpx_handler.py b/litellm/llms/custom_httpx/httpx_handler.py
index 6f684ba01c2..491cd97f7db 100644
--- a/litellm/llms/custom_httpx/httpx_handler.py
+++ b/litellm/llms/custom_httpx/httpx_handler.py
@@ -1,3 +1,4 @@
+import os
from typing import Optional, Union
import httpx
@@ -7,13 +8,22 @@ try:
except Exception:
version = "0.0.0"
-headers = {
- "User-Agent": f"litellm/{version}",
-}
+def get_default_headers() -> dict:
+ """
+ Get default headers for HTTP requests.
+ - Default: `User-Agent: litellm/{version}`
+ - Override: set `LITELLM_USER_AGENT` to fully override the header value.
+ """
+ user_agent = os.environ.get("LITELLM_USER_AGENT")
+ if user_agent is not None:
+ return {"User-Agent": user_agent}
+
+ return {"User-Agent": f"litellm/{version}"}
class HTTPHandler:
def __init__(self, concurrent_limit=1000):
+ headers = get_default_headers()
# Create a client with a connection pool
self.client = httpx.AsyncClient(
limits=httpx.Limits(
diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py
index 2b7f5dd5995..e9ae94307d4 100644
--- a/litellm/llms/databricks/chat/transformation.py
+++ b/litellm/llms/databricks/chat/transformation.py
@@ -298,7 +298,8 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
if "reasoning_effort" in non_default_params and "claude" in model:
optional_params["thinking"] = AnthropicConfig._map_reasoning_effort(
- non_default_params.get("reasoning_effort")
+ reasoning_effort=non_default_params.get("reasoning_effort"),
+ model=model
)
optional_params.pop("reasoning_effort", None)
## handle thinking tokens
diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py
index 86bcd94450f..7ec32fecc46 100644
--- a/litellm/llms/fireworks_ai/chat/transformation.py
+++ b/litellm/llms/fireworks_ai/chat/transformation.py
@@ -236,6 +236,10 @@ class FireworksAIConfig(OpenAIGPTConfig):
disable_add_transform_inline_image_block=disable_add_transform_inline_image_block,
)
filter_value_from_dict(cast(dict, message), "cache_control")
+ # Remove fields not permitted by FireworksAI that may cause:
+ # "Not permitted, field: 'messages[n].provider_specific_fields'"
+ if isinstance(message, dict) and "provider_specific_fields" in message:
+ cast(dict, message).pop("provider_specific_fields", None)
return messages
diff --git a/litellm/llms/gemini/files/transformation.py b/litellm/llms/gemini/files/transformation.py
index 37f1376c2b1..cc799cfd6aa 100644
--- a/litellm/llms/gemini/files/transformation.py
+++ b/litellm/llms/gemini/files/transformation.py
@@ -210,7 +210,7 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
We expect file_id to be the URI (e.g. https://generativelanguage.googleapis.com/v1beta/files/...)
as returned by the upload response.
"""
- api_key = litellm_params.get("api_key")
+ api_key = litellm_params.get("api_key") or self.get_api_key()
if not api_key:
raise ValueError("api_key is required")
@@ -222,7 +222,8 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
api_base = api_base.rstrip("/")
url = "{}/v1beta/{}?key={}".format(api_base, file_id, api_key)
- return url, {"Content-Type": "application/json"}
+ # Return empty params dict - API key is already in URL, no query params needed
+ return url, {}
def transform_retrieve_file_response(
self,
@@ -299,7 +300,7 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
# Extract the file path from full URI
file_name = file_id.split("/v1beta/")[-1]
else:
- file_name = file_id
+ file_name = file_id if file_id.startswith("files/") else f"files/{file_id}"
# Construct the delete URL
url = f"{api_base}/v1beta/{file_name}"
diff --git a/litellm/llms/gemini/image_generation/transformation.py b/litellm/llms/gemini/image_generation/transformation.py
index 63b835df9d0..73aef15e4c7 100644
--- a/litellm/llms/gemini/image_generation/transformation.py
+++ b/litellm/llms/gemini/image_generation/transformation.py
@@ -255,9 +255,11 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
if "inlineData" in part:
inline_data = part["inlineData"]
if "data" in inline_data:
+ thought_sig = part.get("thoughtSignature")
model_response.data.append(ImageObject(
b64_json=inline_data["data"],
url=None,
+ provider_specific_fields={"thought_signature": thought_sig} if thought_sig else None,
))
# Extract usage metadata for Gemini models
diff --git a/litellm/llms/gigachat/chat/transformation.py b/litellm/llms/gigachat/chat/transformation.py
index ba14de1f65d..f546f356e11 100644
--- a/litellm/llms/gigachat/chat/transformation.py
+++ b/litellm/llms/gigachat/chat/transformation.py
@@ -386,33 +386,7 @@ class GigaChatConfig(BaseConfig):
transformed.append(message)
- # Collapse consecutive user messages
- return self._collapse_user_messages(transformed)
-
- def _collapse_user_messages(self, messages: List[dict]) -> List[dict]:
- """Collapse consecutive user messages into one."""
- collapsed: List[dict] = []
- prev_user_msg: Optional[dict] = None
- content_parts: List[str] = []
-
- for msg in messages:
- if msg.get("role") == "user" and prev_user_msg is not None:
- content_parts.append(msg.get("content", ""))
- else:
- if content_parts and prev_user_msg:
- prev_user_msg["content"] = "\n".join(
- [prev_user_msg.get("content", "")] + content_parts
- )
- content_parts = []
- collapsed.append(msg)
- prev_user_msg = msg if msg.get("role") == "user" else None
-
- if content_parts and prev_user_msg:
- prev_user_msg["content"] = "\n".join(
- [prev_user_msg.get("content", "")] + content_parts
- )
-
- return collapsed
+ return transformed
def transform_response(
self,
diff --git a/litellm/llms/github_copilot/chat/transformation.py b/litellm/llms/github_copilot/chat/transformation.py
index 50f18cedf9b..be8ad7d0877 100644
--- a/litellm/llms/github_copilot/chat/transformation.py
+++ b/litellm/llms/github_copilot/chat/transformation.py
@@ -1,11 +1,16 @@
-from typing import Any, Optional, Tuple, cast, List
+from typing import List, Optional, Tuple
+
from litellm.exceptions import AuthenticationError
from litellm.llms.openai.openai import OpenAIConfig
from litellm.types.llms.openai import AllMessageValues
from ..authenticator import Authenticator
-from ..common_utils import GetAPIKeyError, GITHUB_COPILOT_API_BASE
+from ..common_utils import (
+ GITHUB_COPILOT_API_BASE,
+ GetAPIKeyError,
+ get_copilot_default_headers,
+)
class GithubCopilotConfig(OpenAIConfig):
@@ -25,9 +30,7 @@ class GithubCopilotConfig(OpenAIConfig):
api_key: Optional[str],
custom_llm_provider: str,
) -> Tuple[Optional[str], Optional[str], str]:
- dynamic_api_base = (
- self.authenticator.get_api_base() or GITHUB_COPILOT_API_BASE
- )
+ dynamic_api_base = self.authenticator.get_api_base() or GITHUB_COPILOT_API_BASE
try:
dynamic_api_key = self.authenticator.get_api_key()
except GetAPIKeyError as e:
@@ -45,14 +48,24 @@ class GithubCopilotConfig(OpenAIConfig):
):
import litellm
- disable_copilot_system_to_assistant = (
- litellm.disable_copilot_system_to_assistant
- )
- if not disable_copilot_system_to_assistant:
- for message in messages:
- if "role" in message and message["role"] == "system":
- cast(Any, message)["role"] = "assistant"
- return messages
+ # Check if system-to-assistant conversion is disabled
+ if litellm.disable_copilot_system_to_assistant:
+ # GitHub Copilot API now supports system prompts for all models (Claude, GPT, etc.)
+ # No conversion needed - just return messages as-is
+ return messages
+
+ # Default behavior: convert system messages to assistant for compatibility
+ transformed_messages = []
+ for message in messages:
+ if message.get("role") == "system":
+ # Convert system message to assistant message
+ transformed_message = message.copy()
+ transformed_message["role"] = "assistant"
+ transformed_messages.append(transformed_message)
+ else:
+ transformed_messages.append(message)
+
+ return transformed_messages
def validate_environment(
self,
@@ -69,6 +82,14 @@ class GithubCopilotConfig(OpenAIConfig):
headers, model, messages, optional_params, litellm_params, api_key, api_base
)
+ # Add Copilot-specific headers (editor-version, user-agent, etc.)
+ try:
+ copilot_api_key = self.authenticator.get_api_key()
+ copilot_headers = get_copilot_default_headers(copilot_api_key)
+ validated_headers = {**copilot_headers, **validated_headers}
+ except GetAPIKeyError:
+ pass # Will be handled later in the request flow
+
# Add X-Initiator header based on message roles
initiator = self._determine_initiator(messages)
validated_headers["X-Initiator"] = initiator
@@ -87,7 +108,7 @@ class GithubCopilotConfig(OpenAIConfig):
For other models, returns standard OpenAI parameters (which may include reasoning_effort for o-series models).
"""
from litellm.utils import supports_reasoning
-
+
# Get base OpenAI parameters
base_params = super().get_supported_openai_params(model)
@@ -118,7 +139,7 @@ class GithubCopilotConfig(OpenAIConfig):
"""
Check if any message contains vision content (images).
Returns True if any message has content with vision-related types, otherwise False.
-
+
Checks for:
- image_url content type (OpenAI format)
- Content items with type 'image_url'
diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py
index fb00aa28f45..c406f502b45 100644
--- a/litellm/llms/openai/chat/guardrail_translation/handler.py
+++ b/litellm/llms/openai/chat/guardrail_translation/handler.py
@@ -21,7 +21,13 @@ from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.main import stream_chunk_builder
from litellm.types.llms.openai import ChatCompletionToolParam
-from litellm.types.utils import Choices, GenericGuardrailAPIInputs, ModelResponse, ModelResponseStream, StreamingChoices
+from litellm.types.utils import (
+ Choices,
+ GenericGuardrailAPIInputs,
+ ModelResponse,
+ ModelResponseStream,
+ StreamingChoices,
+)
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
@@ -80,9 +86,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
if tool_calls_to_check:
inputs["tool_calls"] = tool_calls_to_check # type: ignore
if messages:
- inputs["structured_messages"] = (
- messages # pass the openai /chat/completions messages to the guardrail, as-is
- )
+ inputs[
+ "structured_messages"
+ ] = messages # pass the openai /chat/completions messages to the guardrail, as-is
# Pass tools (function definitions) to the guardrail
tools = data.get("tools")
if tools:
@@ -362,14 +368,17 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
# check if the stream has ended
has_stream_ended = False
for chunk in responses_so_far:
- if chunk.choices[0].finish_reason is not None:
+ if chunk.choices and chunk.choices[0].finish_reason is not None:
has_stream_ended = True
break
if has_stream_ended:
# convert to model response
model_response = cast(
- ModelResponse, stream_chunk_builder(chunks=responses_so_far, logging_obj=litellm_logging_obj)
+ ModelResponse,
+ stream_chunk_builder(
+ chunks=responses_so_far, logging_obj=litellm_logging_obj
+ ),
)
# run process_output_response
await self.process_output_response(
diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py
index 8bcecd35232..ce470f04aca 100644
--- a/litellm/llms/openai/common_utils.py
+++ b/litellm/llms/openai/common_utils.py
@@ -15,14 +15,12 @@ if TYPE_CHECKING:
from aiohttp import ClientSession
import litellm
-from litellm._logging import verbose_logger
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.custom_httpx.http_handler import (
_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
AsyncHTTPHandler,
get_ssl_configuration,
)
-from litellm.types.utils import LlmProviders
class OpenAIError(BaseLLMException):
@@ -205,67 +203,30 @@ class BaseOpenAILLM:
if litellm.aclient_session is not None:
return litellm.aclient_session
- # Use the global cached client system to prevent memory leaks (issue #14540)
- # This routes through get_async_httpx_client() which provides TTL-based caching
- from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
+ # Get unified SSL configuration
+ ssl_config = get_ssl_configuration()
- try:
- # Get SSL config and include in params for proper cache key
- ssl_config = get_ssl_configuration()
- params = {"ssl_verify": ssl_config} if ssl_config is not None else {}
- params["disable_aiohttp_transport"] = litellm.disable_aiohttp_transport
-
- # Get a cached AsyncHTTPHandler which manages the httpx.AsyncClient
- cached_handler = get_async_httpx_client(
- llm_provider=LlmProviders.OPENAI, # Cache key includes provider
- params=params, # Include SSL config in cache key
+ return httpx.AsyncClient(
+ verify=ssl_config,
+ transport=AsyncHTTPHandler._create_async_transport(
+ ssl_context=ssl_config
+ if isinstance(ssl_config, ssl.SSLContext)
+ else None,
+ ssl_verify=ssl_config if isinstance(ssl_config, bool) else None,
shared_session=shared_session,
- )
- # Return the underlying httpx client from the handler
- return cached_handler.client
- except (ImportError, AttributeError, KeyError) as e:
- # Fallback to creating a client directly if caching system unavailable
- # This preserves backwards compatibility
- verbose_logger.debug(
- f"Client caching unavailable ({type(e).__name__}), using direct client creation"
- )
- ssl_config = get_ssl_configuration()
- return httpx.AsyncClient(
- verify=ssl_config,
- transport=AsyncHTTPHandler._create_async_transport(
- ssl_context=ssl_config
- if isinstance(ssl_config, ssl.SSLContext)
- else None,
- ssl_verify=ssl_config if isinstance(ssl_config, bool) else None,
- shared_session=shared_session,
- ),
- follow_redirects=True,
- )
+ ),
+ follow_redirects=True,
+ )
@staticmethod
def _get_sync_http_client() -> Optional[httpx.Client]:
if litellm.client_session is not None:
return litellm.client_session
- # Use the global cached client system to prevent memory leaks (issue #14540)
- from litellm.llms.custom_httpx.http_handler import _get_httpx_client
+ # Get unified SSL configuration
+ ssl_config = get_ssl_configuration()
- try:
- # Get SSL config and include in params for proper cache key
- ssl_config = get_ssl_configuration()
- params = {"ssl_verify": ssl_config} if ssl_config is not None else None
-
- # Get a cached HTTPHandler which manages the httpx.Client
- cached_handler = _get_httpx_client(params=params)
- # Return the underlying httpx client from the handler
- return cached_handler.client
- except (ImportError, AttributeError, KeyError) as e:
- # Fallback to creating a client directly if caching system unavailable
- verbose_logger.debug(
- f"Client caching unavailable ({type(e).__name__}), using direct client creation"
- )
- ssl_config = get_ssl_configuration()
- return httpx.Client(
- verify=ssl_config,
- follow_redirects=True,
- )
+ return httpx.Client(
+ verify=ssl_config,
+ follow_redirects=True,
+ )
diff --git a/litellm/llms/openai/embeddings/guardrail_translation/__init__.py b/litellm/llms/openai/embeddings/guardrail_translation/__init__.py
new file mode 100644
index 00000000000..a60662282ca
--- /dev/null
+++ b/litellm/llms/openai/embeddings/guardrail_translation/__init__.py
@@ -0,0 +1,13 @@
+"""OpenAI Embeddings handler for Unified Guardrails."""
+
+from litellm.llms.openai.embeddings.guardrail_translation.handler import (
+ OpenAIEmbeddingsHandler,
+)
+from litellm.types.utils import CallTypes
+
+guardrail_translation_mappings = {
+ CallTypes.embedding: OpenAIEmbeddingsHandler,
+ CallTypes.aembedding: OpenAIEmbeddingsHandler,
+}
+
+__all__ = ["guardrail_translation_mappings", "OpenAIEmbeddingsHandler"]
diff --git a/litellm/llms/openai/embeddings/guardrail_translation/handler.py b/litellm/llms/openai/embeddings/guardrail_translation/handler.py
new file mode 100644
index 00000000000..7458020e109
--- /dev/null
+++ b/litellm/llms/openai/embeddings/guardrail_translation/handler.py
@@ -0,0 +1,179 @@
+"""
+OpenAI Embeddings Handler for Unified Guardrails
+
+This module provides guardrail translation support for OpenAI's embeddings endpoint.
+The handler processes the 'input' parameter for guardrails.
+"""
+
+from typing import TYPE_CHECKING, Any, List, Optional, Union
+
+from litellm._logging import verbose_proxy_logger
+from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
+from litellm.types.utils import GenericGuardrailAPIInputs
+
+if TYPE_CHECKING:
+ from litellm.integrations.custom_guardrail import CustomGuardrail
+ from litellm.types.utils import EmbeddingResponse
+
+
+class OpenAIEmbeddingsHandler(BaseTranslation):
+ """
+ Handler for processing OpenAI embeddings requests with guardrails.
+
+ This class provides methods to:
+ 1. Process input text (pre-call hook)
+ 2. Process output response (post-call hook) - embeddings don't typically need output guardrails
+
+ The handler specifically processes the 'input' parameter which can be:
+ - A single string
+ - A list of strings (for batch embeddings)
+ - A list of integers (token IDs - not processed by guardrails)
+ - A list of lists of integers (batch token IDs - not processed by guardrails)
+ """
+
+ async def process_input_messages(
+ self,
+ data: dict,
+ guardrail_to_apply: "CustomGuardrail",
+ litellm_logging_obj: Optional[Any] = None,
+ ) -> Any:
+ """
+ Process input text by applying guardrails to text content.
+
+ Args:
+ data: Request data dictionary containing 'input' parameter
+ guardrail_to_apply: The guardrail instance to apply
+ litellm_logging_obj: Optional logging object
+
+ Returns:
+ Modified data with guardrails applied to input
+ """
+ input_data = data.get("input")
+ if input_data is None:
+ verbose_proxy_logger.debug(
+ "OpenAI Embeddings: No input found in request data"
+ )
+ return data
+
+ if isinstance(input_data, str):
+ data = await self._process_string_input(
+ data, input_data, guardrail_to_apply, litellm_logging_obj
+ )
+ elif isinstance(input_data, list):
+ data = await self._process_list_input(
+ data, input_data, guardrail_to_apply, litellm_logging_obj
+ )
+ else:
+ verbose_proxy_logger.warning(
+ "OpenAI Embeddings: Unexpected input type: %s. Expected string or list.",
+ type(input_data),
+ )
+
+ return data
+
+ async def _process_string_input(
+ self,
+ data: dict,
+ input_data: str,
+ guardrail_to_apply: "CustomGuardrail",
+ litellm_logging_obj: Optional[Any],
+ ) -> dict:
+ """Process a single string input through the guardrail."""
+ inputs = GenericGuardrailAPIInputs(texts=[input_data])
+ if model := data.get("model"):
+ inputs["model"] = model
+
+ guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
+ inputs=inputs,
+ request_data=data,
+ input_type="request",
+ logging_obj=litellm_logging_obj,
+ )
+
+ if guardrailed_texts := guardrailed_inputs.get("texts"):
+ data["input"] = guardrailed_texts[0]
+ verbose_proxy_logger.debug(
+ "OpenAI Embeddings: Applied guardrail to string input. "
+ "Original length: %d, New length: %d",
+ len(input_data),
+ len(data["input"]),
+ )
+
+ return data
+
+ async def _process_list_input(
+ self,
+ data: dict,
+ input_data: List[Union[str, int, List[int]]],
+ guardrail_to_apply: "CustomGuardrail",
+ litellm_logging_obj: Optional[Any],
+ ) -> dict:
+ """Process a list input through the guardrail (if it contains strings)."""
+ if len(input_data) == 0:
+ return data
+
+ first_item = input_data[0]
+
+ # Skip non-text inputs (token IDs)
+ if isinstance(first_item, (int, list)):
+ verbose_proxy_logger.debug(
+ "OpenAI Embeddings: Input is token IDs, skipping guardrail processing"
+ )
+ return data
+
+ if not isinstance(first_item, str):
+ verbose_proxy_logger.warning(
+ "OpenAI Embeddings: Unexpected input list item type: %s",
+ type(first_item),
+ )
+ return data
+
+ # List of strings - apply guardrail
+ inputs = GenericGuardrailAPIInputs(texts=input_data) # type: ignore
+ if model := data.get("model"):
+ inputs["model"] = model
+
+ guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
+ inputs=inputs,
+ request_data=data,
+ input_type="request",
+ logging_obj=litellm_logging_obj,
+ )
+
+ if guardrailed_texts := guardrailed_inputs.get("texts"):
+ data["input"] = guardrailed_texts
+ verbose_proxy_logger.debug(
+ "OpenAI Embeddings: Applied guardrail to %d inputs",
+ len(guardrailed_texts),
+ )
+
+ return data
+
+ async def process_output_response(
+ self,
+ response: "EmbeddingResponse",
+ guardrail_to_apply: "CustomGuardrail",
+ litellm_logging_obj: Optional[Any] = None,
+ user_api_key_dict: Optional[Any] = None,
+ ) -> Any:
+ """
+ Process output response - embeddings responses contain vectors, not text.
+
+ For embeddings, the output is numerical vectors, so there's typically
+ no text content to apply guardrails to. This method is a no-op but
+ is included for interface consistency.
+
+ Args:
+ response: Embedding response object
+ guardrail_to_apply: The guardrail instance to apply
+ litellm_logging_obj: Optional logging object
+ user_api_key_dict: User API key metadata
+
+ Returns:
+ Unmodified response (embeddings don't have text output to guard)
+ """
+ verbose_proxy_logger.debug(
+ "OpenAI Embeddings: Output response processing skipped - "
+ "embeddings contain vectors, not text"
+ )
+ return response
diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py
index fd04ac4d458..ef9cc43c3e1 100644
--- a/litellm/llms/openai/realtime/handler.py
+++ b/litellm/llms/openai/realtime/handler.py
@@ -16,6 +16,62 @@ from ..openai import OpenAIChatCompletion
class OpenAIRealtime(OpenAIChatCompletion):
+ """
+ Base handler for OpenAI-compatible realtime WebSocket connections.
+
+ Subclasses can override template methods to customize:
+ - _get_default_api_base(): Default API base URL
+ - _get_additional_headers(): Extra headers beyond Authorization
+ - _get_ssl_config(): SSL configuration for WebSocket connection
+ """
+
+ def _get_default_api_base(self) -> str:
+ """
+ Get the default API base URL for this provider.
+ Override this in subclasses to set provider-specific defaults.
+ """
+ return "https://api.openai.com/"
+
+ def _get_additional_headers(self, api_key: str) -> dict:
+ """
+ Get additional headers beyond Authorization.
+ Override this in subclasses to customize headers (e.g., remove OpenAI-Beta).
+
+ Args:
+ api_key: API key for authentication
+
+ Returns:
+ Dictionary of additional headers
+ """
+ return {
+ "Authorization": f"Bearer {api_key}",
+ "OpenAI-Beta": "realtime=v1",
+ }
+
+ def _get_ssl_config(self, url: str) -> Any:
+ """
+ Get SSL configuration for WebSocket connection.
+ Override this in subclasses to customize SSL behavior.
+
+ Args:
+ url: WebSocket URL (ws:// or wss://)
+
+ Returns:
+ SSL configuration (None, True, or SSLContext)
+ """
+ if url.startswith("ws://"):
+ return None
+
+ # Use the shared SSL context which respects custom CA certs and SSL settings
+ ssl_config = get_shared_realtime_ssl_context()
+
+ # If ssl_config is False (ssl_verify=False), websockets library needs True instead
+ # to establish connection without verification (False would fail)
+ if ssl_config is False:
+ return True
+
+ return ssl_config
+
def _construct_url(self, api_base: str, query_params: RealtimeQueryParams) -> str:
"""
Construct the backend websocket URL with all query parameters (including 'model').
@@ -45,8 +101,9 @@ class OpenAIRealtime(OpenAIChatCompletion):
):
import websockets
from websockets.asyncio.client import ClientConnection
+
if api_base is None:
- api_base = "https://api.openai.com/"
+ api_base = self._get_default_api_base()
if api_key is None:
raise ValueError("api_key is required for OpenAI realtime calls")
@@ -56,30 +113,27 @@ class OpenAIRealtime(OpenAIChatCompletion):
url = self._construct_url(api_base, query_params)
try:
- # Only use SSL context for secure websocket connections (wss://)
- # websockets library doesn't accept ssl argument for ws:// URIs
- ssl_context = None if url.startswith("ws://") else get_shared_realtime_ssl_context()
+ # Get provider-specific SSL configuration
+ ssl_config = self._get_ssl_config(url)
+
+ # Get provider-specific headers
+ headers = self._get_additional_headers(api_key)
+
# Log a masked request preview consistent with other endpoints.
logging_obj.pre_call(
input=None,
api_key=api_key,
additional_args={
"api_base": url,
- "headers": {
- "Authorization": f"Bearer {api_key}",
- "OpenAI-Beta": "realtime=v1",
- },
+ "headers": headers,
"complete_input_dict": {"query_params": query_params},
},
)
async with websockets.connect( # type: ignore
url,
- additional_headers={
- "Authorization": f"Bearer {api_key}", # type: ignore
- "OpenAI-Beta": "realtime=v1",
- },
+ additional_headers=headers, # type: ignore
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
- ssl=ssl_context,
+ ssl=ssl_config,
) as backend_ws:
realtime_streaming = RealTimeStreaming(
websocket, cast(ClientConnection, backend_ws), logging_obj
diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py
index d943662f9e4..ad3d4c932d4 100644
--- a/litellm/llms/openai/responses/guardrail_translation/handler.py
+++ b/litellm/llms/openai/responses/guardrail_translation/handler.py
@@ -319,9 +319,7 @@ class OpenAIResponsesHandler(BaseTranslation):
return response
if not response_output:
- verbose_proxy_logger.debug(
- "OpenAI Responses API: Empty output in response"
- )
+ verbose_proxy_logger.debug("OpenAI Responses API: Empty output in response")
return response
# Step 1: Extract all text content and tool calls from response output
@@ -427,27 +425,30 @@ class OpenAIResponsesHandler(BaseTranslation):
handle_raw_dict_callback=None,
)
- tool_calls = model_response_choices[0].message.tool_calls
- text = model_response_choices[0].message.content
- guardrail_inputs = GenericGuardrailAPIInputs()
- if text:
- guardrail_inputs["texts"] = [text]
- if tool_calls:
- guardrail_inputs["tool_calls"] = cast(
- List[ChatCompletionToolCallChunk], tool_calls
- )
- # Include model information from the response if available
- response_model = final_chunk.get("response", {}).get("model")
- if response_model:
- guardrail_inputs["model"] = response_model
- if tool_calls or text:
- _guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
- inputs=guardrail_inputs,
- request_data={},
- input_type="response",
- logging_obj=litellm_logging_obj,
- )
- return responses_so_far
+ if model_response_choices:
+ tool_calls = model_response_choices[0].message.tool_calls
+ text = model_response_choices[0].message.content
+ guardrail_inputs = GenericGuardrailAPIInputs()
+ if text:
+ guardrail_inputs["texts"] = [text]
+ if tool_calls:
+ guardrail_inputs["tool_calls"] = cast(
+ List[ChatCompletionToolCallChunk], tool_calls
+ )
+ # Include model information from the response if available
+ response_model = final_chunk.get("response", {}).get("model")
+ if response_model:
+ guardrail_inputs["model"] = response_model
+ if tool_calls or text:
+ _guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
+ inputs=guardrail_inputs,
+ request_data={},
+ input_type="response",
+ logging_obj=litellm_logging_obj,
+ )
+ return responses_so_far
+ else:
+ verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices")
# model_response_stream = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(final_chunk)
# tool_calls = model_response_stream.choices[0].tool_calls
# convert openai response to model response
@@ -513,11 +514,9 @@ class OpenAIResponsesHandler(BaseTranslation):
# Check if it's an OutputText with text
if isinstance(content_item, OutputText):
if content_item.text:
-
return True
elif isinstance(content_item, dict):
if content_item.get("text"):
-
return True
return False
diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py
index 289963e917a..ed4d2d6a740 100644
--- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py
+++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py
@@ -27,6 +27,8 @@ local_cache_obj = Cache(
type=LiteLLMCacheType.LOCAL
) # only used for calling 'get_cache_key' function
+MAX_PAGINATION_PAGES = 100 # Reasonable upper bound for pagination
+
class ContextCachingEndpoints(VertexBase):
"""
@@ -115,7 +117,7 @@ class ContextCachingEndpoints(VertexBase):
- None
"""
- _, url = self._get_token_and_url_context_caching(
+ _, base_url = self._get_token_and_url_context_caching(
gemini_api_key=api_key,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
@@ -123,43 +125,63 @@ class ContextCachingEndpoints(VertexBase):
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
)
- try:
- ## LOGGING
- logging_obj.pre_call(
- input="",
- api_key="",
- additional_args={
- "complete_input_dict": {},
- "api_base": url,
- "headers": headers,
- },
- )
- resp = client.get(url=url, headers=headers)
- resp.raise_for_status()
- except httpx.HTTPStatusError as e:
- if e.response.status_code == 403:
+ page_token: Optional[str] = None
+
+ # Iterate through all pages
+ for _ in range(MAX_PAGINATION_PAGES):
+ # Build URL with pagination token if present
+ if page_token:
+ separator = "&" if "?" in base_url else "?"
+ url = f"{base_url}{separator}pageToken={page_token}"
+ else:
+ url = base_url
+
+ try:
+ ## LOGGING
+ logging_obj.pre_call(
+ input="",
+ api_key="",
+ additional_args={
+ "complete_input_dict": {},
+ "api_base": url,
+ "headers": headers,
+ },
+ )
+
+ resp = client.get(url=url, headers=headers)
+ resp.raise_for_status()
+ except httpx.HTTPStatusError as e:
+ if e.response.status_code == 403:
+ return None
+ raise VertexAIError(
+ status_code=e.response.status_code, message=e.response.text
+ )
+ except Exception as e:
+ raise VertexAIError(status_code=500, message=str(e))
+
+ raw_response = resp.json()
+ logging_obj.post_call(original_response=raw_response)
+
+ if "cachedContents" not in raw_response:
return None
- raise VertexAIError(
- status_code=e.response.status_code, message=e.response.text
- )
- except Exception as e:
- raise VertexAIError(status_code=500, message=str(e))
- raw_response = resp.json()
- logging_obj.post_call(original_response=raw_response)
- if "cachedContents" not in raw_response:
- return None
+ all_cached_items = CachedContentListAllResponseBody(**raw_response)
- all_cached_items = CachedContentListAllResponseBody(**raw_response)
+ if "cachedContents" not in all_cached_items:
+ return None
- if "cachedContents" not in all_cached_items:
- return None
+ # Check current page for matching cache_key
+ for cached_item in all_cached_items["cachedContents"]:
+ display_name = cached_item.get("displayName")
+ if display_name is not None and display_name == cache_key:
+ return cached_item.get("name")
- for cached_item in all_cached_items["cachedContents"]:
- display_name = cached_item.get("displayName")
- if display_name is not None and display_name == cache_key:
- return cached_item.get("name")
+ # Check if there are more pages
+ page_token = all_cached_items.get("nextPageToken")
+ if not page_token:
+ # No more pages, cache not found
+ break
return None
@@ -187,7 +209,7 @@ class ContextCachingEndpoints(VertexBase):
- None
"""
- _, url = self._get_token_and_url_context_caching(
+ _, base_url = self._get_token_and_url_context_caching(
gemini_api_key=api_key,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
@@ -195,43 +217,63 @@ class ContextCachingEndpoints(VertexBase):
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
)
- try:
- ## LOGGING
- logging_obj.pre_call(
- input="",
- api_key="",
- additional_args={
- "complete_input_dict": {},
- "api_base": url,
- "headers": headers,
- },
- )
- resp = await client.get(url=url, headers=headers)
- resp.raise_for_status()
- except httpx.HTTPStatusError as e:
- if e.response.status_code == 403:
+ page_token: Optional[str] = None
+
+ # Iterate through all pages
+ for _ in range(MAX_PAGINATION_PAGES):
+ # Build URL with pagination token if present
+ if page_token:
+ separator = "&" if "?" in base_url else "?"
+ url = f"{base_url}{separator}pageToken={page_token}"
+ else:
+ url = base_url
+
+ try:
+ ## LOGGING
+ logging_obj.pre_call(
+ input="",
+ api_key="",
+ additional_args={
+ "complete_input_dict": {},
+ "api_base": url,
+ "headers": headers,
+ },
+ )
+
+ resp = await client.get(url=url, headers=headers)
+ resp.raise_for_status()
+ except httpx.HTTPStatusError as e:
+ if e.response.status_code == 403:
+ return None
+ raise VertexAIError(
+ status_code=e.response.status_code, message=e.response.text
+ )
+ except Exception as e:
+ raise VertexAIError(status_code=500, message=str(e))
+
+ raw_response = resp.json()
+ logging_obj.post_call(original_response=raw_response)
+
+ if "cachedContents" not in raw_response:
return None
- raise VertexAIError(
- status_code=e.response.status_code, message=e.response.text
- )
- except Exception as e:
- raise VertexAIError(status_code=500, message=str(e))
- raw_response = resp.json()
- logging_obj.post_call(original_response=raw_response)
- if "cachedContents" not in raw_response:
- return None
+ all_cached_items = CachedContentListAllResponseBody(**raw_response)
- all_cached_items = CachedContentListAllResponseBody(**raw_response)
+ if "cachedContents" not in all_cached_items:
+ return None
- if "cachedContents" not in all_cached_items:
- return None
+ # Check current page for matching cache_key
+ for cached_item in all_cached_items["cachedContents"]:
+ display_name = cached_item.get("displayName")
+ if display_name is not None and display_name == cache_key:
+ return cached_item.get("name")
- for cached_item in all_cached_items["cachedContents"]:
- display_name = cached_item.get("displayName")
- if display_name is not None and display_name == cache_key:
- return cached_item.get("name")
+ # Check if there are more pages
+ page_token = all_cached_items.get("nextPageToken")
+ if not page_token:
+ # No more pages, cache not found
+ break
return None
@@ -501,4 +543,4 @@ class ContextCachingEndpoints(VertexBase):
pass
async def async_get_cache(self):
- pass
+ pass
\ No newline at end of file
diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py
index b3612113ec2..2470c59bbac 100644
--- a/litellm/llms/vertex_ai/files/transformation.py
+++ b/litellm/llms/vertex_ai/files/transformation.py
@@ -165,7 +165,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"""
Get the complete url for the request
"""
- bucket_name = litellm_params.get("bucket_name") or os.getenv("GCS_BUCKET_NAME")
+ bucket_name = litellm_params.get("bucket_name") or litellm_params.get("litellm_metadata", {}).pop("gcs_bucket_name", None) or os.getenv("GCS_BUCKET_NAME")
if not bucket_name:
raise ValueError("GCS bucket_name is required")
file_data = data.get("file")
diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py
index a9ac21bb56f..04ae4b6beb8 100644
--- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py
+++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py
@@ -478,6 +478,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
if "type" in tool and tool["type"] == "computer_use":
computer_use_config = {k: v for k, v in tool.items() if k != "type"}
tool = {VertexToolName.COMPUTER_USE.value: computer_use_config}
+ # Handle OpenAI-style web_search and web_search_preview tools
+ # Transform them to Gemini's googleSearch tool
+ elif "type" in tool and tool["type"] in ("web_search", "web_search_preview"):
+ verbose_logger.info(
+ f"Gemini: Transforming OpenAI-style '{tool['type']}' tool to googleSearch"
+ )
+ tool = {VertexToolName.GOOGLE_SEARCH.value: {}}
# Handle tools with 'type' field (OpenAI spec compliance) Ignore this field -> https://github.com/BerriAI/litellm/issues/14644#issuecomment-3342061838
elif "type" in tool:
tool = {k: tool[k] for k in tool if k != "type"}
@@ -1725,6 +1732,52 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
else:
return "stop"
+ @staticmethod
+ def _check_prompt_level_content_filter(
+ processed_chunk: GenerateContentResponseBody,
+ response_id: Optional[str],
+ ) -> Optional["ModelResponseStream"]:
+ """
+ Check if prompt is blocked due to content filtering at the prompt level.
+
+ This handles the case where Vertex AI blocks the prompt before generation begins,
+ indicated by promptFeedback.blockReason being present.
+
+ Args:
+ processed_chunk: The parsed response chunk from Vertex AI
+ response_id: The response ID from the chunk
+
+ Returns:
+ ModelResponseStream with content_filter finish_reason if blocked, None otherwise.
+
+ Note:
+ This is consistent with non-streaming _handle_blocked_response() behavior.
+ Candidate-level content filtering (SAFETY, RECITATION, etc.) is handled
+ separately via _process_candidates() → _check_finish_reason().
+ """
+ from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
+
+ # Check if prompt is blocked due to content filtering
+ prompt_feedback = processed_chunk.get("promptFeedback")
+ if prompt_feedback and "blockReason" in prompt_feedback:
+ verbose_logger.debug(
+ f"Prompt blocked due to: {prompt_feedback.get('blockReason')} - {prompt_feedback.get('blockReasonMessage')}"
+ )
+
+ # Create a content_filter response (consistent with non-streaming _handle_blocked_response)
+ choice = StreamingChoices(
+ finish_reason="content_filter",
+ index=0,
+ delta=Delta(content=None, role="assistant"),
+ logprobs=None,
+ enhancements=None,
+ )
+
+ model_response = ModelResponseStream(choices=[choice], id=response_id)
+ return model_response
+
+ return None
+
@staticmethod
def _calculate_web_search_requests(grounding_metadata: List[dict]) -> Optional[int]:
web_search_requests: Optional[int] = None
@@ -2806,6 +2859,15 @@ class ModelResponseIterator:
processed_chunk = GenerateContentResponseBody(**chunk) # type: ignore
response_id = processed_chunk.get("responseId")
model_response = ModelResponseStream(choices=[], id=response_id)
+
+ # Check if prompt is blocked due to content filtering
+ blocked_response = VertexGeminiConfig._check_prompt_level_content_filter(
+ processed_chunk=processed_chunk,
+ response_id=response_id,
+ )
+ if blocked_response is not None:
+ model_response = blocked_response
+
usage: Optional[Usage] = None
_candidates: Optional[List[Candidates]] = processed_chunk.get("candidates")
grounding_metadata: List[dict] = []
diff --git a/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py b/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py
index 89ed9f1a8a5..ba3df88be14 100644
--- a/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py
+++ b/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py
@@ -295,9 +295,11 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
if "inlineData" in part:
inline_data = part["inlineData"]
if "data" in inline_data:
+ thought_sig = part.get("thoughtSignature")
model_response.data.append(ImageObject(
b64_json=inline_data["data"],
url=None,
+ provider_specific_fields={"thought_signature": thought_sig} if thought_sig else None,
))
if usage_metadata := response_data.get("usageMetadata", None):
diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py
index 9b8ff3ecc2d..918b8ecc225 100644
--- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py
+++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py
@@ -1,5 +1,8 @@
from typing import Any, Dict, List, Optional, Tuple
+from litellm.anthropic_beta_headers_manager import (
+ update_headers_with_filtered_beta,
+)
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
@@ -7,7 +10,6 @@ from litellm.llms.anthropic.experimental_pass_through.messages.transformation im
from litellm.types.llms.anthropic import (
ANTHROPIC_BETA_HEADER_VALUES,
ANTHROPIC_HOSTED_TOOLS,
- ANTHROPIC_PROMPT_CACHING_SCOPE_BETA_HEADER,
)
from litellm.types.llms.anthropic_tool_search import get_tool_search_beta_header
from litellm.types.llms.vertex_ai import VertexPartnerProvider
@@ -65,10 +67,6 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert
existing_beta = headers.get("anthropic-beta")
if existing_beta:
beta_values.update(b.strip() for b in existing_beta.split(","))
-
- # Use the helper to remove unsupported beta headers
- self.remove_unsupported_beta(headers)
- beta_values.discard(ANTHROPIC_PROMPT_CACHING_SCOPE_BETA_HEADER)
# Check for web search tool
for tool in tools:
@@ -84,6 +82,12 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert
if beta_values:
headers["anthropic-beta"] = ",".join(beta_values)
+ # Filter out unsupported beta headers for Vertex AI
+ headers = update_headers_with_filtered_beta(
+ headers=headers,
+ provider="vertex_ai",
+ )
+
return headers, api_base
def get_complete_url(
@@ -128,23 +132,3 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert
) # do not pass output_format in request body to vertex ai - vertex ai does not support output_format as yet
return anthropic_messages_request
-
- def remove_unsupported_beta(self, headers: dict) -> None:
- """
- Helper method to remove unsupported beta headers from the beta headers.
- Modifies headers in place.
- """
- unsupported_beta_headers = [
- ANTHROPIC_PROMPT_CACHING_SCOPE_BETA_HEADER
- ]
- existing_beta = headers.get("anthropic-beta")
- if existing_beta:
- filtered_beta = [
- b.strip()
- for b in existing_beta.split(",")
- if b.strip() not in unsupported_beta_headers
- ]
- if filtered_beta:
- headers["anthropic-beta"] = ",".join(filtered_beta)
- elif "anthropic-beta" in headers:
- del headers["anthropic-beta"]
diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py
index 1df07f405e6..0b728d88e76 100644
--- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py
+++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py
@@ -51,6 +51,40 @@ class VertexAIAnthropicConfig(AnthropicConfig):
def custom_llm_provider(self) -> Optional[str]:
return "vertex_ai"
+ def _add_context_management_beta_headers(
+ self, beta_set: set, context_management: dict
+ ) -> None:
+ """
+ Add context_management beta headers to the beta_set.
+
+ - If any edit has type "compact_20260112", add compact-2026-01-12 header
+ - For all other edits, add context-management-2025-06-27 header
+
+ Args:
+ beta_set: Set of beta headers to modify in-place
+ context_management: The context_management dict from optional_params
+ """
+ from litellm.types.llms.anthropic import ANTHROPIC_BETA_HEADER_VALUES
+
+ edits = context_management.get("edits", [])
+ has_compact = False
+ has_other = False
+
+ for edit in edits:
+ edit_type = edit.get("type", "")
+ if edit_type == "compact_20260112":
+ has_compact = True
+ else:
+ has_other = True
+
+ # Add compact header if any compact edits exist
+ if has_compact:
+ beta_set.add(ANTHROPIC_BETA_HEADER_VALUES.COMPACT_2026_01_12.value)
+
+ # Add context management header if any other edits exist
+ if has_other:
+ beta_set.add(ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value)
+
def transform_request(
self,
model: str,
@@ -86,6 +120,11 @@ class VertexAIAnthropicConfig(AnthropicConfig):
beta_set = set(auto_betas)
if tool_search_used:
beta_set.add("tool-search-tool-2025-10-19") # Vertex requires this header for tool search
+
+ # Add context_management beta headers (compact and/or context-management)
+ context_management = optional_params.get("context_management")
+ if context_management:
+ self._add_context_management_beta_headers(beta_set, context_management)
if beta_set:
data["anthropic_beta"] = list(beta_set)
diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py
index a185370e376..4613b6a5715 100644
--- a/litellm/llms/vertex_ai/vertex_llm_base.py
+++ b/litellm/llms/vertex_ai/vertex_llm_base.py
@@ -20,6 +20,7 @@ from .common_utils import (
_get_vertex_url,
all_gemini_url_modes,
get_vertex_base_model_name,
+ get_vertex_base_url,
is_global_only_vertex_model,
)
@@ -200,12 +201,7 @@ class VertexBase:
) -> str:
if api_base:
return api_base
- elif vertex_location == "global":
- return "https://aiplatform.googleapis.com"
- elif vertex_location:
- return f"https://{vertex_location}-aiplatform.googleapis.com"
- else:
- return f"https://{self.get_default_vertex_location()}-aiplatform.googleapis.com"
+ return get_vertex_base_url(vertex_location or self.get_default_vertex_location())
@staticmethod
def create_vertex_url(
@@ -218,7 +214,8 @@ class VertexBase:
) -> str:
"""Return the base url for the vertex partner models"""
- api_base = api_base or f"https://{vertex_location}-aiplatform.googleapis.com"
+ if api_base is None:
+ api_base = get_vertex_base_url(vertex_location)
if partner == VertexPartnerProvider.llama:
return f"{api_base}/v1/projects/{vertex_project}/locations/{vertex_location}/endpoints/openapi/chat/completions"
elif partner == VertexPartnerProvider.mistralai:
@@ -247,11 +244,13 @@ class VertexBase:
stream: Optional[bool],
model: str,
) -> str:
+ # Use get_vertex_region to handle global-only models
+ resolved_location = self.get_vertex_region(vertex_location, model)
api_base = self.get_api_base(
- api_base=custom_api_base, vertex_location=vertex_location
+ api_base=custom_api_base, vertex_location=resolved_location
)
default_api_base = VertexBase.create_vertex_url(
- vertex_location=vertex_location or "us-central1",
+ vertex_location=resolved_location,
vertex_project=vertex_project or project_id,
partner=partner,
stream=stream,
@@ -274,7 +273,7 @@ class VertexBase:
url=default_api_base,
model=model,
vertex_project=vertex_project or project_id,
- vertex_location=vertex_location or "us-central1",
+ vertex_location=resolved_location,
vertex_api_version="v1", # Partner models typically use v1
)
return api_base
diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py
index 245e10e45c1..21782fc6fbf 100644
--- a/litellm/llms/xai/chat/transformation.py
+++ b/litellm/llms/xai/chat/transformation.py
@@ -4,6 +4,7 @@ import httpx
import litellm
from litellm._logging import verbose_logger
+from litellm.constants import XAI_API_BASE
from litellm.litellm_core_utils.prompt_templates.common_utils import (
filter_value_from_dict,
strip_name_from_messages,
@@ -14,8 +15,6 @@ from litellm.types.utils import Choices, ModelResponse, Usage, PromptTokensDetai
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
-XAI_API_BASE = "https://api.x.ai/v1"
-
class XAIChatConfig(OpenAIGPTConfig):
@property
diff --git a/litellm/llms/xai/realtime/__init__.py b/litellm/llms/xai/realtime/__init__.py
new file mode 100644
index 00000000000..3b0d345f2c2
--- /dev/null
+++ b/litellm/llms/xai/realtime/__init__.py
@@ -0,0 +1,5 @@
+"""xAI Realtime API handler."""
+
+from .handler import XAIRealtime
+
+__all__ = ["XAIRealtime"]
diff --git a/litellm/llms/xai/realtime/handler.py b/litellm/llms/xai/realtime/handler.py
new file mode 100644
index 00000000000..c79477ba1df
--- /dev/null
+++ b/litellm/llms/xai/realtime/handler.py
@@ -0,0 +1,38 @@
+"""
+This file contains the handler for xAI's Grok Voice Agent API `/v1/realtime` endpoint.
+
+xAI's Realtime API is fully OpenAI-compatible, so we inherit from OpenAIRealtime
+and only override the configuration differences.
+
+This requires websockets, and is currently only supported on LiteLLM Proxy.
+"""
+
+from litellm.constants import XAI_API_BASE
+
+from ...openai.realtime.handler import OpenAIRealtime
+
+
+class XAIRealtime(OpenAIRealtime):
+ """
+ Handler for xAI Grok Voice Agent API.
+
+ xAI's Realtime API uses the same WebSocket protocol as OpenAI but with:
+ - Different endpoint: wss://api.x.ai/v1/realtime (via _get_default_api_base)
+ - No OpenAI-Beta header required (via _get_additional_headers)
+ - Model: grok-4-1-fast-non-reasoning
+
+ All WebSocket logic is inherited from OpenAIRealtime.
+ """
+
+ def _get_default_api_base(self) -> str:
+ """xAI uses a different API base URL."""
+ return XAI_API_BASE
+
+ def _get_additional_headers(self, api_key: str) -> dict:
+ """
+ xAI does NOT require the OpenAI-Beta header.
+ Only send Authorization header.
+ """
+ return {
+ "Authorization": f"Bearer {api_key}",
+ }
diff --git a/litellm/llms/xai/responses/transformation.py b/litellm/llms/xai/responses/transformation.py
index 82b4771fb4d..95873aab846 100644
--- a/litellm/llms/xai/responses/transformation.py
+++ b/litellm/llms/xai/responses/transformation.py
@@ -2,6 +2,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
import litellm
from litellm._logging import verbose_logger
+from litellm.constants import XAI_API_BASE
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
@@ -16,8 +17,6 @@ if TYPE_CHECKING:
else:
LiteLLMLoggingObj = Any
-XAI_API_BASE = "https://api.x.ai/v1"
-
class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
"""
diff --git a/litellm/main.py b/litellm/main.py
index 13361c644cb..bca023e65ec 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -1199,6 +1199,13 @@ def completion( # type: ignore # noqa: PLR0915
headers = {}
if extra_headers is not None:
headers.update(extra_headers)
+ # Inject proxy auth headers if configured
+ if litellm.proxy_auth is not None:
+ try:
+ proxy_headers = litellm.proxy_auth.get_auth_headers()
+ headers.update(proxy_headers)
+ except Exception as e:
+ verbose_logger.warning(f"Failed to get proxy auth headers: {e}")
num_retries = kwargs.get(
"num_retries", None
) ## alt. param for 'max_retries'. Use this to pass retries w/ instructor.
@@ -2199,6 +2206,48 @@ def completion( # type: ignore # noqa: PLR0915
logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
client=client,
)
+ elif custom_llm_provider == "a2a":
+ # A2A (Agent-to-Agent) Protocol
+ # Resolve agent configuration from registry if model format is "a2a/"
+ api_base, api_key, headers = litellm.A2AConfig.resolve_agent_config_from_registry(
+ model=model,
+ api_base=api_base,
+ api_key=api_key,
+ headers=headers,
+ optional_params=optional_params,
+ )
+
+ # Fall back to environment variables and defaults
+ api_base = api_base or litellm.api_base or get_secret_str("A2A_API_BASE")
+
+ if api_base is None:
+ raise Exception(
+ "api_base is required for A2A provider. "
+ "Either provide api_base parameter, set A2A_API_BASE environment variable, "
+ "or register the agent in the proxy with model='a2a/'."
+ )
+
+ headers = headers or litellm.headers
+
+ response = base_llm_http_handler.completion(
+ model=model,
+ stream=stream,
+ messages=messages,
+ acompletion=acompletion,
+ api_base=api_base,
+ model_response=model_response,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ custom_llm_provider=custom_llm_provider,
+ timeout=timeout,
+ headers=headers,
+ encoding=_get_encoding(),
+ api_key=api_key,
+ logging_obj=logging,
+ client=client,
+ provider_config=provider_config,
+ )
elif custom_llm_provider == "gigachat":
# GigaChat - Sber AI's LLM (Russia)
api_key = (
@@ -2455,6 +2504,20 @@ def completion( # type: ignore # noqa: PLR0915
headers = headers or litellm.headers
+ # Add GitHub Copilot headers (same as /responses endpoint does)
+ if custom_llm_provider == "github_copilot":
+ from litellm.llms.github_copilot.common_utils import (
+ get_copilot_default_headers,
+ )
+ from litellm.llms.github_copilot.authenticator import Authenticator
+
+ copilot_auth = Authenticator()
+ copilot_api_key = copilot_auth.get_api_key()
+ copilot_headers = get_copilot_default_headers(copilot_api_key)
+ if extra_headers:
+ copilot_headers.update(extra_headers)
+ extra_headers = copilot_headers
+
if extra_headers is not None:
optional_params["extra_headers"] = extra_headers
@@ -3113,8 +3176,8 @@ def completion( # type: ignore # noqa: PLR0915
api_key
or litellm.api_key
or litellm.openrouter_key
- or get_secret("OPENROUTER_API_KEY")
- or get_secret("OR_API_KEY")
+ or get_secret_str("OPENROUTER_API_KEY")
+ or get_secret_str("OR_API_KEY")
)
openrouter_site_url = get_secret("OR_SITE_URL") or "https://litellm.ai"
@@ -4555,6 +4618,13 @@ def embedding( # noqa: PLR0915
headers = {}
if extra_headers is not None:
headers.update(extra_headers)
+ # Inject proxy auth headers if configured
+ if litellm.proxy_auth is not None:
+ try:
+ proxy_headers = litellm.proxy_auth.get_auth_headers()
+ headers.update(proxy_headers)
+ except Exception as e:
+ verbose_logger.warning(f"Failed to get proxy auth headers: {e}")
### CUSTOM MODEL COST ###
input_cost_per_token = kwargs.get("input_cost_per_token", None)
output_cost_per_token = kwargs.get("output_cost_per_token", None)
@@ -4709,11 +4779,11 @@ def embedding( # noqa: PLR0915
litellm_params=litellm_params_dict,
)
elif (
- model in litellm.open_ai_embedding_models
- or custom_llm_provider == "openai"
+ custom_llm_provider == "openai"
or custom_llm_provider == "together_ai"
or custom_llm_provider == "nvidia_nim"
or custom_llm_provider == "litellm_proxy"
+ or (model in litellm.open_ai_embedding_models and custom_llm_provider is None)
):
api_base = (
api_base
@@ -4884,8 +4954,8 @@ def embedding( # noqa: PLR0915
api_key
or litellm.api_key
or litellm.openrouter_key
- or get_secret("OPENROUTER_API_KEY")
- or get_secret("OR_API_KEY")
+ or get_secret_str("OPENROUTER_API_KEY")
+ or get_secret_str("OR_API_KEY")
)
openrouter_site_url = get_secret("OR_SITE_URL") or "https://litellm.ai"
diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json
index 0f84bba941d..0da47634a94 100644
--- a/litellm/model_prices_and_context_window_backup.json
+++ b/litellm/model_prices_and_context_window_backup.json
@@ -744,12 +744,13 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_streaming": true
},
"anthropic.claude-3-5-sonnet-20240620-v1:0": {
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
@@ -758,14 +759,22 @@
"supports_pdf_input": true,
"supports_response_schema": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "input_cost_per_token_above_200k_tokens": 6e-06,
+ "output_cost_per_token_above_200k_tokens": 3e-05,
+ "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
+ "cache_read_input_token_cost_above_200k_tokens": 6e-07,
+ "cache_creation_input_token_cost_above_1hr": 7.5e-06,
+ "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.5e-05,
+ "cache_creation_input_token_cost": 3.75e-06,
+ "cache_read_input_token_cost": 3e-07
},
"anthropic.claude-3-5-sonnet-20241022-v2:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
@@ -777,7 +786,13 @@
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "input_cost_per_token_above_200k_tokens": 6e-06,
+ "output_cost_per_token_above_200k_tokens": 3e-05,
+ "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
+ "cache_read_input_token_cost_above_200k_tokens": 6e-07,
+ "cache_creation_input_token_cost_above_1hr": 7.5e-06,
+ "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.5e-05
},
"anthropic.claude-3-7-sonnet-20240620-v1:0": {
"cache_creation_input_token_cost": 4.5e-06,
@@ -948,6 +963,306 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 159
},
+ "anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "global.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "global.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "us.anthropic.claude-opus-4-6-v1:0": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "us.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "eu.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "eu.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "apac.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "apac.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
"anthropic.claude-sonnet-4-20250514-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
@@ -1429,6 +1744,33 @@
"supports_tool_choice": true,
"supports_vision": true
},
+ "azure_ai/claude-opus-4-6": {
+ "input_cost_per_token": 5e-06,
+ "output_cost_per_token": 2.5e-05,
+ "litellm_provider": "azure_ai",
+ "max_input_tokens": 200000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_1hr": 1e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 159
+ },
"azure_ai/claude-opus-4-1": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_creation_input_token_cost_above_1hr": 3e-05,
@@ -6715,13 +7057,13 @@
"supports_tool_choice": true
},
"cerebras/gpt-oss-120b": {
- "input_cost_per_token": 2.5e-07,
+ "input_cost_per_token": 3.5e-07,
"litellm_provider": "cerebras",
"max_input_tokens": 131072,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 6.9e-07,
+ "output_cost_per_token": 7.5e-07,
"source": "https://www.cerebras.ai/blog/openai-gpt-oss-120b-runs-fastest-on-cerebras",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
@@ -6739,6 +7081,7 @@
"output_cost_per_token": 8e-07,
"source": "https://inference-docs.cerebras.ai/support/pricing",
"supports_function_calling": true,
+ "supports_reasoning": true,
"supports_tool_choice": true
},
"cerebras/zai-glm-4.6": {
@@ -7439,6 +7782,130 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 159
},
+ "claude-opus-4-6": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_creation_input_token_cost_above_1hr": 1e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "anthropic",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "us/claude-opus-4-6": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_creation_input_token_cost_above_1hr": 1.1e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "anthropic",
+ "max_input_tokens": 200000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "claude-opus-4-6-20260205": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_creation_input_token_cost_above_1hr": 1e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "anthropic",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "us/claude-opus-4-6-20260205": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_creation_input_token_cost_above_1hr": 1.1e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "anthropic",
+ "max_input_tokens": 200000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
"claude-sonnet-4-20250514": {
"deprecation_date": "2026-05-14",
"cache_creation_input_token_cost": 3.75e-06,
@@ -10559,6 +11026,32 @@
"/v1/audio/transcriptions"
]
},
+ "elevenlabs/eleven_v3": {
+ "input_cost_per_character": 0.00018,
+ "litellm_provider": "elevenlabs",
+ "metadata": {
+ "calculation": "$0.18/1000 characters (Scale plan pricing, 1 credit per character)",
+ "notes": "ElevenLabs Eleven v3 - most expressive TTS model with 70+ languages and audio tags support"
+ },
+ "mode": "audio_speech",
+ "source": "https://elevenlabs.io/pricing",
+ "supported_endpoints": [
+ "/v1/audio/speech"
+ ]
+ },
+ "elevenlabs/eleven_multilingual_v2": {
+ "input_cost_per_character": 0.00018,
+ "litellm_provider": "elevenlabs",
+ "metadata": {
+ "calculation": "$0.18/1000 characters (Scale plan pricing, 1 credit per character)",
+ "notes": "ElevenLabs Eleven Multilingual v2 - default TTS model with 29 languages support"
+ },
+ "mode": "audio_speech",
+ "source": "https://elevenlabs.io/pricing",
+ "supported_endpoints": [
+ "/v1/audio/speech"
+ ]
+ },
"embed-english-light-v2.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "cohere",
@@ -12835,6 +13328,40 @@
"supports_vision": true,
"supports_web_search": true
},
+ "deep-research-pro-preview-12-2025": {
+ "input_cost_per_image": 0.0011,
+ "input_cost_per_token": 2e-06,
+ "input_cost_per_token_batches": 1e-06,
+ "litellm_provider": "vertex_ai-language-models",
+ "max_input_tokens": 65536,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "image_generation",
+ "output_cost_per_image": 0.134,
+ "output_cost_per_image_token": 0.00012,
+ "output_cost_per_token": 1.2e-05,
+ "output_cost_per_token_batches": 6e-06,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/completions",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "image"
+ ],
+ "supports_function_calling": false,
+ "supports_prompt_caching": true,
+ "supports_response_schema": true,
+ "supports_system_messages": true,
+ "supports_vision": true,
+ "supports_web_search": true
+ },
"gemini-2.5-flash-lite": {
"cache_read_input_token_cost": 1e-08,
"input_cost_per_audio_token": 3e-07,
@@ -13289,7 +13816,8 @@
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_native_streaming": true
},
"vertex_ai/gemini-3-pro-preview": {
"cache_read_input_token_cost": 2e-07,
@@ -13337,7 +13865,8 @@
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_native_streaming": true
},
"vertex_ai/gemini-3-flash-preview": {
"cache_read_input_token_cost": 5e-08,
@@ -13380,7 +13909,8 @@
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_native_streaming": true
},
"gemini-2.5-pro-exp-03-25": {
"cache_read_input_token_cost": 1.25e-07,
@@ -14747,6 +15277,42 @@
"supports_vision": true,
"supports_web_search": true
},
+ "gemini/deep-research-pro-preview-12-2025": {
+ "input_cost_per_image": 0.0011,
+ "input_cost_per_token": 2e-06,
+ "input_cost_per_token_batches": 1e-06,
+ "litellm_provider": "gemini",
+ "max_input_tokens": 65536,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "image_generation",
+ "output_cost_per_image": 0.134,
+ "output_cost_per_image_token": 0.00012,
+ "output_cost_per_token": 1.2e-05,
+ "rpm": 1000,
+ "tpm": 4000000,
+ "output_cost_per_token_batches": 6e-06,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/completions",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "image"
+ ],
+ "supports_function_calling": false,
+ "supports_prompt_caching": true,
+ "supports_response_schema": true,
+ "supports_system_messages": true,
+ "supports_vision": true,
+ "supports_web_search": true
+ },
"gemini/gemini-2.5-flash-lite": {
"cache_read_input_token_cost": 1e-08,
"input_cost_per_audio_token": 3e-07,
@@ -15331,6 +15897,7 @@
"supports_url_context": true,
"supports_vision": true,
"supports_web_search": true,
+ "supports_native_streaming": true,
"tpm": 800000
},
"gemini-3-flash-preview": {
@@ -15376,7 +15943,8 @@
"supports_tool_choice": true,
"supports_url_context": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_native_streaming": true
},
"gemini/gemini-2.5-pro-exp-03-25": {
"cache_read_input_token_cost": 0.0,
@@ -21473,6 +22041,20 @@
"supports_tool_choice": true,
"supports_web_search": true
},
+ "moonshot/kimi-k2.5": {
+ "cache_read_input_token_cost": 1e-07,
+ "input_cost_per_token": 6e-07,
+ "litellm_provider": "moonshot",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 3e-06,
+ "source": "https://platform.moonshot.ai/docs/pricing/chat",
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_vision": true
+ },
"moonshot/kimi-latest": {
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 2e-06,
@@ -24314,6 +24896,31 @@
"supports_tool_choice": true,
"supports_function_calling": true
},
+ "openrouter/qwen/qwen3-235b-a22b-2507": {
+ "input_cost_per_token": 7.1e-08,
+ "litellm_provider": "openrouter",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 1e-07,
+ "source": "https://openrouter.ai/qwen/qwen3-235b-a22b-2507",
+ "supports_function_calling": true,
+ "supports_tool_choice": true
+ },
+ "openrouter/qwen/qwen3-235b-a22b-thinking-2507": {
+ "input_cost_per_token": 1.1e-07,
+ "litellm_provider": "openrouter",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 6e-07,
+ "source": "https://openrouter.ai/qwen/qwen3-235b-a22b-thinking-2507",
+ "supports_function_calling": true,
+ "supports_reasoning": true,
+ "supports_tool_choice": true
+ },
"openrouter/switchpoint/router": {
"input_cost_per_token": 8.5e-07,
"litellm_provider": "openrouter",
@@ -24390,21 +24997,21 @@
"supports_tool_choice": true
},
"openrouter/xiaomi/mimo-v2-flash": {
- "input_cost_per_token": 9e-08,
- "output_cost_per_token": 2.9e-07,
- "cache_creation_input_token_cost": 0.0,
- "cache_read_input_token_cost": 0.0,
- "litellm_provider": "openrouter",
- "max_input_tokens": 262144,
- "max_output_tokens": 16384,
- "max_tokens": 16384,
- "mode": "chat",
- "supports_function_calling": true,
- "supports_tool_choice": true,
- "supports_reasoning": true,
- "supports_vision": false,
- "supports_prompt_caching": false
- },
+ "input_cost_per_token": 9e-08,
+ "output_cost_per_token": 2.9e-07,
+ "cache_creation_input_token_cost": 0.0,
+ "cache_read_input_token_cost": 0.0,
+ "litellm_provider": "openrouter",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 16384,
+ "max_tokens": 16384,
+ "mode": "chat",
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_reasoning": true,
+ "supports_vision": false,
+ "supports_prompt_caching": false
+ },
"openrouter/z-ai/glm-4.7": {
"input_cost_per_token": 4e-07,
"output_cost_per_token": 1.5e-06,
@@ -26319,13 +26926,13 @@
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"mode": "image_edit",
- "output_cost_per_image": 0.40
+ "output_cost_per_image": 0.4
},
"stability.stable-creative-upscale-v1:0": {
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"mode": "image_edit",
- "output_cost_per_image": 0.60
+ "output_cost_per_image": 0.6
},
"stability.stable-fast-upscale-v1:0": {
"litellm_provider": "bedrock",
@@ -27084,6 +27691,34 @@
"supports_reasoning": true,
"supports_tool_choice": true
},
+ "together_ai/zai-org/GLM-4.7": {
+ "input_cost_per_token": 4.5e-07,
+ "litellm_provider": "together_ai",
+ "max_input_tokens": 200000,
+ "max_output_tokens": 200000,
+ "max_tokens": 200000,
+ "mode": "chat",
+ "output_cost_per_token": 2e-06,
+ "source": "https://www.together.ai/models/glm-4-7",
+ "supports_function_calling": true,
+ "supports_parallel_function_calling": true,
+ "supports_reasoning": true,
+ "supports_tool_choice": true
+ },
+ "together_ai/moonshotai/Kimi-K2.5": {
+ "input_cost_per_token": 5e-07,
+ "litellm_provider": "together_ai",
+ "max_input_tokens": 256000,
+ "max_output_tokens": 256000,
+ "max_tokens": 256000,
+ "mode": "chat",
+ "output_cost_per_token": 2.8e-06,
+ "source": "https://www.together.ai/models/kimi-k2-5",
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_reasoning": true
+ },
"together_ai/moonshotai/Kimi-K2-Instruct-0905": {
"input_cost_per_token": 1e-06,
"litellm_provider": "together_ai",
@@ -27800,7 +28435,9 @@
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
- "output_cost_per_token": 3e-07
+ "output_cost_per_token": 3e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/alibaba/qwen3-coder": {
"input_cost_per_token": 4e-07,
@@ -27809,7 +28446,9 @@
"max_output_tokens": 66536,
"max_tokens": 66536,
"mode": "chat",
- "output_cost_per_token": 1.6e-06
+ "output_cost_per_token": 1.6e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/amazon/nova-lite": {
"input_cost_per_token": 6e-08,
@@ -27818,7 +28457,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 2.4e-07
+ "output_cost_per_token": 2.4e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/amazon/nova-micro": {
"input_cost_per_token": 3.5e-08,
@@ -27827,7 +28469,9 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 1.4e-07
+ "output_cost_per_token": 1.4e-07,
+ "supports_function_calling": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/amazon/nova-pro": {
"input_cost_per_token": 8e-07,
@@ -27836,7 +28480,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 3.2e-06
+ "output_cost_per_token": 3.2e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/amazon/titan-embed-text-v2": {
"input_cost_per_token": 2e-08,
@@ -27856,7 +28503,11 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 1.25e-06
+ "output_cost_per_token": 1.25e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-3-opus": {
"cache_creation_input_token_cost": 1.875e-05,
@@ -27867,7 +28518,11 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 7.5e-05
+ "output_cost_per_token": 7.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-3.5-haiku": {
"cache_creation_input_token_cost": 1e-06,
@@ -27878,7 +28533,11 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 4e-06
+ "output_cost_per_token": 4e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-3.5-sonnet": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -27889,7 +28548,11 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-3.7-sonnet": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -27900,7 +28563,11 @@
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-4-opus": {
"cache_creation_input_token_cost": 1.875e-05,
@@ -27911,7 +28578,11 @@
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
- "output_cost_per_token": 7.5e-05
+ "output_cost_per_token": 7.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-4-sonnet": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -27922,7 +28593,9 @@
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/cohere/command-a": {
"input_cost_per_token": 2.5e-06,
@@ -27931,7 +28604,9 @@
"max_output_tokens": 8000,
"max_tokens": 8000,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/cohere/command-r": {
"input_cost_per_token": 1.5e-07,
@@ -27940,7 +28615,9 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 6e-07
+ "output_cost_per_token": 6e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/cohere/command-r-plus": {
"input_cost_per_token": 2.5e-06,
@@ -27949,7 +28626,9 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/cohere/embed-v4.0": {
"input_cost_per_token": 1.2e-07,
@@ -27967,7 +28646,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 2.19e-06
+ "output_cost_per_token": 2.19e-06,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/deepseek/deepseek-r1-distill-llama-70b": {
"input_cost_per_token": 7.5e-07,
@@ -27976,7 +28656,10 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 9.9e-07
+ "output_cost_per_token": 9.9e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/deepseek/deepseek-v3": {
"input_cost_per_token": 9e-07,
@@ -27985,7 +28668,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 9e-07
+ "output_cost_per_token": 9e-07,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/google/gemini-2.0-flash": {
"deprecation_date": "2026-03-31",
@@ -27995,7 +28679,11 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 6e-07
+ "output_cost_per_token": 6e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/google/gemini-2.0-flash-lite": {
"deprecation_date": "2026-03-31",
@@ -28005,7 +28693,11 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 3e-07
+ "output_cost_per_token": 3e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/google/gemini-2.5-flash": {
"input_cost_per_token": 3e-07,
@@ -28014,7 +28706,11 @@
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
- "output_cost_per_token": 2.5e-06
+ "output_cost_per_token": 2.5e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/google/gemini-2.5-pro": {
"input_cost_per_token": 2.5e-06,
@@ -28023,7 +28719,11 @@
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/google/gemini-embedding-001": {
"input_cost_per_token": 1.5e-07,
@@ -28041,7 +28741,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 2e-07
+ "output_cost_per_token": 2e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/google/text-embedding-005": {
"input_cost_per_token": 2.5e-08,
@@ -28077,7 +28780,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 7.9e-07
+ "output_cost_per_token": 7.9e-07,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-3-8b": {
"input_cost_per_token": 5e-08,
@@ -28086,7 +28790,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 8e-08
+ "output_cost_per_token": 8e-08,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-3.1-70b": {
"input_cost_per_token": 7.2e-07,
@@ -28095,7 +28800,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 7.2e-07
+ "output_cost_per_token": 7.2e-07,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-3.1-8b": {
"input_cost_per_token": 5e-08,
@@ -28104,7 +28810,9 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 8e-08
+ "output_cost_per_token": 8e-08,
+ "supports_function_calling": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/meta/llama-3.2-11b": {
"input_cost_per_token": 1.6e-07,
@@ -28113,7 +28821,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 1.6e-07
+ "output_cost_per_token": 1.6e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-3.2-1b": {
"input_cost_per_token": 1e-07,
@@ -28131,7 +28842,9 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 1.5e-07
+ "output_cost_per_token": 1.5e-07,
+ "supports_function_calling": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/meta/llama-3.2-90b": {
"input_cost_per_token": 7.2e-07,
@@ -28140,7 +28853,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 7.2e-07
+ "output_cost_per_token": 7.2e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-3.3-70b": {
"input_cost_per_token": 7.2e-07,
@@ -28149,7 +28865,9 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 7.2e-07
+ "output_cost_per_token": 7.2e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-4-maverick": {
"input_cost_per_token": 2e-07,
@@ -28158,7 +28876,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 6e-07
+ "output_cost_per_token": 6e-07,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-4-scout": {
"input_cost_per_token": 1e-07,
@@ -28167,7 +28886,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 3e-07
+ "output_cost_per_token": 3e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/mistral/codestral": {
"input_cost_per_token": 3e-07,
@@ -28176,7 +28898,9 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 9e-07
+ "output_cost_per_token": 9e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/mistral/codestral-embed": {
"input_cost_per_token": 1.5e-07,
@@ -28194,7 +28918,10 @@
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
- "output_cost_per_token": 2.8e-07
+ "output_cost_per_token": 2.8e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/mistral/magistral-medium": {
"input_cost_per_token": 2e-06,
@@ -28203,7 +28930,10 @@
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
- "output_cost_per_token": 5e-06
+ "output_cost_per_token": 5e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/mistral/magistral-small": {
"input_cost_per_token": 5e-07,
@@ -28212,7 +28942,8 @@
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
- "output_cost_per_token": 1.5e-06
+ "output_cost_per_token": 1.5e-06,
+ "supports_function_calling": true
},
"vercel_ai_gateway/mistral/ministral-3b": {
"input_cost_per_token": 4e-08,
@@ -28221,7 +28952,9 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 4e-08
+ "output_cost_per_token": 4e-08,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/mistral/ministral-8b": {
"input_cost_per_token": 1e-07,
@@ -28230,7 +28963,10 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 1e-07
+ "output_cost_per_token": 1e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/mistral/mistral-embed": {
"input_cost_per_token": 1e-07,
@@ -28248,7 +28984,9 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 6e-06
+ "output_cost_per_token": 6e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/mistral/mistral-saba-24b": {
"input_cost_per_token": 7.9e-07,
@@ -28266,7 +29004,10 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 3e-07
+ "output_cost_per_token": 3e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/mistral/mixtral-8x22b-instruct": {
"input_cost_per_token": 1.2e-06,
@@ -28275,7 +29016,8 @@
"max_output_tokens": 2048,
"max_tokens": 2048,
"mode": "chat",
- "output_cost_per_token": 1.2e-06
+ "output_cost_per_token": 1.2e-06,
+ "supports_function_calling": true
},
"vercel_ai_gateway/mistral/pixtral-12b": {
"input_cost_per_token": 1.5e-07,
@@ -28284,7 +29026,11 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 1.5e-07
+ "output_cost_per_token": 1.5e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/mistral/pixtral-large": {
"input_cost_per_token": 2e-06,
@@ -28293,7 +29039,11 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 6e-06
+ "output_cost_per_token": 6e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/moonshotai/kimi-k2": {
"input_cost_per_token": 5.5e-07,
@@ -28302,7 +29052,9 @@
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
- "output_cost_per_token": 2.2e-06
+ "output_cost_per_token": 2.2e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/morph/morph-v3-fast": {
"input_cost_per_token": 8e-07,
@@ -28329,7 +29081,9 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 1.5e-06
+ "output_cost_per_token": 1.5e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/openai/gpt-3.5-turbo-instruct": {
"input_cost_per_token": 1.5e-06,
@@ -28347,7 +29101,10 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 3e-05
+ "output_cost_per_token": 3e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/openai/gpt-4.1": {
"cache_creation_input_token_cost": 0.0,
@@ -28358,7 +29115,11 @@
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 8e-06
+ "output_cost_per_token": 8e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/gpt-4.1-mini": {
"cache_creation_input_token_cost": 0.0,
@@ -28369,7 +29130,11 @@
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 1.6e-06
+ "output_cost_per_token": 1.6e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/gpt-4.1-nano": {
"cache_creation_input_token_cost": 0.0,
@@ -28380,7 +29145,11 @@
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 4e-07
+ "output_cost_per_token": 4e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/gpt-4o": {
"cache_creation_input_token_cost": 0.0,
@@ -28391,7 +29160,11 @@
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/gpt-4o-mini": {
"cache_creation_input_token_cost": 0.0,
@@ -28402,7 +29175,11 @@
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
- "output_cost_per_token": 6e-07
+ "output_cost_per_token": 6e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/o1": {
"cache_creation_input_token_cost": 0.0,
@@ -28413,7 +29190,11 @@
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
- "output_cost_per_token": 6e-05
+ "output_cost_per_token": 6e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/o3": {
"cache_creation_input_token_cost": 0.0,
@@ -28424,7 +29205,11 @@
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
- "output_cost_per_token": 8e-06
+ "output_cost_per_token": 8e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/o3-mini": {
"cache_creation_input_token_cost": 0.0,
@@ -28435,7 +29220,10 @@
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
- "output_cost_per_token": 4.4e-06
+ "output_cost_per_token": 4.4e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/o4-mini": {
"cache_creation_input_token_cost": 0.0,
@@ -28446,7 +29234,11 @@
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
- "output_cost_per_token": 4.4e-06
+ "output_cost_per_token": 4.4e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/text-embedding-3-large": {
"input_cost_per_token": 1.3e-07,
@@ -28518,7 +29310,10 @@
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/vercel/v0-1.5-md": {
"input_cost_per_token": 3e-06,
@@ -28527,7 +29322,10 @@
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-2": {
"input_cost_per_token": 2e-06,
@@ -28536,7 +29334,9 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-2-vision": {
"input_cost_per_token": 2e-06,
@@ -28545,7 +29345,10 @@
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-3": {
"input_cost_per_token": 3e-06,
@@ -28554,7 +29357,9 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-3-fast": {
"input_cost_per_token": 5e-06,
@@ -28563,7 +29368,8 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 2.5e-05
+ "output_cost_per_token": 2.5e-05,
+ "supports_function_calling": true
},
"vercel_ai_gateway/xai/grok-3-mini": {
"input_cost_per_token": 3e-07,
@@ -28572,7 +29378,9 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 5e-07
+ "output_cost_per_token": 5e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-3-mini-fast": {
"input_cost_per_token": 6e-07,
@@ -28581,7 +29389,9 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 4e-06
+ "output_cost_per_token": 4e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-4": {
"input_cost_per_token": 3e-06,
@@ -28590,7 +29400,9 @@
"max_output_tokens": 256000,
"max_tokens": 256000,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/zai/glm-4.5": {
"input_cost_per_token": 6e-07,
@@ -28599,7 +29411,9 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 2.2e-06
+ "output_cost_per_token": 2.2e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/zai/glm-4.5-air": {
"input_cost_per_token": 2e-07,
@@ -28608,7 +29422,9 @@
"max_output_tokens": 96000,
"max_tokens": 96000,
"mode": "chat",
- "output_cost_per_token": 1.1e-06
+ "output_cost_per_token": 1.1e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/zai/glm-4.6": {
"litellm_provider": "vercel_ai_gateway",
@@ -28676,7 +29492,9 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_streaming": true,
+ "supports_vision": true
},
"vertex_ai/claude-3-5-sonnet": {
"input_cost_per_token": 3e-06,
@@ -28947,7 +29765,38 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_streaming": true
+ },
+ "vertex_ai/claude-opus-4-6": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "vertex_ai-anthropic_models",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
},
"vertex_ai/claude-sonnet-4-5": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -28999,7 +29848,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_native_streaming": true
},
"vertex_ai/claude-opus-4@20250514": {
"cache_creation_input_token_cost": 1.875e-05,
@@ -29281,6 +30131,21 @@
"output_cost_per_token_batches": 6e-06,
"source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
},
+ "vertex_ai/deep-research-pro-preview-12-2025": {
+ "input_cost_per_image": 0.0011,
+ "input_cost_per_token": 2e-06,
+ "input_cost_per_token_batches": 1e-06,
+ "litellm_provider": "vertex_ai-language-models",
+ "max_input_tokens": 65536,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "image_generation",
+ "output_cost_per_image": 0.134,
+ "output_cost_per_image_token": 0.00012,
+ "output_cost_per_token": 1.2e-05,
+ "output_cost_per_token_batches": 6e-06,
+ "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
+ },
"vertex_ai/imagegeneration@006": {
"litellm_provider": "vertex_ai-image-models",
"mode": "image_generation",
@@ -29770,6 +30635,9 @@
"mode": "chat",
"output_cost_per_token": 1e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
+ "supported_regions": [
+ "global"
+ ],
"supports_function_calling": true,
"supports_tool_choice": true
},
@@ -29782,6 +30650,9 @@
"mode": "chat",
"output_cost_per_token": 4e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
+ "supported_regions": [
+ "global"
+ ],
"supports_function_calling": true,
"supports_tool_choice": true
},
@@ -29794,6 +30665,9 @@
"mode": "chat",
"output_cost_per_token": 1.2e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
+ "supported_regions": [
+ "global"
+ ],
"supports_function_calling": true,
"supports_tool_choice": true
},
@@ -29806,6 +30680,9 @@
"mode": "chat",
"output_cost_per_token": 1.2e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
+ "supported_regions": [
+ "global"
+ ],
"supports_function_calling": true,
"supports_tool_choice": true
},
@@ -34754,4 +35631,4 @@
"output_cost_per_token": 0,
"supports_reasoning": true
}
-}
\ No newline at end of file
+}
diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py
index 49d6ac7d898..7e70b5baae4 100644
--- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py
+++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py
@@ -387,6 +387,9 @@ class MCPRequestHandler:
user_api_key_cache,
)
+ verbose_logger.debug(
+ f"MCP team permission lookup: team_id={user_api_key_auth.team_id if user_api_key_auth else None}"
+ )
if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client:
return None
diff --git a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py
new file mode 100644
index 00000000000..e5cb6a0098d
--- /dev/null
+++ b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py
@@ -0,0 +1,250 @@
+"""
+Semantic MCP Tool Filtering using semantic-router
+
+Filters MCP tools semantically for /chat/completions and /responses endpoints.
+"""
+from typing import TYPE_CHECKING, Any, Dict, List, Optional
+
+from litellm._logging import verbose_logger
+
+if TYPE_CHECKING:
+ from semantic_router.routers import SemanticRouter
+
+ from litellm.router import Router
+
+
+class SemanticMCPToolFilter:
+ """Filters MCP tools using semantic similarity to reduce context window size."""
+
+ def __init__(
+ self,
+ embedding_model: str,
+ litellm_router_instance: "Router",
+ top_k: int = 10,
+ similarity_threshold: float = 0.3,
+ enabled: bool = True,
+ ):
+ """
+ Initialize the semantic tool filter.
+
+ Args:
+ embedding_model: Model to use for embeddings (e.g., "text-embedding-3-small")
+ litellm_router_instance: Router instance for embedding generation
+ top_k: Maximum number of tools to return
+ similarity_threshold: Minimum similarity score for filtering
+ enabled: Whether filtering is enabled
+ """
+ self.enabled = enabled
+ self.top_k = top_k
+ self.similarity_threshold = similarity_threshold
+ self.embedding_model = embedding_model
+ self.router_instance = litellm_router_instance
+ self.tool_router: Optional["SemanticRouter"] = None
+ self._tool_map: Dict[str, Any] = {} # MCPTool objects or OpenAI function dicts
+
+ async def build_router_from_mcp_registry(self) -> None:
+ """Build semantic router from all MCP tools in the registry (no auth checks)."""
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
+ global_mcp_server_manager,
+ )
+
+ try:
+ # Get all servers from registry without auth checks
+ registry = global_mcp_server_manager.get_registry()
+ if not registry:
+ verbose_logger.warning("MCP registry is empty")
+ self.tool_router = None
+ return
+
+ # Fetch tools from all servers in parallel
+ all_tools = []
+ for server_id, server in registry.items():
+ try:
+ tools = await global_mcp_server_manager.get_tools_for_server(server_id)
+ all_tools.extend(tools)
+ except Exception as e:
+ verbose_logger.warning(f"Failed to fetch tools from server {server_id}: {e}")
+ continue
+
+ if not all_tools:
+ verbose_logger.warning("No MCP tools found in registry")
+ self.tool_router = None
+ return
+
+ verbose_logger.info(f"Fetched {len(all_tools)} tools from {len(registry)} MCP servers")
+ self._build_router(all_tools)
+
+ except Exception as e:
+ verbose_logger.error(f"Failed to build router from MCP registry: {e}")
+ self.tool_router = None
+ raise
+
+ def _extract_tool_info(self, tool) -> tuple[str, str]:
+ """Extract name and description from MCP tool or OpenAI function dict."""
+ name: str
+ description: str
+
+ if isinstance(tool, dict):
+ # OpenAI function format
+ name = tool.get("name", "")
+ description = tool.get("description", name)
+ else:
+ # MCPTool object
+ name = str(tool.name)
+ description = str(tool.description) if tool.description else str(tool.name)
+
+ return name, description
+
+ def _build_router(self, tools: List) -> None:
+ """Build semantic router with tools (MCPTool objects or OpenAI function dicts)."""
+ from semantic_router.routers import SemanticRouter
+ from semantic_router.routers.base import Route
+
+ from litellm.router_strategy.auto_router.litellm_encoder import (
+ LiteLLMRouterEncoder,
+ )
+
+ if not tools:
+ self.tool_router = None
+ return
+
+ try:
+ # Convert tools to routes
+ routes = []
+ self._tool_map = {}
+
+ for tool in tools:
+ name, description = self._extract_tool_info(tool)
+ self._tool_map[name] = tool
+
+ routes.append(
+ Route(
+ name=name,
+ description=description,
+ utterances=[description],
+ score_threshold=self.similarity_threshold,
+ )
+ )
+
+ self.tool_router = SemanticRouter(
+ routes=routes,
+ encoder=LiteLLMRouterEncoder(
+ litellm_router_instance=self.router_instance,
+ model_name=self.embedding_model,
+ score_threshold=self.similarity_threshold,
+ ),
+ auto_sync="local",
+ )
+
+ verbose_logger.info(
+ f"Built semantic router with {len(routes)} tools"
+ )
+
+ except Exception as e:
+ verbose_logger.error(f"Failed to build semantic router: {e}")
+ self.tool_router = None
+ raise
+
+ async def filter_tools(
+ self,
+ query: str,
+ available_tools: List[Any],
+ top_k: Optional[int] = None,
+ ) -> List[Any]:
+ """
+ Filter tools semantically based on query.
+
+ Args:
+ query: User query to match against tools
+ available_tools: Full list of available MCP tools
+ top_k: Override default top_k (optional)
+
+ Returns:
+ Filtered and ordered list of tools (up to top_k)
+ """
+ # Early returns for cases where we can't/shouldn't filter
+ if not self.enabled:
+ return available_tools
+
+ if not available_tools:
+ return available_tools
+
+ if not query or not query.strip():
+ return available_tools
+
+ # Router should be built on startup - if not, something went wrong
+ if self.tool_router is None:
+ verbose_logger.warning("Router not initialized - was build_router_from_mcp_registry() called on startup?")
+ return available_tools
+
+ # Run semantic filtering
+ try:
+ limit = top_k or self.top_k
+ matches = self.tool_router(text=query, limit=limit)
+ matched_tool_names = self._extract_tool_names_from_matches(matches)
+
+ if not matched_tool_names:
+ return available_tools
+
+ return self._get_tools_by_names(matched_tool_names, available_tools)
+
+ except Exception as e:
+ verbose_logger.error(f"Semantic tool filter failed: {e}", exc_info=True)
+ return available_tools
+
+ def _extract_tool_names_from_matches(self, matches) -> List[str]:
+ """Extract tool names from semantic router match results."""
+ if not matches:
+ return []
+
+ # Handle single match
+ if hasattr(matches, "name") and matches.name:
+ return [matches.name]
+
+ # Handle list of matches
+ if isinstance(matches, list):
+ return [m.name for m in matches if hasattr(m, "name") and m.name]
+
+ return []
+
+ def _get_tools_by_names(
+ self, tool_names: List[str], available_tools: List[Any]
+ ) -> List[Any]:
+ """Get tools from available_tools by their names, preserving order."""
+ # Match tools from available_tools (preserves format - dict or MCPTool)
+ matched_tools = []
+ for tool in available_tools:
+ tool_name, _ = self._extract_tool_info(tool)
+ if tool_name in tool_names:
+ matched_tools.append(tool)
+
+ # Reorder to match semantic router's ordering
+ tool_map = {self._extract_tool_info(t)[0]: t for t in matched_tools}
+ return [tool_map[name] for name in tool_names if name in tool_map]
+
+ def extract_user_query(self, messages: List[Dict[str, Any]]) -> str:
+ """
+ Extract user query from messages for /chat/completions or /responses.
+
+ Args:
+ messages: List of message dictionaries (from 'messages' or 'input' field)
+
+ Returns:
+ Extracted query string
+ """
+ for msg in reversed(messages):
+ if msg.get("role") == "user":
+ content = msg.get("content", "")
+
+ if isinstance(content, str):
+ return content
+
+ if isinstance(content, list):
+ texts = [
+ block.get("text", "") if isinstance(block, dict) else str(block)
+ for block in content
+ if isinstance(block, (dict, str))
+ ]
+ return " ".join(texts)
+
+ return ""
diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py
index 6d54c3871e5..79cd88227a9 100644
--- a/litellm/proxy/_experimental/mcp_server/server.py
+++ b/litellm/proxy/_experimental/mcp_server/server.py
@@ -1840,6 +1840,43 @@ if MCP_AVAILABLE:
raw_headers,
)
+ def _strip_stale_mcp_session_header(
+ scope: Scope,
+ mgr: "StreamableHTTPSessionManager",
+ ) -> None:
+ """
+ Strip stale ``mcp-session-id`` headers so the session manager
+ creates a fresh session instead of returning 404 "Session not found".
+
+ When clients like VSCode reconnect after a reload they may resend a
+ session id that has already been cleaned up. Rather than letting the
+ SDK return a 404 error loop, we detect the stale id and remove the
+ header so a brand-new session is created transparently.
+
+ Fixes https://github.com/BerriAI/litellm/issues/20292
+ """
+ _mcp_session_header = b"mcp-session-id"
+ _session_id: Optional[str] = None
+ for header_name, header_value in scope.get("headers", []):
+ if header_name == _mcp_session_header:
+ _session_id = header_value.decode("utf-8", errors="replace")
+ break
+
+ if _session_id is None:
+ return
+
+ known_sessions = getattr(mgr, "_server_instances", None)
+ if known_sessions is not None and _session_id not in known_sessions:
+ verbose_logger.warning(
+ "MCP session ID '%s' not found in active sessions. "
+ "Stripping stale header to force new session creation.",
+ _session_id,
+ )
+ scope["headers"] = [
+ (k, v) for k, v in scope["headers"]
+ if k != _mcp_session_header
+ ]
+
async def handle_streamable_http_mcp(
scope: Scope, receive: Receive, send: Send
) -> None:
@@ -1896,6 +1933,8 @@ if MCP_AVAILABLE:
# Give it a moment to start up
await asyncio.sleep(0.1)
+ _strip_stale_mcp_session_header(scope, session_manager)
+
await session_manager.handle_request(scope, receive, send)
except Exception as e:
raise e
diff --git a/litellm/proxy/_experimental/out/404.html b/litellm/proxy/_experimental/out/404/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/404.html
rename to litellm/proxy/_experimental/out/404/index.html
diff --git a/litellm/proxy/_experimental/out/api-reference.html b/litellm/proxy/_experimental/out/api-reference/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/api-reference.html
rename to litellm/proxy/_experimental/out/api-reference/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/api-playground.html b/litellm/proxy/_experimental/out/experimental/api-playground/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/api-playground.html
rename to litellm/proxy/_experimental/out/experimental/api-playground/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/budgets.html b/litellm/proxy/_experimental/out/experimental/budgets/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/budgets.html
rename to litellm/proxy/_experimental/out/experimental/budgets/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/caching.html b/litellm/proxy/_experimental/out/experimental/caching/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/caching.html
rename to litellm/proxy/_experimental/out/experimental/caching/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/claude-code-plugins.html b/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/claude-code-plugins.html
rename to litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/old-usage.html b/litellm/proxy/_experimental/out/experimental/old-usage/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/old-usage.html
rename to litellm/proxy/_experimental/out/experimental/old-usage/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/prompts.html b/litellm/proxy/_experimental/out/experimental/prompts/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/prompts.html
rename to litellm/proxy/_experimental/out/experimental/prompts/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/tag-management.html b/litellm/proxy/_experimental/out/experimental/tag-management/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/tag-management.html
rename to litellm/proxy/_experimental/out/experimental/tag-management/index.html
diff --git a/litellm/proxy/_experimental/out/guardrails.html b/litellm/proxy/_experimental/out/guardrails/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/guardrails.html
rename to litellm/proxy/_experimental/out/guardrails/index.html
diff --git a/litellm/proxy/_experimental/out/login.html b/litellm/proxy/_experimental/out/login/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/login.html
rename to litellm/proxy/_experimental/out/login/index.html
diff --git a/litellm/proxy/_experimental/out/logs.html b/litellm/proxy/_experimental/out/logs/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/logs.html
rename to litellm/proxy/_experimental/out/logs/index.html
diff --git a/litellm/proxy/_experimental/out/mcp/oauth/callback.html b/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/mcp/oauth/callback.html
rename to litellm/proxy/_experimental/out/mcp/oauth/callback/index.html
diff --git a/litellm/proxy/_experimental/out/model-hub.html b/litellm/proxy/_experimental/out/model-hub/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/model-hub.html
rename to litellm/proxy/_experimental/out/model-hub/index.html
diff --git a/litellm/proxy/_experimental/out/model_hub.html b/litellm/proxy/_experimental/out/model_hub/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/model_hub.html
rename to litellm/proxy/_experimental/out/model_hub/index.html
diff --git a/litellm/proxy/_experimental/out/model_hub_table.html b/litellm/proxy/_experimental/out/model_hub_table/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/model_hub_table.html
rename to litellm/proxy/_experimental/out/model_hub_table/index.html
diff --git a/litellm/proxy/_experimental/out/models-and-endpoints.html b/litellm/proxy/_experimental/out/models-and-endpoints/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/models-and-endpoints.html
rename to litellm/proxy/_experimental/out/models-and-endpoints/index.html
diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/onboarding.html
rename to litellm/proxy/_experimental/out/onboarding/index.html
diff --git a/litellm/proxy/_experimental/out/organizations.html b/litellm/proxy/_experimental/out/organizations/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/organizations.html
rename to litellm/proxy/_experimental/out/organizations/index.html
diff --git a/litellm/proxy/_experimental/out/playground.html b/litellm/proxy/_experimental/out/playground/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/playground.html
rename to litellm/proxy/_experimental/out/playground/index.html
diff --git a/litellm/proxy/_experimental/out/policies.html b/litellm/proxy/_experimental/out/policies/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/policies.html
rename to litellm/proxy/_experimental/out/policies/index.html
diff --git a/litellm/proxy/_experimental/out/settings/admin-settings.html b/litellm/proxy/_experimental/out/settings/admin-settings/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/settings/admin-settings.html
rename to litellm/proxy/_experimental/out/settings/admin-settings/index.html
diff --git a/litellm/proxy/_experimental/out/settings/logging-and-alerts.html b/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/settings/logging-and-alerts.html
rename to litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html
diff --git a/litellm/proxy/_experimental/out/settings/router-settings.html b/litellm/proxy/_experimental/out/settings/router-settings/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/settings/router-settings.html
rename to litellm/proxy/_experimental/out/settings/router-settings/index.html
diff --git a/litellm/proxy/_experimental/out/settings/ui-theme.html b/litellm/proxy/_experimental/out/settings/ui-theme/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/settings/ui-theme.html
rename to litellm/proxy/_experimental/out/settings/ui-theme/index.html
diff --git a/litellm/proxy/_experimental/out/teams.html b/litellm/proxy/_experimental/out/teams/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/teams.html
rename to litellm/proxy/_experimental/out/teams/index.html
diff --git a/litellm/proxy/_experimental/out/test-key.html b/litellm/proxy/_experimental/out/test-key/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/test-key.html
rename to litellm/proxy/_experimental/out/test-key/index.html
diff --git a/litellm/proxy/_experimental/out/tools/mcp-servers.html b/litellm/proxy/_experimental/out/tools/mcp-servers/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/tools/mcp-servers.html
rename to litellm/proxy/_experimental/out/tools/mcp-servers/index.html
diff --git a/litellm/proxy/_experimental/out/tools/vector-stores.html b/litellm/proxy/_experimental/out/tools/vector-stores/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/tools/vector-stores.html
rename to litellm/proxy/_experimental/out/tools/vector-stores/index.html
diff --git a/litellm/proxy/_experimental/out/usage.html b/litellm/proxy/_experimental/out/usage/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/usage.html
rename to litellm/proxy/_experimental/out/usage/index.html
diff --git a/litellm/proxy/_experimental/out/users.html b/litellm/proxy/_experimental/out/users/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/users.html
rename to litellm/proxy/_experimental/out/users/index.html
diff --git a/litellm/proxy/_experimental/out/virtual-keys.html b/litellm/proxy/_experimental/out/virtual-keys/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/virtual-keys.html
rename to litellm/proxy/_experimental/out/virtual-keys/index.html
diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml
index 13eeae14485..6f527e268b2 100644
--- a/litellm/proxy/_new_secret_config.yaml
+++ b/litellm/proxy/_new_secret_config.yaml
@@ -14,3 +14,14 @@ model_list:
litellm_params:
model: openai/gpt-4.1-mini
+guardrails:
+ - guardrail_name: redact-ssn
+ litellm_params:
+ guardrail: custom_code
+ mode: pre_call
+ custom_code: |
+ def apply_guardrail(inputs, request_data, input_type):
+ for text in inputs["texts"]:
+ if regex_match(text, r"\d{3}-\d{2}-\d{4}"):
+ return block("SSN detected in message")
+ return allow()
\ No newline at end of file
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index bf99347ef6e..f38f94f4c98 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -228,6 +228,7 @@ class KeyManagementRoutes(str, enum.Enum):
KEY_BLOCK = "/key/block"
KEY_UNBLOCK = "/key/unblock"
KEY_BULK_UPDATE = "/key/bulk_update"
+ KEY_RESET_SPEND = "/key/{key_id}/reset_spend"
# info and health routes
KEY_INFO = "/key/info"
@@ -987,6 +988,10 @@ class RegenerateKeyRequest(GenerateKeyRequest):
new_master_key: Optional[str] = None
+class ResetSpendRequest(LiteLLMPydanticObjectBase):
+ reset_to: float
+
+
class KeyRequest(LiteLLMPydanticObjectBase):
keys: Optional[List[str]] = None
key_aliases: Optional[List[str]] = None
@@ -1483,6 +1488,7 @@ class TeamBase(LiteLLMPydanticObjectBase):
# Budget fields
max_budget: Optional[float] = None
+ soft_budget: Optional[float] = None
budget_duration: Optional[str] = None
models: list = []
@@ -1554,6 +1560,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
tpm_limit: Optional[int] = None
rpm_limit: Optional[int] = None
max_budget: Optional[float] = None
+ soft_budget: Optional[float] = None
models: Optional[list] = None
blocked: Optional[bool] = None
budget_duration: Optional[str] = None
@@ -2155,10 +2162,6 @@ class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase):
rotation_interval: Optional[str] = None # How often to rotate (e.g., "30d", "90d")
last_rotation_at: Optional[datetime] = None # When this key was last rotated
key_rotation_at: Optional[datetime] = None # When this key should next be rotated
- router_settings: Optional[
- Dict
- ] = None # Router settings for this key (Key > Team > Global precedence)
-
model_config = ConfigDict(protected_namespaces=())
@@ -2187,6 +2190,7 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken):
team_tpm_limit: Optional[int] = None
team_rpm_limit: Optional[int] = None
team_max_budget: Optional[float] = None
+ team_soft_budget: Optional[float] = None
team_models: List = []
team_blocked: bool = False
soft_budget: Optional[float] = None
@@ -2645,6 +2649,10 @@ class CallInfo(LiteLLMPydanticObjectBase):
projected_exceeded_date: Optional[str] = None
projected_spend: Optional[float] = None
event_group: Litellm_EntityType
+ alert_emails: Optional[List[str]] = Field(
+ default=None,
+ description="Additional email addresses to send alerts to (e.g., from team metadata)",
+ )
class WebhookEvent(CallInfo):
@@ -3672,7 +3680,7 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
team_id_upsert: bool = False
team_ids_jwt_field: Optional[str] = None
upsert_sso_user_to_team: bool = False
- team_allowed_routes: List[str] = ["openai_routes", "info_routes"]
+ team_allowed_routes: List[str] = ["openai_routes", "info_routes", "mcp_routes"]
team_id_default: Optional[str] = Field(
default=None,
description="If no team_id given, default permissions/spend-tracking to this team.s",
diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py
new file mode 100644
index 00000000000..cb277d44ee9
--- /dev/null
+++ b/litellm/proxy/agent_endpoints/a2a_routing.py
@@ -0,0 +1,53 @@
+"""
+A2A Agent Routing
+
+Handles routing for A2A agents (models with "a2a/" prefix).
+Looks up agents in the registry and injects their API base URL.
+"""
+
+from typing import Any, Optional
+
+import litellm
+from litellm._logging import verbose_proxy_logger
+
+
+def route_a2a_agent_request(data: dict, route_type: str) -> Optional[Any]:
+ """
+ Route A2A agent requests directly to litellm with injected API base.
+
+ Returns None if not an A2A request (allows normal routing to continue).
+ """
+ # Import here to avoid circular imports
+ from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
+ from litellm.proxy.route_llm_request import (
+ ROUTE_ENDPOINT_MAPPING,
+ ProxyModelNotFoundError,
+ )
+
+ model_name = data.get("model", "")
+
+ # Check if this is an A2A agent request
+ if not isinstance(model_name, str) or not model_name.startswith("a2a/"):
+ return None
+
+ # Extract agent name (e.g., "a2a/my-agent" -> "my-agent")
+ agent_name = model_name[4:]
+
+ # Look up agent in registry
+ agent = global_agent_registry.get_agent_by_name(agent_name)
+ if agent is None:
+ verbose_proxy_logger.error(f"[A2A] Agent '{agent_name}' not found in registry")
+ route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type)
+ raise ProxyModelNotFoundError(route=route_name, model_name=model_name)
+
+ # Get API base URL from agent config
+ if not agent.agent_card_params or "url" not in agent.agent_card_params:
+ verbose_proxy_logger.error(f"[A2A] Agent '{agent_name}' has no URL configured")
+ route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type)
+ raise ProxyModelNotFoundError(route=route_name, model_name=model_name)
+
+ # Inject API base and route to litellm
+ data["api_base"] = agent.agent_card_params["url"]
+ verbose_proxy_logger.debug(f"[A2A] Routing {model_name} to {data['api_base']}")
+
+ return getattr(litellm, f"{route_type}")(**data)
diff --git a/litellm/proxy/agent_endpoints/model_list_helpers.py b/litellm/proxy/agent_endpoints/model_list_helpers.py
new file mode 100644
index 00000000000..c640300bb8c
--- /dev/null
+++ b/litellm/proxy/agent_endpoints/model_list_helpers.py
@@ -0,0 +1,96 @@
+"""
+Helper functions for appending A2A agents to model lists.
+
+Used by proxy model endpoints to make agents appear in UI alongside models.
+"""
+from typing import List
+
+from litellm._logging import verbose_proxy_logger
+from litellm.proxy._types import UserAPIKeyAuth
+from litellm.types.proxy.management_endpoints.model_management_endpoints import (
+ ModelGroupInfoProxy,
+)
+
+
+async def append_agents_to_model_group(
+ model_groups: List[ModelGroupInfoProxy],
+ user_api_key_dict: UserAPIKeyAuth,
+) -> List[ModelGroupInfoProxy]:
+ """
+ Append A2A agents to model groups list for UI display.
+
+ Converts agents to model format with "a2a/" naming
+ so they appear in playground and work with LiteLLM routing.
+ """
+ try:
+ from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
+ from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
+ AgentRequestHandler,
+ )
+
+ allowed_agent_ids = await AgentRequestHandler.get_allowed_agents(
+ user_api_key_auth=user_api_key_dict
+ )
+
+ for agent_id in allowed_agent_ids:
+ agent = global_agent_registry.get_agent_by_id(agent_id)
+ if agent is not None:
+ model_groups.append(
+ ModelGroupInfoProxy(
+ model_group=f"a2a/{agent.agent_name}",
+ mode="chat",
+ providers=["a2a"],
+ )
+ )
+ except Exception as e:
+ verbose_proxy_logger.debug(
+ f"Error appending agents to model_group/info: {e}"
+ )
+
+ return model_groups
+
+
+async def append_agents_to_model_info(
+ models: List[dict],
+ user_api_key_dict: UserAPIKeyAuth,
+) -> List[dict]:
+ """
+ Append A2A agents to model info list for UI display.
+
+ Converts agents to model format with "a2a/" naming
+ so they appear in models page and work with LiteLLM routing.
+ """
+ try:
+ from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
+ from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
+ AgentRequestHandler,
+ )
+
+ allowed_agent_ids = await AgentRequestHandler.get_allowed_agents(
+ user_api_key_auth=user_api_key_dict
+ )
+
+ for agent_id in allowed_agent_ids:
+ agent = global_agent_registry.get_agent_by_id(agent_id)
+ if agent is not None:
+ models.append({
+ "model_name": f"a2a/{agent.agent_name}",
+ "litellm_params": {
+ "model": f"a2a/{agent.agent_name}",
+ "custom_llm_provider": "a2a",
+ },
+ "model_info": {
+ "id": agent.agent_id,
+ "mode": "chat",
+ "db_model": True,
+ "created_by": agent.created_by,
+ "created_at": agent.created_at,
+ "updated_at": agent.updated_at,
+ },
+ })
+ except Exception as e:
+ verbose_proxy_logger.debug(
+ f"Error appending agents to v2/model/info: {e}"
+ )
+
+ return models
diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py
index 0033deb0766..77bb1f53e62 100644
--- a/litellm/proxy/anthropic_endpoints/endpoints.py
+++ b/litellm/proxy/anthropic_endpoints/endpoints.py
@@ -80,6 +80,7 @@ async def anthropic_response( # noqa: PLR0915
# Create Anthropic-formatted response with violation message
import uuid
+
from litellm.types.utils import AnthropicMessagesResponse
_anthropic_response = AnthropicMessagesResponse(
@@ -240,3 +241,19 @@ async def count_tokens(
raise HTTPException(
status_code=500, detail={"error": f"Internal server error: {str(e)}"}
)
+
+
+@router.post(
+ "/api/event_logging/batch",
+ tags=["[beta] Anthropic Event Logging"],
+)
+async def event_logging_batch(
+ request: Request,
+):
+ """
+ Stubbed endpoint for Anthropic event logging batch requests.
+
+ This endpoint accepts event logging requests but does nothing with them.
+ It exists to prevent 404 errors from Claude Code clients that send telemetry.
+ """
+ return {"status": "ok"}
diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py
index e0b056d450f..c6093172932 100644
--- a/litellm/proxy/auth/auth_checks.py
+++ b/litellm/proxy/auth/auth_checks.py
@@ -75,6 +75,97 @@ db_cache_expiry = DEFAULT_IN_MEMORY_TTL # refresh every 5s
all_routes = LiteLLMRoutes.openai_routes.value + LiteLLMRoutes.management_routes.value
+def _log_budget_lookup_failure(entity: str, error: Exception) -> None:
+ """
+ Log a warning when budget lookup fails; cache will not be populated.
+
+ Skips logging for expected "user not found" cases (bare Exception from
+ get_user_object when user_id_upsert=False). Adds a schema migration hint
+ when the error appears schema-related.
+ """
+ # Skip logging for expected "user not found" - not caching is correct
+ if str(error) == "" and type(error).__name__ == "Exception":
+ return
+ err_str = str(error).lower()
+ hint = ""
+ if any(
+ x in err_str
+ for x in ("column", "schema", "does not exist", "prisma", "migrate")
+ ):
+ hint = " Run `prisma db push` or `prisma migrate deploy` to fix schema mismatches."
+ verbose_proxy_logger.error(
+ f"Budget lookup failed for {entity}; cache will not be populated. "
+ f"Each request will hit the database. Error: {error}.{hint}"
+ )
+
+def _is_model_cost_zero(
+ model: Optional[Union[str, List[str]]], llm_router: Optional[Router]
+) -> bool:
+ """
+ Check if a model has zero cost (no configured pricing).
+
+ Uses the router's get_model_group_info method to get pricing information.
+
+ Args:
+ model: The model name or list of model names
+ llm_router: The LiteLLM router instance
+
+ Returns:
+ bool: True if all costs for the model are zero, False otherwise
+ """
+ if model is None or llm_router is None:
+ return False
+
+ # Handle list of models
+ model_list = [model] if isinstance(model, str) else model
+
+ for model_name in model_list:
+ try:
+ # Use router's get_model_group_info method directly for better reliability
+ model_group_info = llm_router.get_model_group_info(model_group=model_name)
+
+ if model_group_info is None:
+ # Model not found or no pricing info available
+ # Conservative approach: assume it has cost
+ verbose_proxy_logger.debug(
+ f"No model group info found for {model_name}, assuming it has cost"
+ )
+ return False
+
+ # Check costs for this model
+ # Only allow bypass if BOTH costs are explicitly set to 0 (not None)
+ input_cost = model_group_info.input_cost_per_token
+ output_cost = model_group_info.output_cost_per_token
+
+ # If costs are not explicitly configured (None), assume it has cost
+ if input_cost is None or output_cost is None:
+ verbose_proxy_logger.debug(
+ f"Model {model_name} has undefined cost (input: {input_cost}, output: {output_cost}), assuming it has cost"
+ )
+ return False
+
+ # If either cost is non-zero, return False
+ if input_cost > 0 or output_cost > 0:
+ verbose_proxy_logger.debug(
+ f"Model {model_name} has non-zero cost (input: {input_cost}, output: {output_cost})"
+ )
+ return False
+
+ # This model has zero cost explicitly configured
+ verbose_proxy_logger.debug(
+ f"Model {model_name} has zero cost explicitly configured (input: {input_cost}, output: {output_cost})"
+ )
+
+ except Exception as e:
+ # If we can't determine the cost, assume it has cost (conservative approach)
+ verbose_proxy_logger.debug(
+ f"Error checking cost for model {model_name}: {str(e)}, assuming it has cost"
+ )
+ return False
+
+ # All models checked have zero cost
+ return True
+
async def common_checks(
request_body: dict,
@@ -88,6 +179,7 @@ async def common_checks(
proxy_logging_obj: ProxyLogging,
valid_token: Optional[UserAPIKeyAuth],
request: Request,
+ skip_budget_checks: bool = False,
) -> bool:
"""
Common checks across jwt + key-based auth.
@@ -139,64 +231,73 @@ async def common_checks(
user_object=user_object,
)
- # 3. If team is in budget
- await _team_max_budget_check(
- team_object=team_object,
- proxy_logging_obj=proxy_logging_obj,
- valid_token=valid_token,
- )
+ # If this is a free model, skip all budget checks
+ if not skip_budget_checks:
+ # 3. If team is in budget
+ await _team_max_budget_check(
+ team_object=team_object,
+ proxy_logging_obj=proxy_logging_obj,
+ valid_token=valid_token,
+ )
- # 3.1. If organization is in budget
- await _organization_max_budget_check(
- valid_token=valid_token,
- team_object=team_object,
- prisma_client=prisma_client,
- user_api_key_cache=user_api_key_cache,
- proxy_logging_obj=proxy_logging_obj,
- )
+ # 3.0.5. If team is over soft budget (alert only, doesn't block)
+ await _team_soft_budget_check(
+ team_object=team_object,
+ proxy_logging_obj=proxy_logging_obj,
+ valid_token=valid_token,
+ )
- await _tag_max_budget_check(
- request_body=request_body,
- prisma_client=prisma_client,
- user_api_key_cache=user_api_key_cache,
- proxy_logging_obj=proxy_logging_obj,
- valid_token=valid_token,
- )
+ # 3.1. If organization is in budget
+ await _organization_max_budget_check(
+ valid_token=valid_token,
+ team_object=team_object,
+ prisma_client=prisma_client,
+ user_api_key_cache=user_api_key_cache,
+ proxy_logging_obj=proxy_logging_obj,
+ )
- # 4. If user is in budget
- ## 4.1 check personal budget, if personal key
- if (
- (team_object is None or team_object.team_id is None)
- and user_object is not None
- and user_object.max_budget is not None
- ):
- user_budget = user_object.max_budget
- if user_budget < user_object.spend:
- raise litellm.BudgetExceededError(
- current_cost=user_object.spend,
- max_budget=user_budget,
- message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_object.spend}, Budget={user_budget}",
- )
+ await _tag_max_budget_check(
+ request_body=request_body,
+ prisma_client=prisma_client,
+ user_api_key_cache=user_api_key_cache,
+ proxy_logging_obj=proxy_logging_obj,
+ valid_token=valid_token,
+ )
- ## 4.2 check team member budget, if team key
- await _check_team_member_budget(
- team_object=team_object,
- user_object=user_object,
- valid_token=valid_token,
- prisma_client=prisma_client,
- user_api_key_cache=user_api_key_cache,
- proxy_logging_obj=proxy_logging_obj,
- )
+ # 4. If user is in budget
+ ## 4.1 check personal budget, if personal key
+ if (
+ (team_object is None or team_object.team_id is None)
+ and user_object is not None
+ and user_object.max_budget is not None
+ ):
+ user_budget = user_object.max_budget
+ if user_budget < user_object.spend:
+ raise litellm.BudgetExceededError(
+ current_cost=user_object.spend,
+ max_budget=user_budget,
+ message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_object.spend}, Budget={user_budget}",
+ )
- # 5. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
- if end_user_object is not None and end_user_object.litellm_budget_table is not None:
- end_user_budget = end_user_object.litellm_budget_table.max_budget
- if end_user_budget is not None and end_user_object.spend > end_user_budget:
- raise litellm.BudgetExceededError(
- current_cost=end_user_object.spend,
- max_budget=end_user_budget,
- message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}",
- )
+ ## 4.2 check team member budget, if team key
+ await _check_team_member_budget(
+ team_object=team_object,
+ user_object=user_object,
+ valid_token=valid_token,
+ prisma_client=prisma_client,
+ user_api_key_cache=user_api_key_cache,
+ proxy_logging_obj=proxy_logging_obj,
+ )
+
+ # 5. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
+ if end_user_object is not None and end_user_object.litellm_budget_table is not None:
+ end_user_budget = end_user_object.litellm_budget_table.max_budget
+ if end_user_budget is not None and end_user_object.spend > end_user_budget:
+ raise litellm.BudgetExceededError(
+ current_cost=end_user_object.spend,
+ max_budget=end_user_budget,
+ message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}",
+ )
# 6. [OPTIONAL] If 'enforce_user_param' enabled - did developer pass in 'user' param for openai endpoints
if (
@@ -247,6 +348,7 @@ async def common_checks(
# 7. [OPTIONAL] If 'litellm.max_budget' is set (>0), is proxy under budget
if (
litellm.max_budget > 0
+ and not skip_budget_checks
and global_proxy_spend is not None
# only run global budget checks for OpenAI routes
# Reason - the Admin UI should continue working if the proxy crosses it's global budget
@@ -1135,6 +1237,7 @@ async def get_user_object(
return _response
except Exception as e: # if user not in db
+ _log_budget_lookup_failure("user", e)
raise ValueError(
f"User doesn't exist in db. 'user_id'={user_id}. Create user via `/user/new` call. Got error - {e}"
)
@@ -2348,6 +2451,75 @@ async def _team_max_budget_check(
)
+async def _team_soft_budget_check(
+ team_object: Optional[LiteLLM_TeamTable],
+ valid_token: Optional[UserAPIKeyAuth],
+ proxy_logging_obj: ProxyLogging,
+):
+ """
+ Triggers a budget alert if the team is over it's soft budget.
+ """
+ if (
+ team_object is not None
+ and team_object.soft_budget is not None
+ and team_object.spend is not None
+ and team_object.spend >= team_object.soft_budget
+ ):
+ verbose_proxy_logger.debug(
+ "Crossed Soft Budget for team %s, spend %s, soft_budget %s",
+ team_object.team_id,
+ team_object.spend,
+ team_object.soft_budget,
+ )
+ if valid_token:
+ # Extract alert emails from team metadata
+ alert_emails: Optional[List[str]] = None
+ if team_object.metadata is not None and isinstance(team_object.metadata, dict):
+ soft_budget_alert_emails = team_object.metadata.get("soft_budget_alerting_emails")
+ if soft_budget_alert_emails is not None:
+ if isinstance(soft_budget_alert_emails, list):
+ alert_emails = [email for email in soft_budget_alert_emails if isinstance(email, str) and email.strip()]
+ elif isinstance(soft_budget_alert_emails, str):
+ # Handle comma-separated string
+ alert_emails = [email.strip() for email in soft_budget_alert_emails.split(",") if email.strip()]
+ # Filter out empty strings
+ if alert_emails:
+ alert_emails = [email for email in alert_emails if email]
+ else:
+ alert_emails = None
+
+ # Only send team soft budget alerts if alert_emails are configured
+ # Team soft budget alerts are sent via metadata.soft_budget_alerting_emails, not global alerting
+ if alert_emails is None or len(alert_emails) == 0:
+ verbose_proxy_logger.debug(
+ "Skipping team soft budget alert for team %s: no alert_emails configured in metadata.soft_budget_alerting_emails",
+ team_object.team_id,
+ )
+ return
+
+ call_info = CallInfo(
+ token=valid_token.token,
+ spend=team_object.spend,
+ max_budget=team_object.max_budget,
+ soft_budget=team_object.soft_budget,
+ user_id=valid_token.user_id,
+ team_id=valid_token.team_id,
+ team_alias=valid_token.team_alias,
+ organization_id=valid_token.org_id,
+ user_email=None, # Team-level alert, no specific user email
+ key_alias=valid_token.key_alias,
+ event_group=Litellm_EntityType.TEAM,
+ alert_emails=alert_emails,
+ )
+
+ asyncio.create_task(
+ proxy_logging_obj.budget_alerts(
+ type="soft_budget",
+ user_info=call_info,
+ )
+ )
+
+
async def _organization_max_budget_check(
valid_token: Optional[UserAPIKeyAuth],
team_object: Optional[LiteLLM_TeamTable],
diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py
index 33667b5d8d9..584be0a9496 100644
--- a/litellm/proxy/auth/handle_jwt.py
+++ b/litellm/proxy/auth/handle_jwt.py
@@ -976,6 +976,9 @@ class JWTAuthManager:
user_route=route,
litellm_proxy_roles=jwt_handler.litellm_jwtauth,
)
+ verbose_proxy_logger.debug(
+ f"JWT team route check: team_id={team_id}, route={route}, is_allowed={is_allowed}"
+ )
if is_allowed:
return team_id, team_object
except Exception:
diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py
index 939cfefadcc..4df773dec2b 100644
--- a/litellm/proxy/auth/login_utils.py
+++ b/litellm/proxy/auth/login_utils.py
@@ -34,59 +34,6 @@ from litellm.secret_managers.main import get_secret_bool
from litellm.types.proxy.ui_sso import ReturnedUITokenObject
-async def expire_previous_ui_session_tokens(
- user_id: str, prisma_client: Optional[PrismaClient]
-) -> None:
- """
- Expire (block) all other valid UI session tokens for a user.
-
- This prevents accumulation of multiple valid UI session tokens that
- are supposed to be short-lived test keys. Only affects keys with
- team_id = "litellm-dashboard" and that haven't expired yet.
-
- Args:
- user_id: The user ID whose previous UI session tokens should be expired
- prisma_client: Database client for performing the update
- """
- if prisma_client is None:
- return
-
- try:
- from datetime import datetime, timezone
-
- current_time = datetime.now(timezone.utc)
-
- # Find all unblocked AND non-expired UI session tokens for this user
- ui_session_tokens = await prisma_client.db.litellm_verificationtoken.find_many(
- where={
- "user_id": user_id,
- "team_id": "litellm-dashboard",
- "OR": [
- {"blocked": None}, # Tokens that have never been blocked (null)
- {"blocked": False}, # Tokens explicitly set to not blocked
- ],
- "expires": {"gt": current_time}, # Only get tokens that haven't expired
- }
- )
-
- if not ui_session_tokens:
- return
-
- # Block all the found tokens
- tokens_to_block = [token.token for token in ui_session_tokens if token.token]
-
- if tokens_to_block:
- await prisma_client.db.litellm_verificationtoken.update_many(
- where={"token": {"in": tokens_to_block}},
- data={"blocked": True}
- )
-
- except Exception:
- # Silently fail - don't block login if cleanup fails
- # This is a best-effort operation
- pass
-
-
def get_ui_credentials(master_key: Optional[str]) -> tuple[str, str]:
"""
Get UI username and password from environment variables or master key.
@@ -227,10 +174,6 @@ async def authenticate_user( # noqa: PLR0915
)
if os.getenv("DATABASE_URL") is not None:
- # Expire any previous UI session tokens for this user
- await expire_previous_ui_session_tokens(
- user_id=key_user_id, prisma_client=prisma_client
- )
response = await generate_key_helper_fn(
request_type="key",
**{
@@ -317,11 +260,6 @@ async def authenticate_user( # noqa: PLR0915
password.encode("utf-8"), _password.encode("utf-8")
) or secrets.compare_digest(hash_password.encode("utf-8"), _password.encode("utf-8")):
if os.getenv("DATABASE_URL") is not None:
- # Expire any previous UI session tokens for this user
- await expire_previous_ui_session_tokens(
- user_id=user_id, prisma_client=prisma_client
- )
-
response = await generate_key_helper_fn(
request_type="key",
**{ # type: ignore
diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py
index 7290528cb5a..05eeab3f611 100644
--- a/litellm/proxy/auth/user_api_key_auth.py
+++ b/litellm/proxy/auth/user_api_key_auth.py
@@ -604,6 +604,21 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
if team_object is not None
else None,
)
+
+ # Check if model has zero cost - if so, skip all budget checks
+ model = get_model_from_request(request_data, route)
+ skip_budget_checks = False
+ if model is not None and llm_router is not None:
+ from litellm.proxy.auth.auth_checks import _is_model_cost_zero
+
+ skip_budget_checks = _is_model_cost_zero(
+ model=model, llm_router=llm_router
+ )
+ if skip_budget_checks:
+ verbose_proxy_logger.info(
+ f"Skipping all budget checks for zero-cost model: {model}"
+ )
+
# run through common checks
_ = await common_checks(
request=request,
@@ -617,6 +632,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
+ skip_budget_checks=skip_budget_checks,
)
# return UserAPIKeyAuth object
@@ -1008,8 +1024,22 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
)
user_obj = None
+ # Check 2a. Check if model has zero cost - if so, skip all budget checks
+ model = get_model_from_request(request_data, route)
+ skip_budget_checks = False
+ if model is not None and llm_router is not None:
+ from litellm.proxy.auth.auth_checks import _is_model_cost_zero
+
+ skip_budget_checks = _is_model_cost_zero(
+ model=model, llm_router=llm_router
+ )
+ if skip_budget_checks:
+ verbose_proxy_logger.info(
+ f"Skipping all budget checks for zero-cost model: {model}"
+ )
+
# Check 3. Check if user is in their team budget
- if valid_token.team_member_spend is not None:
+ if not skip_budget_checks and valid_token.team_member_spend is not None:
if prisma_client is not None:
_cache_key = f"{valid_token.team_id}_{valid_token.user_id}"
@@ -1073,51 +1103,53 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
param=abbreviate_api_key(api_key=api_key),
)
- # Check 4. Token Spend is under budget
- if RouteChecks.is_llm_api_route(route=route):
- await _virtual_key_max_budget_check(
+ if not skip_budget_checks:
+ # Check 4. Token Spend is under budget
+ if RouteChecks.is_llm_api_route(route=route):
+ await _virtual_key_max_budget_check(
+ valid_token=valid_token,
+ proxy_logging_obj=proxy_logging_obj,
+ user_obj=user_obj,
+ )
+
+ # Check 5. Max Budget Alert Check
+ await _virtual_key_max_budget_alert_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
- # Check 5. Max Budget Alert Check
- await _virtual_key_max_budget_alert_check(
- valid_token=valid_token,
- proxy_logging_obj=proxy_logging_obj,
- user_obj=user_obj,
- )
-
- # Check 6. Soft Budget Check
- await _virtual_key_soft_budget_check(
- valid_token=valid_token,
- proxy_logging_obj=proxy_logging_obj,
- user_obj=user_obj,
- )
-
- # Check 5. Token Model Spend is under Model budget
- max_budget_per_model = valid_token.model_max_budget
- current_model = request_data.get("model", None)
-
- if (
- max_budget_per_model is not None
- and isinstance(max_budget_per_model, dict)
- and len(max_budget_per_model) > 0
- and prisma_client is not None
- and current_model is not None
- and valid_token.token is not None
- ):
- ## GET THE SPEND FOR THIS MODEL
- await model_max_budget_limiter.is_key_within_model_budget(
- user_api_key_dict=valid_token,
- model=current_model,
+ # Check 6. Soft Budget Check
+ await _virtual_key_soft_budget_check(
+ valid_token=valid_token,
+ proxy_logging_obj=proxy_logging_obj,
+ user_obj=user_obj,
)
+ # Check 5. Token Model Spend is under Model budget
+ max_budget_per_model = valid_token.model_max_budget
+ current_model = request_data.get("model", None)
+
+ if (
+ max_budget_per_model is not None
+ and isinstance(max_budget_per_model, dict)
+ and len(max_budget_per_model) > 0
+ and prisma_client is not None
+ and current_model is not None
+ and valid_token.token is not None
+ ):
+ ## GET THE SPEND FOR THIS MODEL
+ await model_max_budget_limiter.is_key_within_model_budget(
+ user_api_key_dict=valid_token,
+ model=current_model,
+ )
+
# Check 6: Additional Common Checks across jwt + key auth
if valid_token.team_id is not None:
_team_obj: Optional[LiteLLM_TeamTable] = LiteLLM_TeamTable(
team_id=valid_token.team_id,
max_budget=valid_token.team_max_budget,
+ soft_budget=valid_token.team_soft_budget,
spend=valid_token.team_spend,
tpm_limit=valid_token.team_tpm_limit,
rpm_limit=valid_token.team_rpm_limit,
@@ -1171,6 +1203,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
+ skip_budget_checks=skip_budget_checks,
)
# Token passed all checks
if valid_token is None:
diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py
index f47e2e1667b..06800cb4524 100644
--- a/litellm/proxy/batches_endpoints/endpoints.py
+++ b/litellm/proxy/batches_endpoints/endpoints.py
@@ -24,10 +24,12 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
decode_model_from_file_id,
encode_file_id_with_model,
+ get_batch_from_database,
get_credentials_for_model,
get_models_from_unified_file_id,
get_original_file_id,
prepare_data_with_credentials,
+ update_batch_in_database,
)
from litellm.proxy.utils import handle_exception_on_proxy, is_known_model
from litellm.types.llms.openai import LiteLLMBatchCreateRequest
@@ -357,6 +359,57 @@ async def retrieve_batch(
route_type="aretrieve_batch",
)
+ # FIX: First, try to read from ManagedObjectTable for consistent state
+ managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files")
+ from litellm.proxy.proxy_server import prisma_client
+
+ db_batch_object, response = await get_batch_from_database(
+ batch_id=batch_id,
+ unified_batch_id=unified_batch_id,
+ managed_files_obj=managed_files_obj,
+ prisma_client=prisma_client,
+ verbose_proxy_logger=verbose_proxy_logger,
+ )
+
+ # If batch is in a terminal state, return immediately
+ if response is not None and response.status in ["completed", "failed", "cancelled", "expired"]:
+ # Call hooks and return
+ response = await proxy_logging_obj.post_call_success_hook(
+ data=data, user_api_key_dict=user_api_key_dict, response=response
+ )
+
+ asyncio.create_task(
+ proxy_logging_obj.update_request_status(
+ litellm_call_id=data.get("litellm_call_id", ""), status="success"
+ )
+ )
+
+ hidden_params = getattr(response, "_hidden_params", {}) or {}
+ model_id = hidden_params.get("model_id", None) or ""
+ cache_key = hidden_params.get("cache_key", None) or ""
+ api_base = hidden_params.get("api_base", None) or ""
+
+ fastapi_response.headers.update(
+ ProxyBaseLLMRequestProcessing.get_custom_headers(
+ user_api_key_dict=user_api_key_dict,
+ model_id=model_id,
+ cache_key=cache_key,
+ api_base=api_base,
+ version=version,
+ model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
+ request_data=data,
+ )
+ )
+
+ return response
+
+ # If batch is still processing, sync with provider to get latest state
+ if response is not None:
+ verbose_proxy_logger.debug(
+ f"Batch {batch_id} is in non-terminal state {response.status}, syncing with provider"
+ )
+
+ # Retrieve from provider (for non-terminal states or if DB lookup failed)
# SCENARIO 1: Batch ID is encoded with model info
if model_from_id is not None:
credentials = get_credentials_for_model(
@@ -408,6 +461,18 @@ async def retrieve_batch(
response = await litellm.aretrieve_batch(
custom_llm_provider=custom_llm_provider, **data # type: ignore
)
+
+ # FIX: Update the database with the latest state from provider
+ await update_batch_in_database(
+ batch_id=batch_id,
+ unified_batch_id=unified_batch_id,
+ response=response,
+ managed_files_obj=managed_files_obj,
+ prisma_client=prisma_client,
+ verbose_proxy_logger=verbose_proxy_logger,
+ db_batch_object=db_batch_object,
+ operation="retrieve",
+ )
### CALL HOOKS ### - modify outgoing data
response = await proxy_logging_obj.post_call_success_hook(
@@ -769,6 +834,20 @@ async def cancel_batch(
**_cancel_batch_data,
)
+ # FIX: Update the database with the new cancelled state
+ managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files")
+ from litellm.proxy.proxy_server import prisma_client
+
+ await update_batch_in_database(
+ batch_id=batch_id,
+ unified_batch_id=unified_batch_id,
+ response=response,
+ managed_files_obj=managed_files_obj,
+ prisma_client=prisma_client,
+ verbose_proxy_logger=verbose_proxy_logger,
+ operation="cancel",
+ )
+
### CALL HOOKS ### - modify outgoing data
response = await proxy_logging_obj.post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py
index 429e56c805b..dc928921425 100644
--- a/litellm/proxy/db/db_spend_update_writer.py
+++ b/litellm/proxy/db/db_spend_update_writer.py
@@ -1187,119 +1187,130 @@ class DBSpendUpdateWriter:
)
break
- async with prisma_client.db.batch_() as batcher:
- for _, transaction in transactions_to_process.items():
- entity_id = transaction.get(entity_id_field)
+ try:
+ async with prisma_client.db.batch_() as batcher:
+ for _, transaction in transactions_to_process.items():
+ entity_id = transaction.get(entity_id_field)
- # Construct the where clause dynamically
- where_clause = {
- unique_constraint_name: {
+ # Construct the where clause dynamically
+ where_clause = {
+ unique_constraint_name: {
+ entity_id_field: entity_id,
+ "date": transaction["date"],
+ "api_key": transaction["api_key"],
+ "model": transaction["model"],
+ "custom_llm_provider": transaction.get(
+ "custom_llm_provider"
+ )
+ or "",
+ "mcp_namespaced_tool_name": transaction.get(
+ "mcp_namespaced_tool_name"
+ )
+ or "",
+ "endpoint": transaction.get("endpoint") or "",
+ }
+ }
+
+ # Get the table dynamically
+ table = getattr(batcher, table_name)
+
+ # Common data structure for both create and update
+ common_data = {
entity_id_field: entity_id,
"date": transaction["date"],
"api_key": transaction["api_key"],
- "model": transaction["model"],
- "custom_llm_provider": transaction.get(
- "custom_llm_provider"
- )
- or "",
+ "model": transaction.get("model"),
+ "model_group": transaction.get("model_group"),
"mcp_namespaced_tool_name": transaction.get(
"mcp_namespaced_tool_name"
)
or "",
+ "custom_llm_provider": transaction.get(
+ "custom_llm_provider"
+ ),
"endpoint": transaction.get("endpoint") or "",
+ "prompt_tokens": transaction["prompt_tokens"],
+ "completion_tokens": transaction["completion_tokens"],
+ "spend": transaction["spend"],
+ "api_requests": transaction["api_requests"],
+ "successful_requests": transaction[
+ "successful_requests"
+ ],
+ "failed_requests": transaction["failed_requests"],
}
- }
- # Get the table dynamically
- table = getattr(batcher, table_name)
-
- # Common data structure for both create and update
- common_data = {
- entity_id_field: entity_id,
- "date": transaction["date"],
- "api_key": transaction["api_key"],
- "model": transaction.get("model"),
- "model_group": transaction.get("model_group"),
- "mcp_namespaced_tool_name": transaction.get(
- "mcp_namespaced_tool_name"
- )
- or "",
- "custom_llm_provider": transaction.get(
- "custom_llm_provider"
- ),
- "endpoint": transaction.get("endpoint"),
- "prompt_tokens": transaction["prompt_tokens"],
- "completion_tokens": transaction["completion_tokens"],
- "spend": transaction["spend"],
- "api_requests": transaction["api_requests"],
- "successful_requests": transaction[
- "successful_requests"
- ],
- "failed_requests": transaction["failed_requests"],
- }
-
- # Add cache-related fields if they exist
- if "cache_read_input_tokens" in transaction:
- common_data["cache_read_input_tokens"] = (
- transaction.get("cache_read_input_tokens", 0)
- )
- if "cache_creation_input_tokens" in transaction:
- common_data["cache_creation_input_tokens"] = (
- transaction.get("cache_creation_input_tokens", 0)
- )
-
- if entity_type == "tag" and "request_id" in transaction:
- common_data["request_id"] = transaction.get(
- "request_id"
- )
-
- # Create update data structure
- update_data = {
- "prompt_tokens": {
- "increment": transaction["prompt_tokens"]
- },
- "completion_tokens": {
- "increment": transaction["completion_tokens"]
- },
- "spend": {"increment": transaction["spend"]},
- "api_requests": {
- "increment": transaction["api_requests"]
- },
- "successful_requests": {
- "increment": transaction["successful_requests"]
- },
- "failed_requests": {
- "increment": transaction["failed_requests"]
- },
- }
-
- # Add cache-related fields to update if they exist
- if "cache_read_input_tokens" in transaction:
- update_data["cache_read_input_tokens"] = {
- "increment": transaction.get(
- "cache_read_input_tokens", 0
+ # Add cache-related fields if they exist
+ if "cache_read_input_tokens" in transaction:
+ common_data["cache_read_input_tokens"] = (
+ transaction.get("cache_read_input_tokens", 0)
)
- }
- if "cache_creation_input_tokens" in transaction:
- update_data["cache_creation_input_tokens"] = {
- "increment": transaction.get(
- "cache_creation_input_tokens", 0
+ if "cache_creation_input_tokens" in transaction:
+ common_data["cache_creation_input_tokens"] = (
+ transaction.get("cache_creation_input_tokens", 0)
)
+
+ if entity_type == "tag" and "request_id" in transaction:
+ common_data["request_id"] = transaction.get(
+ "request_id"
+ )
+
+ # Create update data structure
+ update_data = {
+ "prompt_tokens": {
+ "increment": transaction["prompt_tokens"]
+ },
+ "completion_tokens": {
+ "increment": transaction["completion_tokens"]
+ },
+ "spend": {"increment": transaction["spend"]},
+ "api_requests": {
+ "increment": transaction["api_requests"]
+ },
+ "successful_requests": {
+ "increment": transaction["successful_requests"]
+ },
+ "failed_requests": {
+ "increment": transaction["failed_requests"]
+ },
}
- if entity_type == "tag" and "request_id" in transaction:
- update_data["request_id"] = transaction.get("request_id")
+ # Add cache-related fields to update if they exist
+ if "cache_read_input_tokens" in transaction:
+ update_data["cache_read_input_tokens"] = {
+ "increment": transaction.get(
+ "cache_read_input_tokens", 0
+ )
+ }
+ if "cache_creation_input_tokens" in transaction:
+ update_data["cache_creation_input_tokens"] = {
+ "increment": transaction.get(
+ "cache_creation_input_tokens", 0
+ )
+ }
- # Add endpoint to update_data so existing rows get their endpoint field updated
- update_data["endpoint"] = transaction.get("endpoint") or ""
+ if entity_type == "tag" and "request_id" in transaction:
+ update_data["request_id"] = transaction.get("request_id")
- table.upsert(
- where=where_clause,
- data={
- "create": common_data,
- "update": update_data,
- },
- )
+ # Add endpoint to update_data so existing rows get their endpoint field updated
+ update_data["endpoint"] = transaction.get("endpoint") or ""
+
+ table.upsert(
+ where=where_clause,
+ data={
+ "create": common_data,
+ "update": update_data,
+ },
+ )
+ except Exception as batch_error:
+ # Log detailed error information for debugging batch upsert failures
+ # This helps diagnose issues like unique constraint violations
+ verbose_proxy_logger.exception(
+ f"Daily {entity_type} spend batch upsert failed. "
+ f"Table: {table_name}, Constraint: {unique_constraint_name}, "
+ f"Batch size: {len(transactions_to_process)}, "
+ f"Error: {str(batch_error)}"
+ )
+ raise
verbose_proxy_logger.debug(
f"Processed {len(transactions_to_process)} daily {entity_type} transactions in {time.time() - start_time:.2f}s"
diff --git a/litellm/proxy/example_config_yaml/otel_test_config.yaml b/litellm/proxy/example_config_yaml/otel_test_config.yaml
index 714875d56ce..7ddb5d40c0c 100644
--- a/litellm/proxy/example_config_yaml/otel_test_config.yaml
+++ b/litellm/proxy/example_config_yaml/otel_test_config.yaml
@@ -1,7 +1,7 @@
model_list:
- model_name: fake-openai-endpoint
litellm_params:
- model: openai/fake
+ model: openai/gpt-3.5-turbo-0301
api_key: fake-key
api_base: https://exampleopenaiendpoint-production.up.railway.app/
tags: ["teamA"]
@@ -9,7 +9,7 @@ model_list:
id: "team-a-model"
- model_name: fake-openai-endpoint
litellm_params:
- model: openai/fake
+ model: openai/gpt-3.5-turbo-0301
api_key: fake-key
api_base: https://exampleopenaiendpoint-production.up.railway.app/
tags: ["teamB"]
diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py
index 3ce819439cb..a825ce22b25 100644
--- a/litellm/proxy/guardrails/guardrail_endpoints.py
+++ b/litellm/proxy/guardrails/guardrail_endpoints.py
@@ -1236,6 +1236,275 @@ async def get_provider_specific_params():
return provider_params
+class TestCustomCodeGuardrailRequest(BaseModel):
+ """Request model for testing custom code guardrails."""
+
+ custom_code: str
+ """The Python-like code containing the apply_guardrail function."""
+
+ test_input: Dict[str, Any]
+ """The test input to pass to the guardrail. Should contain 'texts', optionally 'images', 'tools', etc."""
+
+ input_type: str = "request"
+ """Whether this is a 'request' or 'response' input type."""
+
+ request_data: Optional[Dict[str, Any]] = None
+ """Optional mock request_data (model, user_id, team_id, metadata, etc.)."""
+
+
+class TestCustomCodeGuardrailResponse(BaseModel):
+ """Response model for testing custom code guardrails."""
+
+ success: bool
+ """Whether the test executed successfully (no errors)."""
+
+ result: Optional[Dict[str, Any]] = None
+ """The guardrail result: action (allow/block/modify), reason, modified_texts, etc."""
+
+ error: Optional[str] = None
+ """Error message if execution failed."""
+
+ error_type: Optional[str] = None
+ """Type of error: 'compilation' or 'execution'."""
+
+
+@router.post(
+ "/guardrails/test_custom_code",
+ tags=["Guardrails"],
+ dependencies=[Depends(user_api_key_auth)],
+ response_model=TestCustomCodeGuardrailResponse,
+)
+async def test_custom_code_guardrail(request: TestCustomCodeGuardrailRequest):
+ """
+ Test custom code guardrail logic without creating a guardrail.
+
+ This endpoint allows admins to experiment with custom code guardrails by:
+ 1. Compiling the provided code in a sandbox
+ 2. Executing the apply_guardrail function with test input
+ 3. Returning the result (allow/block/modify)
+
+ 👉 [Custom Code Guardrail docs](https://docs.litellm.ai/docs/proxy/guardrails/custom_code_guardrail)
+
+ Example Request:
+ ```bash
+ curl -X POST "http://localhost:4000/guardrails/test_custom_code" \\
+ -H "Authorization: Bearer " \\
+ -H "Content-Type: application/json" \\
+ -d '{
+ "custom_code": "def apply_guardrail(inputs, request_data, input_type):\\n for text in inputs[\\"texts\\"]:\\n if regex_match(text, r\\"\\\\d{3}-\\\\d{2}-\\\\d{4}\\"):\\n return block(\\"SSN detected\\")\\n return allow()",
+ "test_input": {
+ "texts": ["My SSN is 123-45-6789"]
+ },
+ "input_type": "request"
+ }'
+ ```
+
+ Example Success Response (blocked):
+ ```json
+ {
+ "success": true,
+ "result": {
+ "action": "block",
+ "reason": "SSN detected"
+ },
+ "error": null,
+ "error_type": null
+ }
+ ```
+
+ Example Success Response (allowed):
+ ```json
+ {
+ "success": true,
+ "result": {
+ "action": "allow"
+ },
+ "error": null,
+ "error_type": null
+ }
+ ```
+
+ Example Success Response (modified):
+ ```json
+ {
+ "success": true,
+ "result": {
+ "action": "modify",
+ "texts": ["My SSN is [REDACTED]"]
+ },
+ "error": null,
+ "error_type": null
+ }
+ ```
+
+ Example Error Response (compilation error):
+ ```json
+ {
+ "success": false,
+ "result": null,
+ "error": "Syntax error in custom code: invalid syntax (, line 1)",
+ "error_type": "compilation"
+ }
+ ```
+ """
+ import concurrent.futures
+ import re
+
+ from litellm.proxy.guardrails.guardrail_hooks.custom_code.primitives import (
+ get_custom_code_primitives,
+ )
+
+ # Security validation patterns
+ FORBIDDEN_PATTERNS = [
+ # Import statements
+ (r"\bimport\s+", "import statements are not allowed"),
+ (r"\bfrom\s+\w+\s+import\b", "from...import statements are not allowed"),
+ (r"__import__\s*\(", "__import__() is not allowed"),
+ # Dangerous builtins
+ (r"\bexec\s*\(", "exec() is not allowed"),
+ (r"\beval\s*\(", "eval() is not allowed"),
+ (r"\bcompile\s*\(", "compile() is not allowed"),
+ (r"\bopen\s*\(", "open() is not allowed"),
+ (r"\bgetattr\s*\(", "getattr() is not allowed"),
+ (r"\bsetattr\s*\(", "setattr() is not allowed"),
+ (r"\bdelattr\s*\(", "delattr() is not allowed"),
+ (r"\bglobals\s*\(", "globals() is not allowed"),
+ (r"\blocals\s*\(", "locals() is not allowed"),
+ (r"\bvars\s*\(", "vars() is not allowed"),
+ (r"\bdir\s*\(", "dir() is not allowed"),
+ (r"\bbreakpoint\s*\(", "breakpoint() is not allowed"),
+ (r"\binput\s*\(", "input() is not allowed"),
+ # Dangerous dunder access
+ (r"__builtins__", "__builtins__ access is not allowed"),
+ (r"__globals__", "__globals__ access is not allowed"),
+ (r"__code__", "__code__ access is not allowed"),
+ (r"__subclasses__", "__subclasses__ access is not allowed"),
+ (r"__bases__", "__bases__ access is not allowed"),
+ (r"__mro__", "__mro__ access is not allowed"),
+ (r"__class__", "__class__ access is not allowed"),
+ (r"__dict__", "__dict__ access is not allowed"),
+ (r"__getattribute__", "__getattribute__ access is not allowed"),
+ (r"__reduce__", "__reduce__ access is not allowed"),
+ (r"__reduce_ex__", "__reduce_ex__ access is not allowed"),
+ # OS/system access
+ (r"\bos\.", "os module access is not allowed"),
+ (r"\bsys\.", "sys module access is not allowed"),
+ (r"\bsubprocess\.", "subprocess module access is not allowed"),
+ ]
+
+ EXECUTION_TIMEOUT_SECONDS = 5
+
+ try:
+ # Step 0: Security validation - check for forbidden patterns
+ code = request.custom_code
+ for pattern, error_msg in FORBIDDEN_PATTERNS:
+ if re.search(pattern, code):
+ return TestCustomCodeGuardrailResponse(
+ success=False,
+ error=f"Security violation: {error_msg}",
+ error_type="compilation",
+ )
+
+ # Step 1: Compile the custom code with restricted environment
+ exec_globals = get_custom_code_primitives().copy()
+
+ # Remove access to builtins to prevent escape
+ exec_globals["__builtins__"] = {}
+
+ try:
+ exec(compile(request.custom_code, "", "exec"), exec_globals)
+ except SyntaxError as e:
+ return TestCustomCodeGuardrailResponse(
+ success=False,
+ error=f"Syntax error in custom code: {e}",
+ error_type="compilation",
+ )
+ except Exception as e:
+ return TestCustomCodeGuardrailResponse(
+ success=False,
+ error=f"Failed to compile custom code: {e}",
+ error_type="compilation",
+ )
+
+ # Step 2: Verify apply_guardrail function exists
+ if "apply_guardrail" not in exec_globals:
+ return TestCustomCodeGuardrailResponse(
+ success=False,
+ error="Custom code must define an 'apply_guardrail' function. "
+ "Expected signature: apply_guardrail(inputs, request_data, input_type)",
+ error_type="compilation",
+ )
+
+ apply_fn = exec_globals["apply_guardrail"]
+ if not callable(apply_fn):
+ return TestCustomCodeGuardrailResponse(
+ success=False,
+ error="'apply_guardrail' must be a callable function",
+ error_type="compilation",
+ )
+
+ # Step 3: Prepare test inputs
+ test_inputs = request.test_input
+ if "texts" not in test_inputs:
+ test_inputs["texts"] = []
+
+ # Prepare mock request_data
+ mock_request_data = request.request_data or {}
+ safe_request_data = {
+ "model": mock_request_data.get("model", "test-model"),
+ "user_id": mock_request_data.get("user_id"),
+ "team_id": mock_request_data.get("team_id"),
+ "end_user_id": mock_request_data.get("end_user_id"),
+ "metadata": mock_request_data.get("metadata", {}),
+ }
+
+ # Step 4: Execute the function with timeout protection
+
+ def execute_guardrail():
+ return apply_fn(test_inputs, safe_request_data, request.input_type)
+
+ try:
+ with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor:
+ future = executor.submit(execute_guardrail)
+ try:
+ result = future.result(timeout=EXECUTION_TIMEOUT_SECONDS)
+ except concurrent.futures.TimeoutError:
+ return TestCustomCodeGuardrailResponse(
+ success=False,
+ error=f"Execution timeout: code took longer than {EXECUTION_TIMEOUT_SECONDS} seconds",
+ error_type="execution",
+ )
+ except Exception as e:
+ return TestCustomCodeGuardrailResponse(
+ success=False,
+ error=f"Execution error: {e}",
+ error_type="execution",
+ )
+
+ # Step 5: Validate and return result
+ if not isinstance(result, dict):
+ return TestCustomCodeGuardrailResponse(
+ success=True,
+ result={
+ "action": "allow",
+ "warning": f"Expected dict result, got {type(result).__name__}. Treating as allow.",
+ },
+ )
+
+ return TestCustomCodeGuardrailResponse(
+ success=True,
+ result=result,
+ )
+
+ except Exception as e:
+ verbose_proxy_logger.exception(f"Error testing custom code guardrail: {e}")
+ return TestCustomCodeGuardrailResponse(
+ success=False,
+ error=f"Unexpected error: {e}",
+ error_type="execution",
+ )
+
+
@router.post("/guardrails/apply_guardrail", response_model=ApplyGuardrailResponse)
@router.post("/apply_guardrail", response_model=ApplyGuardrailResponse)
async def apply_guardrail(
diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/__init__.py
new file mode 100644
index 00000000000..747b188feea
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/__init__.py
@@ -0,0 +1,65 @@
+"""Custom code guardrail integration for LiteLLM.
+
+This module allows users to write custom guardrail logic using Python-like code
+that runs in a sandboxed environment with access to LiteLLM-provided primitives.
+"""
+
+from typing import TYPE_CHECKING
+
+from litellm.types.guardrails import SupportedGuardrailIntegrations
+
+from .custom_code_guardrail import CustomCodeGuardrail
+
+if TYPE_CHECKING:
+ from litellm.types.guardrails import Guardrail, LitellmParams
+
+
+def initialize_guardrail(
+ litellm_params: "LitellmParams", guardrail: "Guardrail"
+) -> CustomCodeGuardrail:
+ """
+ Initialize a custom code guardrail.
+
+ Args:
+ litellm_params: Configuration parameters including the custom code
+ guardrail: The guardrail configuration dict
+
+ Returns:
+ CustomCodeGuardrail instance
+ """
+ import litellm
+
+ guardrail_name = guardrail.get("guardrail_name")
+ if not guardrail_name:
+ raise ValueError("Custom code guardrail requires a guardrail_name")
+
+ # Get the custom code from litellm_params
+ custom_code = getattr(litellm_params, "custom_code", None)
+ if not custom_code:
+ raise ValueError(
+ "Custom code guardrail requires 'custom_code' in litellm_params"
+ )
+
+ custom_code_guardrail = CustomCodeGuardrail(
+ guardrail_name=guardrail_name,
+ custom_code=custom_code,
+ event_hook=litellm_params.mode,
+ default_on=litellm_params.default_on,
+ )
+
+ litellm.logging_callback_manager.add_litellm_callback(custom_code_guardrail)
+ return custom_code_guardrail
+
+
+guardrail_initializer_registry = {
+ SupportedGuardrailIntegrations.CUSTOM_CODE.value: initialize_guardrail,
+}
+
+guardrail_class_registry = {
+ SupportedGuardrailIntegrations.CUSTOM_CODE.value: CustomCodeGuardrail,
+}
+
+__all__ = [
+ "CustomCodeGuardrail",
+ "initialize_guardrail",
+]
diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py
new file mode 100644
index 00000000000..a0ca324411c
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py
@@ -0,0 +1,372 @@
+"""
+Custom code guardrail for LiteLLM.
+
+This module provides a guardrail that executes user-defined Python-like code
+to implement custom guardrail logic. The code runs in a sandboxed environment
+with access to LiteLLM-provided primitives for common guardrail operations.
+
+Example custom code:
+
+ def apply_guardrail(inputs, request_data, input_type):
+ '''Block messages containing SSNs'''
+ for text in inputs["texts"]:
+ if regex_match(text, r"\\d{3}-\\d{2}-\\d{4}"):
+ return block("Social Security Number detected")
+ return allow()
+"""
+
+import threading
+from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Type, cast
+
+from fastapi import HTTPException
+
+from litellm._logging import verbose_proxy_logger
+from litellm.integrations.custom_guardrail import CustomGuardrail
+from litellm.types.guardrails import GuardrailEventHooks
+from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
+from litellm.types.utils import GenericGuardrailAPIInputs
+
+from .primitives import get_custom_code_primitives
+
+if TYPE_CHECKING:
+ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
+
+
+class CustomCodeGuardrailError(Exception):
+ """Raised when custom code guardrail execution fails."""
+
+ def __init__(self, message: str, details: Optional[Dict[str, Any]] = None) -> None:
+ super().__init__(message)
+ self.details = details or {}
+
+
+class CustomCodeCompilationError(CustomCodeGuardrailError):
+ """Raised when custom code fails to compile."""
+
+
+class CustomCodeExecutionError(CustomCodeGuardrailError):
+ """Raised when custom code fails during execution."""
+
+
+class CustomCodeGuardrailConfigModel(GuardrailConfigModel):
+ """Configuration parameters for the custom code guardrail."""
+
+ custom_code: str
+ """The Python-like code containing the apply_guardrail function."""
+
+
+class CustomCodeGuardrail(CustomGuardrail):
+ """
+ Guardrail that executes user-defined Python-like code.
+
+ The code runs in a sandboxed environment that provides:
+ - Access to LiteLLM primitives (regex_match, json_parse, etc.)
+ - No file I/O or network access
+ - No imports allowed
+
+ Users write an `apply_guardrail(inputs, request_data, input_type)` function
+ that returns one of:
+ - allow() - let the request/response through
+ - block(reason) - reject with a message
+ - modify(texts=...) - transform the content
+
+ Example:
+ def apply_guardrail(inputs, request_data, input_type):
+ for text in inputs["texts"]:
+ if regex_match(text, r"password"):
+ return block("Sensitive content detected")
+ return allow()
+ """
+
+ def __init__(
+ self,
+ custom_code: str,
+ guardrail_name: Optional[str] = "custom_code",
+ **kwargs: Any,
+ ) -> None:
+ """
+ Initialize the custom code guardrail.
+
+ Args:
+ custom_code: The source code containing apply_guardrail function
+ guardrail_name: Name of this guardrail instance
+ **kwargs: Additional arguments passed to CustomGuardrail
+ """
+ self.custom_code = custom_code
+ self._compiled_function: Optional[Any] = None
+ self._compile_lock = threading.Lock()
+ self._compile_error: Optional[str] = None
+
+ supported_event_hooks = [
+ GuardrailEventHooks.pre_call,
+ GuardrailEventHooks.during_call,
+ GuardrailEventHooks.post_call,
+ ]
+
+ super().__init__(
+ guardrail_name=guardrail_name,
+ supported_event_hooks=supported_event_hooks,
+ **kwargs,
+ )
+
+ # Compile the code on initialization
+ self._compile_custom_code()
+
+ @staticmethod
+ def get_config_model() -> Optional[Type[GuardrailConfigModel]]:
+ """Returns the config model for the UI."""
+ return CustomCodeGuardrailConfigModel
+
+ def _compile_custom_code(self) -> None:
+ """
+ Compile the custom code and extract the apply_guardrail function.
+
+ The code runs in a sandboxed environment with only the allowed primitives.
+ """
+ with self._compile_lock:
+ if self._compiled_function is not None:
+ return
+
+ try:
+ # Create a restricted execution environment
+ # Only include our safe primitives
+ exec_globals = get_custom_code_primitives().copy()
+
+ # Execute the user code in the restricted environment
+ exec(compile(self.custom_code, "", "exec"), exec_globals)
+
+ # Extract the apply_guardrail function
+ if "apply_guardrail" not in exec_globals:
+ raise CustomCodeCompilationError(
+ "Custom code must define an 'apply_guardrail' function. "
+ "Expected signature: apply_guardrail(inputs, request_data, input_type)"
+ )
+
+ apply_fn = exec_globals["apply_guardrail"]
+ if not callable(apply_fn):
+ raise CustomCodeCompilationError(
+ "'apply_guardrail' must be a callable function"
+ )
+
+ self._compiled_function = apply_fn
+ verbose_proxy_logger.debug(
+ f"Custom code guardrail '{self.guardrail_name}' compiled successfully"
+ )
+
+ except SyntaxError as e:
+ self._compile_error = f"Syntax error in custom code: {e}"
+ raise CustomCodeCompilationError(self._compile_error) from e
+ except CustomCodeCompilationError:
+ raise
+ except Exception as e:
+ self._compile_error = f"Failed to compile custom code: {e}"
+ raise CustomCodeCompilationError(self._compile_error) from e
+
+ async def apply_guardrail(
+ self,
+ inputs: GenericGuardrailAPIInputs,
+ request_data: dict,
+ input_type: Literal["request", "response"],
+ logging_obj: Optional["LiteLLMLoggingObj"] = None,
+ ) -> GenericGuardrailAPIInputs:
+ """
+ Apply the custom code guardrail to the inputs.
+
+ This method calls the user-defined apply_guardrail function and
+ processes its result to determine the appropriate action.
+
+ Args:
+ inputs: Dictionary containing texts, images, tool_calls
+ request_data: The original request data with metadata
+ input_type: "request" for pre-call, "response" for post-call
+ logging_obj: Optional logging object
+
+ Returns:
+ GenericGuardrailAPIInputs - possibly modified
+
+ Raises:
+ HTTPException: If content is blocked
+ CustomCodeExecutionError: If execution fails
+ """
+ if self._compiled_function is None:
+ if self._compile_error:
+ raise CustomCodeExecutionError(
+ f"Custom code guardrail not compiled: {self._compile_error}"
+ )
+ raise CustomCodeExecutionError("Custom code guardrail not compiled")
+
+ try:
+ # Prepare inputs dict for the function
+
+ # Prepare request_data with safe subset of information
+ safe_request_data = self._prepare_safe_request_data(request_data)
+
+ # Execute the custom function
+ result = self._compiled_function(inputs, safe_request_data, input_type)
+
+ # Process the result
+ return self._process_result(
+ result=result,
+ inputs=inputs,
+ request_data=request_data,
+ input_type=input_type,
+ )
+
+ except HTTPException:
+ # Re-raise HTTP exceptions (from block action)
+ raise
+ except Exception as e:
+ verbose_proxy_logger.error(
+ f"Custom code guardrail '{self.guardrail_name}' execution error: {e}"
+ )
+ raise CustomCodeExecutionError(
+ f"Custom code guardrail execution failed: {e}",
+ details={
+ "guardrail_name": self.guardrail_name,
+ "input_type": input_type,
+ },
+ ) from e
+
+ def _prepare_safe_request_data(self, request_data: dict) -> Dict[str, Any]:
+ """
+ Prepare a safe subset of request_data for code execution.
+
+ This filters out sensitive information and provides only what's
+ needed for guardrail logic.
+
+ Args:
+ request_data: The full request data
+
+ Returns:
+ Safe subset of request data
+ """
+ return {
+ "model": request_data.get("model"),
+ "user_id": request_data.get("user_api_key_user_id"),
+ "team_id": request_data.get("user_api_key_team_id"),
+ "end_user_id": request_data.get("user_api_key_end_user_id"),
+ "metadata": request_data.get("metadata", {}),
+ }
+
+ def _process_result(
+ self,
+ result: Any,
+ inputs: GenericGuardrailAPIInputs,
+ request_data: dict,
+ input_type: Literal["request", "response"],
+ ) -> GenericGuardrailAPIInputs:
+ """
+ Process the result from the custom code function.
+
+ Args:
+ result: The return value from apply_guardrail
+ inputs: The original inputs
+ request_data: The request data
+ input_type: "request" or "response"
+
+ Returns:
+ GenericGuardrailAPIInputs - possibly modified
+
+ Raises:
+ HTTPException: If action is "block"
+ """
+ if not isinstance(result, dict):
+ verbose_proxy_logger.warning(
+ f"Custom code guardrail '{self.guardrail_name}': "
+ f"Expected dict result, got {type(result).__name__}. Treating as allow."
+ )
+ return inputs
+
+ action = result.get("action", "allow")
+
+ if action == "allow":
+ verbose_proxy_logger.debug(
+ f"Custom code guardrail '{self.guardrail_name}': Allowing {input_type}"
+ )
+ return inputs
+
+ elif action == "block":
+ reason = result.get("reason", "Blocked by custom code guardrail")
+ detection_info = result.get("detection_info", {})
+
+ verbose_proxy_logger.info(
+ f"Custom code guardrail '{self.guardrail_name}': Blocking {input_type} - {reason}"
+ )
+
+ is_output = input_type == "response"
+
+ # For pre-call, raise passthrough exception to return synthetic response
+ if not is_output:
+ self.raise_passthrough_exception(
+ violation_message=reason,
+ request_data=request_data,
+ detection_info=detection_info,
+ )
+
+ # For post-call, raise HTTP exception
+ raise HTTPException(
+ status_code=400,
+ detail={
+ "error": reason,
+ "guardrail": self.guardrail_name,
+ "detection_info": detection_info,
+ },
+ )
+
+ elif action == "modify":
+ verbose_proxy_logger.debug(
+ f"Custom code guardrail '{self.guardrail_name}': Modifying {input_type}"
+ )
+
+ # Apply modifications
+ modified_inputs = dict(inputs)
+
+ if "texts" in result and result["texts"] is not None:
+ modified_inputs["texts"] = result["texts"]
+
+ if "images" in result and result["images"] is not None:
+ modified_inputs["images"] = result["images"]
+
+ if "tool_calls" in result and result["tool_calls"] is not None:
+ modified_inputs["tool_calls"] = result["tool_calls"]
+
+ return cast(GenericGuardrailAPIInputs, modified_inputs)
+
+ else:
+ verbose_proxy_logger.warning(
+ f"Custom code guardrail '{self.guardrail_name}': "
+ f"Unknown action '{action}'. Treating as allow."
+ )
+ return inputs
+
+ def update_custom_code(self, new_code: str) -> None:
+ """
+ Update the custom code and recompile.
+
+ This method allows hot-reloading of guardrail logic without
+ restarting the server.
+
+ Args:
+ new_code: The new source code
+
+ Raises:
+ CustomCodeCompilationError: If the new code fails to compile
+ """
+ with self._compile_lock:
+ # Reset state
+ old_function = self._compiled_function
+ old_code = self.custom_code
+ self._compiled_function = None
+ self._compile_error = None
+
+ try:
+ self.custom_code = new_code
+ self._compile_custom_code()
+ verbose_proxy_logger.info(
+ f"Custom code guardrail '{self.guardrail_name}': Code updated successfully"
+ )
+ except CustomCodeCompilationError:
+ # Rollback on failure
+ self.custom_code = old_code
+ self._compiled_function = old_function
+ raise
diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py
new file mode 100644
index 00000000000..695e59977c8
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py
@@ -0,0 +1,602 @@
+"""
+Built-in primitives provided to custom code guardrails.
+
+These functions are injected into the custom code execution environment
+and provide safe, sandboxed functionality for common guardrail operations.
+"""
+
+import json
+import re
+from typing import Any, Dict, List, Optional, Tuple, Type, Union
+from urllib.parse import urlparse
+
+from litellm._logging import verbose_proxy_logger
+
+# =============================================================================
+# Result Types - Used by Starlark code to return guardrail decisions
+# =============================================================================
+
+
+def allow() -> Dict[str, Any]:
+ """
+ Allow the request/response to proceed unchanged.
+
+ Returns:
+ Dict indicating the request should be allowed
+ """
+ return {"action": "allow"}
+
+
+def block(
+ reason: str, detection_info: Optional[Dict[str, Any]] = None
+) -> Dict[str, Any]:
+ """
+ Block the request/response with a reason.
+
+ Args:
+ reason: Human-readable reason for blocking
+ detection_info: Optional additional detection metadata
+
+ Returns:
+ Dict indicating the request should be blocked
+ """
+ result: Dict[str, Any] = {"action": "block", "reason": reason}
+ if detection_info:
+ result["detection_info"] = detection_info
+ return result
+
+
+def modify(
+ texts: Optional[List[str]] = None,
+ images: Optional[List[Any]] = None,
+ tool_calls: Optional[List[Any]] = None,
+) -> Dict[str, Any]:
+ """
+ Modify the request/response content.
+
+ Args:
+ texts: Modified text content (if None, keeps original)
+ images: Modified image content (if None, keeps original)
+ tool_calls: Modified tool calls (if None, keeps original)
+
+ Returns:
+ Dict indicating the content should be modified
+ """
+ result: Dict[str, Any] = {"action": "modify"}
+ if texts is not None:
+ result["texts"] = texts
+ if images is not None:
+ result["images"] = images
+ if tool_calls is not None:
+ result["tool_calls"] = tool_calls
+ return result
+
+
+# =============================================================================
+# Regex Primitives
+# =============================================================================
+
+
+def regex_match(text: str, pattern: str, flags: int = 0) -> bool:
+ """
+ Check if a regex pattern matches anywhere in the text.
+
+ Args:
+ text: The text to search in
+ pattern: The regex pattern to match
+ flags: Optional regex flags (default: 0)
+
+ Returns:
+ True if pattern matches, False otherwise
+ """
+ try:
+ return bool(re.search(pattern, text, flags))
+ except re.error as e:
+ verbose_proxy_logger.warning(f"Starlark regex_match error: {e}")
+ return False
+
+
+def regex_match_all(text: str, pattern: str, flags: int = 0) -> bool:
+ """
+ Check if a regex pattern matches the entire text.
+
+ Args:
+ text: The text to match
+ pattern: The regex pattern
+ flags: Optional regex flags
+
+ Returns:
+ True if pattern matches entire text, False otherwise
+ """
+ try:
+ return bool(re.fullmatch(pattern, text, flags))
+ except re.error as e:
+ verbose_proxy_logger.warning(f"Starlark regex_match_all error: {e}")
+ return False
+
+
+def regex_replace(text: str, pattern: str, replacement: str, flags: int = 0) -> str:
+ """
+ Replace all occurrences of a pattern in text.
+
+ Args:
+ text: The text to modify
+ pattern: The regex pattern to find
+ replacement: The replacement string
+ flags: Optional regex flags
+
+ Returns:
+ The text with replacements applied
+ """
+ try:
+ return re.sub(pattern, replacement, text, flags=flags)
+ except re.error as e:
+ verbose_proxy_logger.warning(f"Starlark regex_replace error: {e}")
+ return text
+
+
+def regex_find_all(text: str, pattern: str, flags: int = 0) -> List[str]:
+ """
+ Find all occurrences of a pattern in text.
+
+ Args:
+ text: The text to search
+ pattern: The regex pattern to find
+ flags: Optional regex flags
+
+ Returns:
+ List of all matches
+ """
+ try:
+ return re.findall(pattern, text, flags)
+ except re.error as e:
+ verbose_proxy_logger.warning(f"Starlark regex_find_all error: {e}")
+ return []
+
+
+# =============================================================================
+# JSON Primitives
+# =============================================================================
+
+
+def json_parse(text: str) -> Optional[Any]:
+ """
+ Parse a JSON string into a Python object.
+
+ Args:
+ text: The JSON string to parse
+
+ Returns:
+ Parsed Python object, or None if parsing fails
+ """
+ try:
+ return json.loads(text)
+ except (json.JSONDecodeError, TypeError) as e:
+ verbose_proxy_logger.debug(f"Starlark json_parse error: {e}")
+ return None
+
+
+def json_stringify(obj: Any) -> str:
+ """
+ Convert a Python object to a JSON string.
+
+ Args:
+ obj: The object to serialize
+
+ Returns:
+ JSON string representation
+ """
+ try:
+ return json.dumps(obj)
+ except (TypeError, ValueError) as e:
+ verbose_proxy_logger.warning(f"Starlark json_stringify error: {e}")
+ return ""
+
+
+def json_schema_valid(obj: Any, schema: Dict[str, Any]) -> bool:
+ """
+ Validate an object against a JSON schema.
+
+ Args:
+ obj: The object to validate
+ schema: The JSON schema to validate against
+
+ Returns:
+ True if valid, False otherwise
+ """
+ try:
+ # Try to import jsonschema, fall back to basic validation if not available
+ try:
+ import jsonschema
+
+ jsonschema.validate(instance=obj, schema=schema)
+ return True
+ except ImportError:
+ # Basic validation without jsonschema library
+ return _basic_json_schema_validate(obj, schema)
+ except Exception as validation_error:
+ # Catch jsonschema.ValidationError and other validation errors
+ if "ValidationError" in type(validation_error).__name__:
+ return False
+ raise
+ except Exception as e:
+ verbose_proxy_logger.warning(f"Custom code json_schema_valid error: {e}")
+ return False
+
+
+def _basic_json_schema_validate(
+ obj: Any, schema: Dict[str, Any], max_depth: int = 50
+) -> bool:
+ """
+ Basic JSON schema validation without external library.
+ Handles: type, required, properties
+
+ Uses an iterative approach with a stack to avoid recursion limits.
+ max_depth limits nesting to prevent infinite loops from circular schemas.
+ """
+ type_map: Dict[str, Union[Type, Tuple[Type, ...]]] = {
+ "object": dict,
+ "array": list,
+ "string": str,
+ "number": (int, float),
+ "integer": int,
+ "boolean": bool,
+ "null": type(None),
+ }
+
+ # Stack of (obj, schema, depth) tuples to process
+ stack: List[Tuple[Any, Dict[str, Any], int]] = [(obj, schema, 0)]
+
+ while stack:
+ current_obj, current_schema, depth = stack.pop()
+
+ # Circuit breaker: stop if we've gone too deep
+ if depth > max_depth:
+ return False
+
+ # Check type
+ schema_type = current_schema.get("type")
+ if schema_type:
+ expected_type = type_map.get(schema_type)
+ if expected_type is not None and not isinstance(current_obj, expected_type):
+ return False
+
+ # Check required fields and properties for dicts
+ if isinstance(current_obj, dict):
+ required = current_schema.get("required", [])
+ for field in required:
+ if field not in current_obj:
+ return False
+
+ # Queue property validations
+ properties = current_schema.get("properties", {})
+ for prop_name, prop_schema in properties.items():
+ if prop_name in current_obj:
+ stack.append((current_obj[prop_name], prop_schema, depth + 1))
+
+ return True
+
+
+# =============================================================================
+# URL Primitives
+# =============================================================================
+
+
+# Common URL pattern for extraction
+_URL_PATTERN = re.compile(
+ r"https?://(?:[-\w.]|(?:%[\da-fA-F]{2}))+[^\s]*", re.IGNORECASE
+)
+
+
+def extract_urls(text: str) -> List[str]:
+ """
+ Extract all URLs from text.
+
+ Args:
+ text: The text to search for URLs
+
+ Returns:
+ List of URLs found in the text
+ """
+ return _URL_PATTERN.findall(text)
+
+
+def is_valid_url(url: str) -> bool:
+ """
+ Check if a URL is syntactically valid.
+
+ Args:
+ url: The URL to validate
+
+ Returns:
+ True if the URL is valid, False otherwise
+ """
+ try:
+ result = urlparse(url)
+ return all([result.scheme, result.netloc])
+ except Exception:
+ return False
+
+
+def all_urls_valid(text: str) -> bool:
+ """
+ Check if all URLs in text are valid.
+
+ Args:
+ text: The text containing URLs
+
+ Returns:
+ True if all URLs are valid (or no URLs), False otherwise
+ """
+ urls = extract_urls(text)
+ return all(is_valid_url(url) for url in urls)
+
+
+def get_url_domain(url: str) -> Optional[str]:
+ """
+ Extract the domain from a URL.
+
+ Args:
+ url: The URL to parse
+
+ Returns:
+ The domain, or None if invalid
+ """
+ try:
+ result = urlparse(url)
+ return result.netloc if result.netloc else None
+ except Exception:
+ return None
+
+
+# =============================================================================
+# Code Detection Primitives
+# =============================================================================
+
+
+# Common code patterns for detection
+_CODE_PATTERNS = {
+ "sql": [
+ r"\b(SELECT|INSERT|UPDATE|DELETE|DROP|CREATE|ALTER|TRUNCATE)\b.*\b(FROM|INTO|TABLE|SET|WHERE)\b",
+ r"\b(SELECT)\s+[\w\*,\s]+\s+FROM\s+\w+",
+ r"\b(INSERT\s+INTO|UPDATE\s+\w+\s+SET|DELETE\s+FROM)\b",
+ ],
+ "python": [
+ r"^\s*(def|class|import|from|if|for|while|try|except|with)\s+",
+ r"^\s*@\w+", # decorators
+ r"\b(print|len|range|str|int|float|list|dict|set)\s*\(",
+ ],
+ "javascript": [
+ r"\b(function|const|let|var|class|import|export)\s+",
+ r"=>", # arrow functions
+ r"\b(console\.(log|error|warn))\s*\(",
+ ],
+ "typescript": [
+ r":\s*(string|number|boolean|any|void|never)\b",
+ r"\b(interface|type|enum)\s+\w+",
+ r"<[A-Z]\w*>", # generics
+ ],
+ "java": [
+ r"\b(public|private|protected)\s+(static\s+)?(class|void|int|String)\b",
+ r"\bSystem\.(out|err)\.print",
+ ],
+ "go": [
+ r"\bfunc\s+\w+\s*\(",
+ r"\b(package|import)\s+",
+ r":=", # short variable declaration
+ ],
+ "rust": [
+ r"\b(fn|let|mut|impl|struct|enum|pub|mod)\s+",
+ r"->", # return type
+ r"\b(println!|format!)\s*\(",
+ ],
+ "shell": [
+ r"^#!.*\b(bash|sh|zsh)\b",
+ r"\b(echo|grep|sed|awk|cat|ls|cd|mkdir|rm)\s+",
+ r"\$\{?\w+\}?", # variable expansion
+ ],
+ "html": [
+ r"<\s*(html|head|body|div|span|p|a|img|script|style)\b[^>]*>",
+ r"\s*(html|head|body|div|span|p|a|script|style)\s*>",
+ ],
+ "css": [
+ r"\{[^}]*:\s*[^}]+;[^}]*\}",
+ r"@(media|keyframes|import|font-face)\b",
+ ],
+}
+
+
+def detect_code(text: str) -> bool:
+ """
+ Check if text contains code of any language.
+
+ Args:
+ text: The text to check
+
+ Returns:
+ True if code is detected, False otherwise
+ """
+ return len(detect_code_languages(text)) > 0
+
+
+def detect_code_languages(text: str) -> List[str]:
+ """
+ Detect which programming languages are present in text.
+
+ Args:
+ text: The text to analyze
+
+ Returns:
+ List of detected language names
+ """
+ detected = []
+ for lang, patterns in _CODE_PATTERNS.items():
+ for pattern in patterns:
+ try:
+ if re.search(pattern, text, re.IGNORECASE | re.MULTILINE):
+ detected.append(lang)
+ break # Only add each language once
+ except re.error:
+ continue
+ return detected
+
+
+def contains_code_language(text: str, languages: List[str]) -> bool:
+ """
+ Check if text contains code from specific languages.
+
+ Args:
+ text: The text to check
+ languages: List of language names to check for
+
+ Returns:
+ True if any of the specified languages are detected
+ """
+ detected = detect_code_languages(text)
+ return any(lang.lower() in [d.lower() for d in detected] for lang in languages)
+
+
+# =============================================================================
+# Text Utility Primitives
+# =============================================================================
+
+
+def contains(text: str, substring: str) -> bool:
+ """
+ Check if text contains a substring.
+
+ Args:
+ text: The text to search in
+ substring: The substring to find
+
+ Returns:
+ True if substring is found, False otherwise
+ """
+ return substring in text
+
+
+def contains_any(text: str, substrings: List[str]) -> bool:
+ """
+ Check if text contains any of the given substrings.
+
+ Args:
+ text: The text to search in
+ substrings: List of substrings to find
+
+ Returns:
+ True if any substring is found, False otherwise
+ """
+ return any(s in text for s in substrings)
+
+
+def contains_all(text: str, substrings: List[str]) -> bool:
+ """
+ Check if text contains all of the given substrings.
+
+ Args:
+ text: The text to search in
+ substrings: List of substrings to find
+
+ Returns:
+ True if all substrings are found, False otherwise
+ """
+ return all(s in text for s in substrings)
+
+
+def word_count(text: str) -> int:
+ """
+ Count the number of words in text.
+
+ Args:
+ text: The text to count words in
+
+ Returns:
+ Number of words
+ """
+ return len(text.split())
+
+
+def char_count(text: str) -> int:
+ """
+ Count the number of characters in text.
+
+ Args:
+ text: The text to count characters in
+
+ Returns:
+ Number of characters
+ """
+ return len(text)
+
+
+def lower(text: str) -> str:
+ """Convert text to lowercase."""
+ return text.lower()
+
+
+def upper(text: str) -> str:
+ """Convert text to uppercase."""
+ return text.upper()
+
+
+def trim(text: str) -> str:
+ """Remove leading and trailing whitespace."""
+ return text.strip()
+
+
+# =============================================================================
+# Primitives Registry
+# =============================================================================
+
+
+def get_custom_code_primitives() -> Dict[str, Any]:
+ """
+ Get all primitives to inject into the custom code environment.
+
+ Returns:
+ Dict of function name to function
+ """
+ return {
+ # Result types
+ "allow": allow,
+ "block": block,
+ "modify": modify,
+ # Regex
+ "regex_match": regex_match,
+ "regex_match_all": regex_match_all,
+ "regex_replace": regex_replace,
+ "regex_find_all": regex_find_all,
+ # JSON
+ "json_parse": json_parse,
+ "json_stringify": json_stringify,
+ "json_schema_valid": json_schema_valid,
+ # URL
+ "extract_urls": extract_urls,
+ "is_valid_url": is_valid_url,
+ "all_urls_valid": all_urls_valid,
+ "get_url_domain": get_url_domain,
+ # Code detection
+ "detect_code": detect_code,
+ "detect_code_languages": detect_code_languages,
+ "contains_code_language": contains_code_language,
+ # Text utilities
+ "contains": contains,
+ "contains_any": contains_any,
+ "contains_all": contains_all,
+ "word_count": word_count,
+ "char_count": char_count,
+ "lower": lower,
+ "upper": upper,
+ "trim": trim,
+ # Python builtins (safe subset)
+ "len": len,
+ "str": str,
+ "int": int,
+ "float": float,
+ "bool": bool,
+ "list": list,
+ "dict": dict,
+ "True": True,
+ "False": False,
+ "None": None,
+ }
diff --git a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py
index 2a852cbda08..90f689ed23c 100644
--- a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py
+++ b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py
@@ -9,8 +9,10 @@ from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
+ ModifyResponseException
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
+from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@@ -21,6 +23,8 @@ from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
+GRAYSWAN_BLOCK_ERROR_MSG = "Blocked by Gray Swan Guardrail"
+
class GraySwanGuardrailMissingSecrets(Exception):
"""Raised when the Gray Swan API key is missing."""
@@ -205,9 +209,13 @@ class GraySwanGuardrail(CustomGuardrail):
# Get dynamic params from request metadata
dynamic_body = self.get_guardrail_dynamic_request_body_params(request_data) or {}
+ if dynamic_body:
+ verbose_proxy_logger.debug(
+ "Gray Swan Guardrail: dynamic extra_body=%s", safe_dumps(dynamic_body)
+ )
# Prepare and send payload
- payload = self._prepare_payload(messages, dynamic_body)
+ payload = self._prepare_payload(messages, dynamic_body, request_data)
if payload is None:
return inputs
@@ -223,6 +231,8 @@ class GraySwanGuardrail(CustomGuardrail):
)
return result
except Exception as exc:
+ if self._is_grayswan_exception(exc):
+ raise
end_time = time.time()
status_code = getattr(exc, "status_code", None) or getattr(
exc, "exception_status_code", None
@@ -240,8 +250,20 @@ class GraySwanGuardrail(CustomGuardrail):
exc,
)
return inputs
+ if isinstance(exc, GraySwanGuardrailAPIError):
+ raise exc
raise GraySwanGuardrailAPIError(str(exc), status_code=status_code) from exc
+ def _is_grayswan_exception(self, exc: Exception) -> bool:
+ # Guardrail decision (passthrough) should always propagate,
+ # regardless of fail_open.
+ if isinstance(exc, ModifyResponseException):
+ return True
+ detail = getattr(exc, "detail", None)
+ if isinstance(detail, dict):
+ return detail.get("error") == GRAYSWAN_BLOCK_ERROR_MSG
+ return False
+
# ------------------------------------------------------------------
# Legacy Test Interface (for backward compatibility)
# ------------------------------------------------------------------
@@ -324,7 +346,7 @@ class GraySwanGuardrail(CustomGuardrail):
raise HTTPException(
status_code=400,
detail={
- "error": "Blocked by Gray Swan Guardrail",
+ "error": GRAYSWAN_BLOCK_ERROR_MSG,
"violation_location": violation_location,
"violation": violation_score,
"violated_rules": violated_rules,
@@ -445,7 +467,7 @@ class GraySwanGuardrail(CustomGuardrail):
raise HTTPException(
status_code=400,
detail={
- "error": "Blocked by Gray Swan Guardrail",
+ "error": GRAYSWAN_BLOCK_ERROR_MSG,
"violation_location": violation_location,
"violation": violation_score,
"violated_rules": violated_rules,
@@ -494,7 +516,7 @@ class GraySwanGuardrail(CustomGuardrail):
}
def _prepare_payload(
- self, messages: List[Dict[str, str]], dynamic_body: dict
+ self, messages: List[Dict[str, str]], dynamic_body: dict, request_data: dict
) -> Optional[Dict[str, Any]]:
payload: Dict[str, Any] = {"messages": messages}
@@ -510,6 +532,18 @@ class GraySwanGuardrail(CustomGuardrail):
if reasoning_mode:
payload["reasoning_mode"] = reasoning_mode
+ # Pass through arbitrary metadata when provided via dynamic extra_body.
+ if "metadata" in dynamic_body:
+ payload["metadata"] = dynamic_body["metadata"]
+
+ litellm_metadata = request_data.get("litellm_metadata")
+ if isinstance(litellm_metadata, dict) and litellm_metadata:
+ cleaned_litellm_metadata = dict(litellm_metadata)
+ # cleaned_litellm_metadata.pop("user_api_key_auth", None)
+ sanitized = safe_json_loads(safe_dumps(cleaned_litellm_metadata), default={})
+ if isinstance(sanitized, dict) and sanitized:
+ payload["litellm_metadata"] = sanitized
+
return payload
def _format_violation_message(
diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py
index a12eb2486d2..38462094b11 100644
--- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py
+++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py
@@ -421,6 +421,13 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
)
else "success"
)
+
+ # Add guardrail to applied_guardrails BEFORE potential blocking
+ # This ensures guardrail is recorded even when it blocks the request
+ add_guardrail_to_applied_guardrails_header(
+ request_data=data, guardrail_name=self.guardrail_name
+ )
+
# Check if content should be blocked
if self._should_block_content(
armor_response, allow_sanitization=self.mask_request_content
@@ -456,11 +463,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
if self.optional_params.get("fail_on_error", True):
raise
- # Add guardrail to headers
- add_guardrail_to_applied_guardrails_header(
- request_data=data, guardrail_name=self.guardrail_name
- )
-
return data
@log_guardrail_information
@@ -517,6 +519,12 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
else "success"
)
+ # Add guardrail to applied_guardrails BEFORE potential blocking
+ # This ensures guardrail is recorded even when it blocks the request
+ add_guardrail_to_applied_guardrails_header(
+ request_data=data, guardrail_name=self.guardrail_name
+ )
+
# Check if content should be blocked
if self._should_block_content(
armor_response, allow_sanitization=self.mask_request_content
@@ -550,11 +558,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
if self.optional_params.get("fail_on_error", True):
raise
- # Add guardrail to headers
- add_guardrail_to_applied_guardrails_header(
- request_data=data, guardrail_name=self.guardrail_name
- )
-
return data
@log_guardrail_information
@@ -622,6 +625,12 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
guardrail_response=standard_logging_guardrail_information,
)
+ # Add guardrail to applied_guardrails BEFORE potential blocking
+ # This ensures guardrail is recorded even when it blocks the request
+ add_guardrail_to_applied_guardrails_header(
+ request_data=data, guardrail_name=self.guardrail_name
+ )
+
# Check if content should be blocked
if self._should_block_content(
armor_response, allow_sanitization=self.mask_response_content
@@ -654,11 +663,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
if self.optional_params.get("fail_on_error", True):
raise
- # Add guardrail to headers
- add_guardrail_to_applied_guardrails_header(
- request_data=data, guardrail_name=self.guardrail_name
- )
-
return response
async def async_post_call_streaming_iterator_hook(
@@ -703,6 +707,16 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
else "success"
)
+ # Add guardrail to applied_guardrails BEFORE potential blocking
+ # This ensures guardrail is recorded even when it blocks the request
+ from litellm.proxy.common_utils.callback_utils import (
+ add_guardrail_to_applied_guardrails_header,
+ )
+
+ add_guardrail_to_applied_guardrails_header(
+ request_data=request_data, guardrail_name=self.guardrail_name
+ )
+
# Check if blocked
if self._should_block_content(armor_response):
raise HTTPException(
diff --git a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py
index 852cf01bc09..030b6036815 100644
--- a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py
+++ b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py
@@ -22,16 +22,18 @@ from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
+from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
+from litellm.types.utils import GenericGuardrailAPIInputs
from .base import OpenAIGuardrailBase
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
- from litellm.types.llms.openai import AllMessageValues, OpenAIModerationResponse
+ from litellm.types.llms.openai import OpenAIModerationResponse
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
from litellm.types.utils import ModelResponse, ModelResponseStream
@@ -178,170 +180,59 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
},
)
- def _extract_user_content_from_data(self, data: Dict[str, Any]) -> Optional[str]:
+ async def apply_guardrail(
+ self,
+ inputs: GenericGuardrailAPIInputs,
+ request_data: dict,
+ input_type: Literal["request", "response"],
+ logging_obj: Optional["LiteLLMLoggingObj"] = None,
+ ) -> GenericGuardrailAPIInputs:
"""
- Extract user content from request data, supporting both Chat Completions and Responses API.
+ Apply OpenAI moderation guardrail using the unified guardrail interface.
- For Chat Completions: extracts from 'messages' field
- For Responses API: extracts from 'input' field
+ This method is called by the UnifiedLLMGuardrails system for all endpoint types
+ (chat completions, embeddings, responses API, etc.).
+ Args:
+ inputs: GenericGuardrailAPIInputs containing texts and/or structured_messages
+ request_data: The original request data
+ input_type: Whether this is a "request" (pre-call) or "response" (post-call)
+ logging_obj: Optional logging object
+
Returns:
- The extracted user content string, or None if no content found
+ The inputs unchanged (moderation doesn't modify content, only blocks)
+
+ Raises:
+ HTTPException: If content violates moderation policy
"""
- # Try to get messages first (Chat Completions API)
- messages: Optional[List["AllMessageValues"]] = data.get("messages")
- if messages is not None:
- return self.get_user_prompt(messages)
+ # Extract text to moderate from inputs
+ text_to_moderate: Optional[str] = None
- # Try to get input (Responses API)
- input_data = data.get("input")
- if input_data is not None:
- # input can be a string or a list of message-like objects
- if isinstance(input_data, str):
- return input_data
- elif isinstance(input_data, list):
- # Treat input as messages and extract user content
- return self.get_user_prompt(input_data)
+ # Prefer structured_messages if available (has role context)
+ if structured_messages := inputs.get("structured_messages"):
+ text_to_moderate = self.get_user_prompt(structured_messages)
- return None
-
- @log_guardrail_information
- async def async_pre_call_hook(
- self,
- user_api_key_dict: "UserAPIKeyAuth",
- cache: Any,
- data: Dict[str, Any],
- call_type: Literal[
- "completion",
- "text_completion",
- "embeddings",
- "image_generation",
- "moderation",
- "audio_transcription",
- "pass_through_endpoint",
- "rerank",
- "mcp_call",
- ],
- ) -> Optional[Dict[str, Any]]:
- """
- Pre-call hook to scan user prompts before sending to LLM.
-
- Raises HTTPException if content should be blocked.
- """
- verbose_proxy_logger.debug(
- "OpenAI Moderation: Running pre-call prompt scan, on call_type: %s",
- call_type,
- )
+ # Fall back to texts
+ if not text_to_moderate:
+ if texts := inputs.get("texts"):
+ # Join all texts for moderation
+ text_to_moderate = "\n".join(texts)
- # Skip moderation calls to avoid infinite recursion
- if call_type == "moderation":
- return data
-
- user_prompt = self._extract_user_content_from_data(data)
-
- if user_prompt is None:
- verbose_proxy_logger.warning(
- "OpenAI Moderation: not running guardrail. No messages or input in data"
- )
- return data
-
- if user_prompt:
+ if not text_to_moderate:
verbose_proxy_logger.debug(
- f"OpenAI Moderation: User prompt: {user_prompt[:100]}..." # Log first 100 chars for debugging
+ "OpenAI Moderation: No text content to moderate in inputs"
)
-
- moderation_response = await self.async_make_request(
- input_text=user_prompt,
- )
-
- # Check if content is flagged and raise exception if needed
- self._check_moderation_result(moderation_response)
- else:
- verbose_proxy_logger.warning(
- "OpenAI Moderation: No user prompt found"
- )
-
- return data
-
- @log_guardrail_information
- async def async_moderation_hook(
- self,
- data: Dict[str, Any],
- user_api_key_dict: "UserAPIKeyAuth",
- call_type: Literal[
- "completion",
- "embeddings",
- "image_generation",
- "moderation",
- "audio_transcription",
- "responses",
- "mcp_call",
- ],
- ) -> Optional[Dict[str, Any]]:
- """
- Moderation hook to scan user prompts during call processing.
-
- Raises HTTPException if content should be blocked.
- """
- verbose_proxy_logger.debug(
- "OpenAI Moderation: Running moderation hook, on call_type: %s",
- call_type,
- )
+ return inputs
+
+ # Make moderation request
+ moderation_response = await self.async_make_request(input_text=text_to_moderate)
- # Skip moderation calls to avoid infinite recursion
- if call_type == "moderation":
- return data
-
- # Extract user content from either messages or input field
- user_prompt = self._extract_user_content_from_data(data)
+ # Check if content is flagged and raise exception if needed
+ self._check_moderation_result(moderation_response)
- if user_prompt is None:
- verbose_proxy_logger.warning(
- "OpenAI Moderation: not running guardrail. No messages or input in data"
- )
- return data
+ # Moderation doesn't modify content, just blocks - return inputs unchanged
+ return inputs
- if user_prompt:
- moderation_response = await self.async_make_request(
- input_text=user_prompt,
- )
-
- # Check if content is flagged and raise exception if needed
- self._check_moderation_result(moderation_response)
-
- return data
-
- @log_guardrail_information
- async def async_post_call_hook(
- self,
- data: Dict[str, Any],
- user_api_key_dict: "UserAPIKeyAuth",
- response: "ModelResponse",
- ) -> "ModelResponse":
- """
- Post-call hook to scan LLM responses before returning to user.
-
- Raises HTTPException if response should be blocked.
- """
- verbose_proxy_logger.debug(
- "OpenAI Moderation: Running post-call response scan"
- )
-
- # Extract response text for moderation
- response_text = self._extract_response_text(response)
- if response_text:
- verbose_proxy_logger.debug(
- f"OpenAI Moderation: Response text: {response_text[:100]}..." # Log first 100 chars
- )
-
- moderation_response = await self.async_make_request(
- input_text=response_text,
- )
-
- # Check if content is flagged and raise exception if needed
- self._check_moderation_result(moderation_response)
-
- return response
@log_guardrail_information
async def async_post_call_streaming_iterator_hook(
diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py
index 80f9860bdff..f07f65d10f5 100644
--- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py
+++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py
@@ -6,6 +6,7 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint
3. Implements a way to call /applyGuardrail endpoint for `/chat/completions` + `/v1/messages` requests on async_post_call_streaming_iterator_hook
"""
+import copy
from typing import Any, AsyncGenerator, List, Optional, Union
from litellm._logging import verbose_proxy_logger
@@ -349,22 +350,26 @@ class UnifiedLLMGuardrails(CustomLogger):
guardrail_to_apply.guardrail_name,
)
+ # Deep-copy the current chunk before guardrail processing.
+ # process_output_streaming_response modifies responses_so_far
+ # in-place: it puts the combined guardrailed text in the first
+ # chunk and clears all subsequent chunks to "". Without this
+ # copy, yielding processed_items[-1] would yield an empty
+ # string, permanently losing this chunk's content.
+ original_item = copy.deepcopy(item)
+
endpoint_translation = endpoint_guardrail_translation_mappings[
CallTypes(call_type)
]()
- processed_items = (
- await endpoint_translation.process_output_streaming_response(
- responses_so_far=responses_so_far,
- guardrail_to_apply=guardrail_to_apply,
- litellm_logging_obj=request_data.get("litellm_logging_obj"),
- user_api_key_dict=user_api_key_dict,
- )
+ await endpoint_translation.process_output_streaming_response(
+ responses_so_far=responses_so_far,
+ guardrail_to_apply=guardrail_to_apply,
+ litellm_logging_obj=request_data.get("litellm_logging_obj"),
+ user_api_key_dict=user_api_key_dict,
)
- last_item = processed_items[-1]
-
- yield last_item
+ yield original_item
else:
yield item
diff --git a/litellm/proxy/hooks/mcp_semantic_filter/ARCHITECTURE.md b/litellm/proxy/hooks/mcp_semantic_filter/ARCHITECTURE.md
new file mode 100644
index 00000000000..f2f9a1d4856
--- /dev/null
+++ b/litellm/proxy/hooks/mcp_semantic_filter/ARCHITECTURE.md
@@ -0,0 +1,96 @@
+# MCP Semantic Tool Filter Architecture
+
+## Why Filter MCP Tools
+
+When multiple MCP servers are connected, the proxy may expose hundreds of tools. Sending all tools in every request wastes context window tokens and increases cost. The semantic filter keeps only the top-K most relevant tools based on embedding similarity.
+
+```mermaid
+sequenceDiagram
+ participant Client
+ participant Hook as SemanticToolFilterHook
+ participant Filter as SemanticMCPToolFilter
+ participant Router as semantic-router
+ participant LLM
+
+ Client->>Hook: POST /chat/completions
+ Note over Client,Hook: tools: [100+ MCP tools]
+ Note over Client,Hook: messages: [{"role": "user", "content": "Get my Jira issues"}]
+
+ rect rgb(240, 240, 240)
+ Note over Hook: 1. Extract User Query
+ Hook->>Filter: filter_tools("Get my Jira issues", tools)
+ end
+
+ rect rgb(240, 240, 240)
+ Note over Filter: 2. Convert Tools → Routes
+ Note over Filter: Tool name + description → Route
+ end
+
+ rect rgb(240, 240, 240)
+ Note over Filter: 3. Semantic Matching
+ Filter->>Router: router(query)
+ Router->>Router: Embeddings + similarity
+ Router-->>Filter: [top 10 matches]
+ end
+
+ rect rgb(240, 240, 240)
+ Note over Filter: 4. Return Filtered Tools
+ Filter-->>Hook: [10 relevant tools]
+ end
+
+ Hook->>LLM: POST /chat/completions
+ Note over Hook,LLM: tools: [10 Jira-related tools] ← FILTERED
+ Note over Hook,LLM: messages: [...] ← UNCHANGED
+
+ LLM-->>Client: Response (unchanged)
+```
+
+## Filter Operations
+
+The hook intercepts requests before they reach the LLM:
+
+| Operation | Description |
+|-----------|-------------|
+| **Extract query** | Get user message from `messages[-1]` |
+| **Convert to Routes** | Transform MCP tools into semantic-router Routes |
+| **Semantic match** | Use `semantic-router` to find top-K similar tools |
+| **Filter tools** | Replace request `tools` with filtered subset |
+
+## Trigger Conditions
+
+The filter only runs when:
+- Call type is `completion` or `acompletion`
+- Request contains `tools` field
+- Request contains `messages` field
+- Filter is enabled in config
+
+## What Does NOT Change
+
+- Request messages
+- Response body
+- Non-tool parameters
+
+## Integration with semantic-router
+
+Reuses existing LiteLLM infrastructure:
+- `semantic-router` - Already an optional dependency
+- `LiteLLMRouterEncoder` - Wraps `Router.aembedding()` for embeddings
+- `SemanticRouter` - Handles similarity calculation and top-K selection
+
+## Configuration
+
+```yaml
+litellm_settings:
+ mcp_semantic_tool_filter:
+ enabled: true
+ embedding_model: "openai/text-embedding-3-small"
+ top_k: 10
+ similarity_threshold: 0.3
+```
+
+## Error Handling
+
+The filter fails gracefully:
+- If filtering fails → Return all tools (no impact on functionality)
+- If query extraction fails → Skip filtering
+- If no matches found → Return all tools
diff --git a/litellm/proxy/hooks/mcp_semantic_filter/__init__.py b/litellm/proxy/hooks/mcp_semantic_filter/__init__.py
new file mode 100644
index 00000000000..36d357d560f
--- /dev/null
+++ b/litellm/proxy/hooks/mcp_semantic_filter/__init__.py
@@ -0,0 +1,9 @@
+"""
+MCP Semantic Tool Filter Hook
+
+Semantic filtering for MCP tools to reduce context window size
+and improve tool selection accuracy.
+"""
+from litellm.proxy.hooks.mcp_semantic_filter.hook import SemanticToolFilterHook
+
+__all__ = ["SemanticToolFilterHook"]
diff --git a/litellm/proxy/hooks/mcp_semantic_filter/hook.py b/litellm/proxy/hooks/mcp_semantic_filter/hook.py
new file mode 100644
index 00000000000..fc9349c2a42
--- /dev/null
+++ b/litellm/proxy/hooks/mcp_semantic_filter/hook.py
@@ -0,0 +1,353 @@
+"""
+Semantic Tool Filter Hook
+
+Pre-call hook that filters MCP tools semantically before LLM inference.
+Reduces context window size and improves tool selection accuracy.
+"""
+from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
+
+from litellm._logging import verbose_proxy_logger
+from litellm.constants import (
+ DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL,
+ DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD,
+ DEFAULT_MCP_SEMANTIC_FILTER_TOP_K,
+)
+from litellm.integrations.custom_logger import CustomLogger
+
+if TYPE_CHECKING:
+ from litellm.caching.caching import DualCache
+ from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
+ SemanticMCPToolFilter,
+ )
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.router import Router
+
+
+class SemanticToolFilterHook(CustomLogger):
+ """
+ Pre-call hook that filters MCP tools semantically.
+
+ This hook:
+ 1. Extracts the user query from messages
+ 2. Filters tools based on semantic similarity to the query
+ 3. Returns only the top-k most relevant tools to the LLM
+ """
+
+ def __init__(self, semantic_filter: "SemanticMCPToolFilter"):
+ """
+ Initialize the hook.
+
+ Args:
+ semantic_filter: SemanticMCPToolFilter instance
+ """
+ super().__init__()
+ self.filter = semantic_filter
+
+ verbose_proxy_logger.debug(
+ f"Initialized SemanticToolFilterHook with filter: "
+ f"enabled={semantic_filter.enabled}, top_k={semantic_filter.top_k}"
+ )
+
+ def _should_expand_mcp_tools(self, tools: List[Any]) -> bool:
+ """
+ Check if tools contain MCP references with server_url="litellm_proxy".
+
+ Only expands MCP tools pointing to litellm proxy, not external MCP servers.
+ """
+ from litellm.responses.mcp.litellm_proxy_mcp_handler import (
+ LiteLLM_Proxy_MCP_Handler,
+ )
+
+ return LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools)
+
+ async def _expand_mcp_tools(
+ self,
+ tools: List[Any],
+ user_api_key_dict: "UserAPIKeyAuth",
+ ) -> List[Dict[str, Any]]:
+ """
+ Expand MCP references to actual tool definitions.
+
+ Reuses LiteLLM_Proxy_MCP_Handler._process_mcp_tools_to_openai_format
+ which internally does: parse -> fetch -> filter -> deduplicate -> transform
+ """
+ from litellm.responses.mcp.litellm_proxy_mcp_handler import (
+ LiteLLM_Proxy_MCP_Handler,
+ )
+
+ # Parse to separate MCP tools from other tools
+ mcp_tools, _ = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools)
+
+ if not mcp_tools:
+ return []
+
+ # Use single combined method instead of 3 separate calls
+ # This already handles: fetch -> filter by allowed_tools -> deduplicate -> transform
+ openai_tools, _ = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_to_openai_format(
+ user_api_key_auth=user_api_key_dict,
+ mcp_tools_with_litellm_proxy=mcp_tools
+ )
+
+ # Convert Pydantic models to dicts for compatibility
+ openai_tools_as_dicts = []
+ for tool in openai_tools:
+ if hasattr(tool, "model_dump"):
+ tool_dict = tool.model_dump(exclude_none=True)
+ verbose_proxy_logger.debug(f"Converted Pydantic tool to dict: {type(tool).__name__} -> dict with keys: {list(tool_dict.keys())}")
+ openai_tools_as_dicts.append(tool_dict)
+ elif hasattr(tool, "dict"):
+ tool_dict = tool.dict(exclude_none=True)
+ verbose_proxy_logger.debug(f"Converted Pydantic tool (v1) to dict: {type(tool).__name__} -> dict")
+ openai_tools_as_dicts.append(tool_dict)
+ elif isinstance(tool, dict):
+ verbose_proxy_logger.debug(f"Tool is already a dict with keys: {list(tool.keys())}")
+ openai_tools_as_dicts.append(tool)
+ else:
+ verbose_proxy_logger.warning(f"Tool is unknown type: {type(tool)}, passing as-is")
+ openai_tools_as_dicts.append(tool)
+
+ verbose_proxy_logger.debug(
+ f"Expanded {len(mcp_tools)} MCP reference(s) to {len(openai_tools_as_dicts)} tools (all as dicts)"
+ )
+
+ return openai_tools_as_dicts
+
+ def _get_metadata_variable_name(self, data: dict) -> str:
+ if "litellm_metadata" in data:
+ return "litellm_metadata"
+ return "metadata"
+
+ async def async_pre_call_hook(
+ self,
+ user_api_key_dict: "UserAPIKeyAuth",
+ cache: "DualCache",
+ data: dict,
+ call_type: str,
+ ) -> Optional[Union[Exception, str, dict]]:
+ """
+ Filter tools before LLM call based on user query.
+
+ This hook is called before the LLM request is made. It filters the
+ tools list to only include semantically relevant tools.
+
+ Args:
+ user_api_key_dict: User authentication
+ cache: Cache instance
+ data: Request data containing messages and tools
+ call_type: Type of call (completion, acompletion, etc.)
+
+ Returns:
+ Modified data dict with filtered tools, or None if no changes
+ """
+ # Only filter endpoints that support tools
+ if call_type not in ("completion", "acompletion", "aresponses"):
+ verbose_proxy_logger.debug(
+ f"Skipping semantic filter for call_type={call_type}"
+ )
+ return None
+
+ # Check if tools are present
+ tools = data.get("tools")
+ if not tools:
+ verbose_proxy_logger.debug("No tools in request, skipping semantic filter")
+ return None
+
+ original_tool_count = len(tools)
+
+ # Check for MCP references (server_url="litellm_proxy") and expand them
+ if self._should_expand_mcp_tools(tools):
+ verbose_proxy_logger.debug(
+ "Detected litellm_proxy MCP references, expanding before semantic filtering"
+ )
+
+ try:
+ expanded_tools = await self._expand_mcp_tools(
+ tools, user_api_key_dict
+ )
+
+ if not expanded_tools:
+ verbose_proxy_logger.warning(
+ "No tools expanded from MCP references"
+ )
+ return None
+
+ verbose_proxy_logger.info(
+ f"Expanded {len(tools)} MCP reference(s) to {len(expanded_tools)} tools"
+ )
+
+ # Update tools for filtering
+ tools = expanded_tools
+ original_tool_count = len(tools)
+
+ except Exception as e:
+ verbose_proxy_logger.error(
+ f"Failed to expand MCP references: {e}", exc_info=True
+ )
+ return None
+
+ # Check if messages are present (try both "messages" and "input" for responses API)
+ messages = data.get("messages", [])
+ if not messages:
+ messages = data.get("input", [])
+ if not messages:
+ verbose_proxy_logger.debug("No messages in request, skipping semantic filter")
+ return None
+
+ # Check if filter is enabled
+ if not self.filter.enabled:
+ verbose_proxy_logger.debug("Semantic filter disabled, skipping")
+ return None
+
+ try:
+ # Extract user query from messages
+ user_query = self.filter.extract_user_query(messages)
+ if not user_query:
+ verbose_proxy_logger.debug("No user query found, skipping semantic filter")
+ return None
+
+ verbose_proxy_logger.debug(
+ f"Applying semantic filter to {len(tools)} tools "
+ f"with query: '{user_query[:50]}...'"
+ )
+
+ # Filter tools semantically
+ filtered_tools = await self.filter.filter_tools(
+ query=user_query,
+ available_tools=tools, # type: ignore
+ )
+
+ # Always update tools and emit header (even if count unchanged)
+ data["tools"] = filtered_tools
+
+ # Store filter stats and tool names for response header
+ filter_stats = f"{original_tool_count}->{len(filtered_tools)}"
+ tool_names_csv = self._get_tool_names_csv(filtered_tools)
+
+ _metadata_variable_name = self._get_metadata_variable_name(data)
+ data[_metadata_variable_name]["litellm_semantic_filter_stats"] = filter_stats
+ data[_metadata_variable_name]["litellm_semantic_filter_tools"] = tool_names_csv
+
+ verbose_proxy_logger.info(
+ f"Semantic tool filter: {filter_stats} tools"
+ )
+
+ return data
+
+ except Exception as e:
+ verbose_proxy_logger.warning(
+ f"Semantic tool filter hook failed: {e}. Proceeding with all tools."
+ )
+ return None
+
+ async def async_post_call_response_headers_hook(
+ self,
+ data: dict,
+ user_api_key_dict: "UserAPIKeyAuth",
+ response: Any,
+ request_headers: Optional[Dict[str, str]] = None,
+ ) -> Optional[Dict[str, str]]:
+ """Add semantic filter stats and tool names to response headers."""
+ from litellm.constants import MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH
+
+ _metadata_variable_name = self._get_metadata_variable_name(data)
+ metadata = data[_metadata_variable_name]
+
+ filter_stats = metadata.get("litellm_semantic_filter_stats")
+ if not filter_stats:
+ return None
+
+ headers = {"x-litellm-semantic-filter": filter_stats}
+
+ # Add CSV of filtered tool names (nginx-safe length)
+ tool_names_csv = metadata.get("litellm_semantic_filter_tools", "")
+ if tool_names_csv:
+ if len(tool_names_csv) > MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH:
+ tool_names_csv = tool_names_csv[:MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH - 3] + "..."
+
+ headers["x-litellm-semantic-filter-tools"] = tool_names_csv
+
+ return headers
+
+ def _get_tool_names_csv(self, tools: List[Any]) -> str:
+ """Extract tool names and return as CSV string."""
+ if not tools:
+ return ""
+
+ tool_names = []
+ for tool in tools:
+ name = tool.get("name", "") if isinstance(tool, dict) else getattr(tool, "name", "")
+ if name:
+ tool_names.append(name)
+
+ return ",".join(tool_names)
+
+ @staticmethod
+ async def initialize_from_config(
+ config: Optional[Dict[str, Any]],
+ llm_router: Optional["Router"],
+ ) -> Optional["SemanticToolFilterHook"]:
+ """
+ Initialize semantic tool filter from proxy config.
+
+ Args:
+ config: Proxy configuration dict (litellm_settings.mcp_semantic_tool_filter)
+ llm_router: LiteLLM router instance for embeddings
+
+ Returns:
+ SemanticToolFilterHook instance if enabled, None otherwise
+ """
+ from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
+ SemanticMCPToolFilter,
+ )
+ if not config or not config.get("enabled", False):
+ verbose_proxy_logger.debug("Semantic tool filter not enabled in config")
+ return None
+
+ if llm_router is None:
+ verbose_proxy_logger.warning(
+ "Cannot initialize semantic filter: llm_router is None"
+ )
+ return None
+
+ try:
+
+ embedding_model = config.get(
+ "embedding_model", DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL
+ )
+ top_k = config.get("top_k", DEFAULT_MCP_SEMANTIC_FILTER_TOP_K)
+ similarity_threshold = config.get(
+ "similarity_threshold", DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD
+ )
+
+ semantic_filter = SemanticMCPToolFilter(
+ embedding_model=embedding_model,
+ litellm_router_instance=llm_router,
+ top_k=top_k,
+ similarity_threshold=similarity_threshold,
+ enabled=True,
+ )
+
+ # Build router from MCP registry on startup
+ await semantic_filter.build_router_from_mcp_registry()
+
+ hook = SemanticToolFilterHook(semantic_filter)
+
+ verbose_proxy_logger.info(
+ f"✅ MCP Semantic Tool Filter enabled: "
+ f"embedding_model={embedding_model}, top_k={top_k}, "
+ f"similarity_threshold={similarity_threshold}"
+ )
+
+ return hook
+
+ except ImportError as e:
+ verbose_proxy_logger.warning(
+ f"semantic-router not installed. Install with: "
+ f"pip install 'litellm[semantic-router]'. Error: {e}"
+ )
+ return None
+ except Exception as e:
+ verbose_proxy_logger.exception(
+ f"Failed to initialize MCP semantic tool filter: {e}"
+ )
+ return None
diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py
index c52491efc7c..99a732f9efb 100644
--- a/litellm/proxy/management_endpoints/common_daily_activity.py
+++ b/litellm/proxy/management_endpoints/common_daily_activity.py
@@ -1,5 +1,5 @@
-from datetime import datetime
-from typing import Any, Callable, Dict, List, Optional, Set, Union
+from datetime import datetime, timedelta
+from typing import Any, Callable, Dict, List, Optional, Set, Tuple, Union
from fastapi import HTTPException, status
@@ -336,6 +336,46 @@ async def get_api_key_metadata(
}
+def _adjust_dates_for_timezone(
+ start_date: str,
+ end_date: str,
+ timezone_offset_minutes: Optional[int],
+) -> Tuple[str, str]:
+ """
+ Adjust date range to account for timezone differences.
+
+ The database stores dates in UTC. When a user in a different timezone
+ selects a local date range, we need to expand the UTC query range to
+ capture all records that fall within their local date range.
+
+ Args:
+ start_date: Start date in YYYY-MM-DD format (user's local date)
+ end_date: End date in YYYY-MM-DD format (user's local date)
+ timezone_offset_minutes: Minutes behind UTC (positive = west of UTC)
+ This matches JavaScript's Date.getTimezoneOffset() convention.
+ For example: PST = +480 (8 hours * 60 = 480 minutes behind UTC)
+
+ Returns:
+ Tuple of (adjusted_start_date, adjusted_end_date) in YYYY-MM-DD format
+ """
+ if timezone_offset_minutes is None or timezone_offset_minutes == 0:
+ return start_date, end_date
+
+ start = datetime.strptime(start_date, "%Y-%m-%d")
+ end = datetime.strptime(end_date, "%Y-%m-%d")
+
+ if timezone_offset_minutes > 0:
+ # West of UTC (Americas): local evening extends into next UTC day
+ # e.g., Feb 4 23:59 PST = Feb 5 07:59 UTC
+ end = end + timedelta(days=1)
+ else:
+ # East of UTC (Asia/Europe): local morning starts in previous UTC day
+ # e.g., Feb 4 00:00 IST = Feb 3 18:30 UTC
+ start = start - timedelta(days=1)
+
+ return start.strftime("%Y-%m-%d"), end.strftime("%Y-%m-%d")
+
+
def _build_where_conditions(
*,
entity_id_field: str,
@@ -345,12 +385,18 @@ def _build_where_conditions(
model: Optional[str],
api_key: Optional[Union[str, List[str]]],
exclude_entity_ids: Optional[List[str]] = None,
+ timezone_offset_minutes: Optional[int] = None,
) -> Dict[str, Any]:
"""Build prisma where clause for daily activity queries."""
+ # Adjust dates for timezone if provided
+ adjusted_start, adjusted_end = _adjust_dates_for_timezone(
+ start_date, end_date, timezone_offset_minutes
+ )
+
where_conditions: Dict[str, Any] = {
"date": {
- "gte": start_date,
- "lte": end_date,
+ "gte": adjusted_start,
+ "lte": adjusted_end,
}
}
@@ -453,6 +499,7 @@ async def get_daily_activity(
page_size: int,
exclude_entity_ids: Optional[List[str]] = None,
metadata_metrics_func: Optional[Callable[[List[Any]], SpendMetrics]] = None,
+ timezone_offset_minutes: Optional[int] = None,
) -> SpendAnalyticsPaginatedResponse:
"""Common function to get daily activity for any entity type."""
@@ -477,6 +524,7 @@ async def get_daily_activity(
model=model,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
+ timezone_offset_minutes=timezone_offset_minutes,
)
# Get total count for pagination
@@ -542,6 +590,7 @@ async def get_daily_activity_aggregated(
model: Optional[str],
api_key: Optional[str],
exclude_entity_ids: Optional[List[str]] = None,
+ timezone_offset_minutes: Optional[int] = None,
) -> SpendAnalyticsPaginatedResponse:
"""Aggregated variant that returns the full result set (no pagination).
@@ -568,6 +617,7 @@ async def get_daily_activity_aggregated(
model=model,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
+ timezone_offset_minutes=timezone_offset_minutes,
)
# Fetch all matching results (no pagination)
diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py
index 38a867d031b..c0285407855 100644
--- a/litellm/proxy/management_endpoints/internal_user_endpoints.py
+++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py
@@ -813,9 +813,12 @@ def _update_internal_user_params(
data_json: dict, data: Union[UpdateUserRequest, UpdateUserRequestNoUserIDorEmail]
) -> dict:
non_default_values = {}
+ fields_set = data.fields_set() if hasattr(data, 'fields_set') else set()
+
for k, v in data_json.items():
if k == "max_budget":
- non_default_values[k] = v
+ if "max_budget" in fields_set:
+ non_default_values[k] = v
elif (
v is not None
and v
@@ -1914,6 +1917,11 @@ async def get_user_daily_activity(
page_size: int = fastapi.Query(
default=50, description="Items per page", ge=1, le=1000
),
+ timezone: Optional[int] = fastapi.Query(
+ default=None,
+ description="Timezone offset in minutes from UTC (e.g., 480 for PST). "
+ "Matches JavaScript's Date.getTimezoneOffset() convention.",
+ ),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> SpendAnalyticsPaginatedResponse:
"""
@@ -1963,6 +1971,7 @@ async def get_user_daily_activity(
api_key=api_key,
page=page,
page_size=page_size,
+ timezone_offset_minutes=timezone,
)
except Exception as e:
@@ -1999,6 +2008,11 @@ async def get_user_daily_activity_aggregated(
default=None,
description="Filter by specific API key",
),
+ timezone: Optional[int] = fastapi.Query(
+ default=None,
+ description="Timezone offset in minutes from UTC (e.g., 480 for PST). "
+ "Matches JavaScript's Date.getTimezoneOffset() convention.",
+ ),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> SpendAnalyticsPaginatedResponse:
"""
@@ -2034,6 +2048,7 @@ async def get_user_daily_activity_aggregated(
end_date=end_date,
model=model,
api_key=api_key,
+ timezone_offset_minutes=timezone,
)
except Exception as e:
diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py
index 278971a91a5..9dadffca351 100644
--- a/litellm/proxy/management_endpoints/key_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/key_management_endpoints.py
@@ -37,13 +37,6 @@ from litellm.proxy._experimental.mcp_server.db import (
)
from litellm.proxy._types import *
from litellm.proxy._types import LiteLLM_VerificationToken
-from litellm.types.proxy.management_endpoints.key_management_endpoints import (
- BulkUpdateKeyRequest,
- BulkUpdateKeyRequestItem,
- BulkUpdateKeyResponse,
- FailedKeyUpdate,
- SuccessfulKeyUpdate,
-)
from litellm.proxy.auth.auth_checks import (
_cache_key_object,
_delete_cache_key_object,
@@ -82,6 +75,13 @@ from litellm.proxy.utils import (
)
from litellm.router import Router
from litellm.secret_managers.main import get_secret
+from litellm.types.proxy.management_endpoints.key_management_endpoints import (
+ BulkUpdateKeyRequest,
+ BulkUpdateKeyRequestItem,
+ BulkUpdateKeyResponse,
+ FailedKeyUpdate,
+ SuccessfulKeyUpdate,
+)
from litellm.types.router import Deployment
from litellm.types.utils import (
BudgetConfig,
@@ -2381,6 +2381,10 @@ async def info_key_fn(
# if using pydantic v1
key_info = key_info.dict()
key_info.pop("token")
+
+ # Attach object_permission if object_permission_id is set
+ key_info = await attach_object_permission_to_dict(key_info, prisma_client)
+
return {"key": key, "info": key_info}
except Exception as e:
raise handle_exception_on_proxy(e)
@@ -3373,6 +3377,163 @@ async def regenerate_key_fn(
raise handle_exception_on_proxy(e)
+async def _check_proxy_or_team_admin_for_key(
+ key_in_db: LiteLLM_VerificationToken,
+ user_api_key_dict: UserAPIKeyAuth,
+ prisma_client: PrismaClient,
+ user_api_key_cache: DualCache,
+) -> None:
+ if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
+ return
+
+ if key_in_db.team_id is not None:
+ team_table = await get_team_object(
+ team_id=key_in_db.team_id,
+ prisma_client=prisma_client,
+ user_api_key_cache=user_api_key_cache,
+ check_db_only=True,
+ )
+ if team_table is not None:
+ if _is_user_team_admin(
+ user_api_key_dict=user_api_key_dict,
+ team_obj=team_table,
+ ):
+ return
+
+ raise HTTPException(
+ status_code=status.HTTP_403_FORBIDDEN,
+ detail={"error": "You must be a proxy admin or team admin to reset key spend"},
+ )
+
+
+def _validate_reset_spend_value(
+ reset_to: Any, key_in_db: LiteLLM_VerificationToken
+) -> float:
+ if not isinstance(reset_to, (int, float)):
+ raise HTTPException(
+ status_code=status.HTTP_400_BAD_REQUEST,
+ detail={"error": "reset_to must be a float"},
+ )
+
+ reset_to = float(reset_to)
+
+ if reset_to < 0:
+ raise HTTPException(
+ status_code=status.HTTP_400_BAD_REQUEST,
+ detail={"error": "reset_to must be >= 0"},
+ )
+
+ current_spend = key_in_db.spend or 0.0
+ if reset_to > current_spend:
+ raise HTTPException(
+ status_code=status.HTTP_400_BAD_REQUEST,
+ detail={"error": f"reset_to ({reset_to}) must be <= current spend ({current_spend})"},
+ )
+
+ max_budget = key_in_db.max_budget
+ if key_in_db.litellm_budget_table is not None:
+ budget_max_budget = getattr(key_in_db.litellm_budget_table, "max_budget", None)
+ if budget_max_budget is not None:
+ if max_budget is None or budget_max_budget < max_budget:
+ max_budget = budget_max_budget
+
+ if max_budget is not None and reset_to > max_budget:
+ raise HTTPException(
+ status_code=status.HTTP_400_BAD_REQUEST,
+ detail={"error": f"reset_to ({reset_to}) must be <= budget ({max_budget})"},
+ )
+
+ return reset_to
+
+
+@router.post(
+ "/key/{key:path}/reset_spend",
+ tags=["key management"],
+ dependencies=[Depends(user_api_key_auth)],
+)
+@management_endpoint_wrapper
+async def reset_key_spend_fn(
+ key: str,
+ data: ResetSpendRequest,
+ user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
+ litellm_changed_by: Optional[str] = Header(
+ None,
+ description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
+ ),
+) -> Dict[str, Any]:
+ try:
+ from litellm.proxy.proxy_server import (
+ hash_token,
+ prisma_client,
+ proxy_logging_obj,
+ user_api_key_cache,
+ )
+
+ if prisma_client is None:
+ raise HTTPException(
+ status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
+ detail={"error": "DB not connected. prisma_client is None"},
+ )
+
+ if "sk" not in key:
+ hashed_api_key = key
+ else:
+ hashed_api_key = hash_token(key)
+
+ _key_in_db = await prisma_client.db.litellm_verificationtoken.find_unique(
+ where={"token": hashed_api_key},
+ include={"litellm_budget_table": True},
+ )
+ if _key_in_db is None:
+ raise HTTPException(
+ status_code=status.HTTP_404_NOT_FOUND,
+ detail={"error": f"Key {key} not found."},
+ )
+
+ current_spend = _key_in_db.spend or 0.0
+ reset_to = _validate_reset_spend_value(data.reset_to, _key_in_db)
+
+ await _check_proxy_or_team_admin_for_key(
+ key_in_db=_key_in_db,
+ user_api_key_dict=user_api_key_dict,
+ prisma_client=prisma_client,
+ user_api_key_cache=user_api_key_cache,
+ )
+
+ updated_key = await prisma_client.db.litellm_verificationtoken.update(
+ where={"token": hashed_api_key},
+ data={"spend": reset_to},
+ )
+
+ if updated_key is None:
+ raise HTTPException(
+ status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
+ detail={"error": "Failed to update key spend"},
+ )
+
+ await _delete_cache_key_object(
+ hashed_token=hashed_api_key,
+ user_api_key_cache=user_api_key_cache,
+ proxy_logging_obj=proxy_logging_obj,
+ )
+
+ max_budget = updated_key.max_budget
+ budget_reset_at = updated_key.budget_reset_at
+
+ return {
+ "key_hash": hashed_api_key,
+ "spend": reset_to,
+ "previous_spend": current_spend,
+ "max_budget": max_budget,
+ "budget_reset_at": budget_reset_at,
+ }
+ except HTTPException:
+ raise
+ except Exception as e:
+ verbose_proxy_logger.exception("Error resetting key spend: %s", e)
+ raise handle_exception_on_proxy(e)
+
+
async def validate_key_list_check(
user_api_key_dict: UserAPIKeyAuth,
user_id: Optional[str],
diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py
index 0965198bad9..e67e1eae745 100644
--- a/litellm/proxy/management_endpoints/scim/scim_v2.py
+++ b/litellm/proxy/management_endpoints/scim/scim_v2.py
@@ -410,6 +410,308 @@ async def set_scim_content_type(response: Response):
response.headers["Content-Type"] = "application/scim+json"
+def _get_resource_types(base_url: str = "/scim/v2") -> list:
+ """Return the list of SCIM ResourceType definitions per RFC 7643 Section 6."""
+ return [
+ SCIMResourceType(
+ id="User",
+ name="User",
+ description="User Account",
+ endpoint="/Users",
+ schema_="urn:ietf:params:scim:schemas:core:2.0:User",
+ meta={
+ "location": f"{base_url}/ResourceTypes/User",
+ "resourceType": "ResourceType",
+ },
+ ),
+ SCIMResourceType(
+ id="Group",
+ name="Group",
+ description="Group",
+ endpoint="/Groups",
+ schema_="urn:ietf:params:scim:schemas:core:2.0:Group",
+ meta={
+ "location": f"{base_url}/ResourceTypes/Group",
+ "resourceType": "ResourceType",
+ },
+ ),
+ ]
+
+
+def _get_schemas() -> list:
+ """Return the list of SCIM Schema definitions per RFC 7643 Section 7."""
+ return [
+ SCIMSchema(
+ id="urn:ietf:params:scim:schemas:core:2.0:User",
+ name="User",
+ description="User Account",
+ attributes=[
+ SCIMSchemaAttribute(
+ name="userName",
+ type="string",
+ multiValued=False,
+ description="Unique identifier for the User.",
+ required=True,
+ mutability="readWrite",
+ returned="default",
+ uniqueness="server",
+ ),
+ SCIMSchemaAttribute(
+ name="name",
+ type="complex",
+ multiValued=False,
+ description="The components of the user's real name.",
+ required=False,
+ subAttributes=[
+ SCIMSchemaAttribute(
+ name="givenName",
+ type="string",
+ description="The given name of the User.",
+ ),
+ SCIMSchemaAttribute(
+ name="familyName",
+ type="string",
+ description="The family name of the User.",
+ ),
+ SCIMSchemaAttribute(
+ name="formatted",
+ type="string",
+ description="The full name.",
+ ),
+ ],
+ ),
+ SCIMSchemaAttribute(
+ name="displayName",
+ type="string",
+ multiValued=False,
+ description="The name of the User, suitable for display.",
+ ),
+ SCIMSchemaAttribute(
+ name="emails",
+ type="complex",
+ multiValued=True,
+ description="Email addresses for the user.",
+ subAttributes=[
+ SCIMSchemaAttribute(
+ name="value",
+ type="string",
+ description="Email address value.",
+ ),
+ SCIMSchemaAttribute(
+ name="type",
+ type="string",
+ description="Type of email (work, home, etc.).",
+ ),
+ SCIMSchemaAttribute(
+ name="primary",
+ type="boolean",
+ description="Whether this is the primary email.",
+ ),
+ ],
+ ),
+ SCIMSchemaAttribute(
+ name="active",
+ type="boolean",
+ multiValued=False,
+ description="Whether the user account is active.",
+ ),
+ SCIMSchemaAttribute(
+ name="groups",
+ type="complex",
+ multiValued=True,
+ description="Groups to which the user belongs.",
+ mutability="readOnly",
+ subAttributes=[
+ SCIMSchemaAttribute(
+ name="value",
+ type="string",
+ description="Group identifier.",
+ ),
+ SCIMSchemaAttribute(
+ name="display",
+ type="string",
+ description="Group display name.",
+ ),
+ ],
+ ),
+ ],
+ meta={
+ "location": "/scim/v2/Schemas/urn:ietf:params:scim:schemas:core:2.0:User",
+ "resourceType": "Schema",
+ },
+ ),
+ SCIMSchema(
+ id="urn:ietf:params:scim:schemas:core:2.0:Group",
+ name="Group",
+ description="Group",
+ attributes=[
+ SCIMSchemaAttribute(
+ name="displayName",
+ type="string",
+ multiValued=False,
+ description="A human-readable name for the Group.",
+ required=True,
+ mutability="readWrite",
+ returned="default",
+ uniqueness="none",
+ ),
+ SCIMSchemaAttribute(
+ name="members",
+ type="complex",
+ multiValued=True,
+ description="A list of members of the Group.",
+ subAttributes=[
+ SCIMSchemaAttribute(
+ name="value",
+ type="string",
+ description="Member identifier.",
+ ),
+ SCIMSchemaAttribute(
+ name="display",
+ type="string",
+ description="Member display name.",
+ ),
+ ],
+ ),
+ ],
+ meta={
+ "location": "/scim/v2/Schemas/urn:ietf:params:scim:schemas:core:2.0:Group",
+ "resourceType": "Schema",
+ },
+ ),
+ ]
+
+
+@scim_router.get(
+ "",
+ status_code=200,
+ dependencies=[Depends(user_api_key_auth), Depends(set_scim_content_type)],
+)
+@scim_router.get(
+ "/",
+ status_code=200,
+ include_in_schema=False,
+ dependencies=[Depends(user_api_key_auth), Depends(set_scim_content_type)],
+)
+async def get_scim_base(request: Request):
+ """
+ Base SCIM v2 endpoint for resource discovery per RFC 7644 Section 4.
+
+ Returns a ListResponse of ResourceTypes supported by this SCIM service provider.
+ Identity providers (Okta, Azure AD, etc.) use this endpoint for resource discovery.
+ """
+ verbose_proxy_logger.debug(
+ "SCIM base resource discovery request: method=%s url=%s",
+ request.method,
+ request.url,
+ )
+ base_url = str(request.base_url).rstrip("/") + "/scim/v2"
+ resource_types = _get_resource_types(base_url)
+ return {
+ "schemas": ["urn:ietf:params:scim:api:messages:2.0:ListResponse"],
+ "totalResults": len(resource_types),
+ "Resources": [rt.model_dump() for rt in resource_types],
+ }
+
+
+@scim_router.get(
+ "/ResourceTypes",
+ status_code=200,
+ dependencies=[Depends(user_api_key_auth), Depends(set_scim_content_type)],
+)
+async def get_resource_types(request: Request):
+ """
+ SCIM ResourceTypes endpoint per RFC 7644 Section 4.
+
+ Returns a ListResponse of all resource types supported by this service provider.
+ """
+ verbose_proxy_logger.debug(
+ "SCIM ResourceTypes request: method=%s url=%s",
+ request.method,
+ request.url,
+ )
+ base_url = str(request.base_url).rstrip("/") + "/scim/v2"
+ resource_types = _get_resource_types(base_url)
+ return {
+ "schemas": ["urn:ietf:params:scim:api:messages:2.0:ListResponse"],
+ "totalResults": len(resource_types),
+ "Resources": [rt.model_dump() for rt in resource_types],
+ }
+
+
+@scim_router.get(
+ "/ResourceTypes/{resource_type_id}",
+ status_code=200,
+ dependencies=[Depends(user_api_key_auth), Depends(set_scim_content_type)],
+)
+async def get_resource_type(
+ request: Request,
+ resource_type_id: str = Path(..., title="ResourceType ID"),
+):
+ """
+ Get a single ResourceType by ID per RFC 7644.
+ """
+ verbose_proxy_logger.debug(
+ "SCIM ResourceType request for id=%s", resource_type_id
+ )
+ base_url = str(request.base_url).rstrip("/") + "/scim/v2"
+ resource_types = _get_resource_types(base_url)
+ for rt in resource_types:
+ if rt.id == resource_type_id:
+ return rt.model_dump()
+ raise HTTPException(
+ status_code=404,
+ detail={"error": f"ResourceType not found: {resource_type_id}"},
+ )
+
+
+@scim_router.get(
+ "/Schemas",
+ status_code=200,
+ dependencies=[Depends(user_api_key_auth), Depends(set_scim_content_type)],
+)
+async def get_schemas(request: Request):
+ """
+ SCIM Schemas endpoint per RFC 7643 Section 7.
+
+ Returns a ListResponse of all schemas supported by this service provider.
+ """
+ verbose_proxy_logger.debug(
+ "SCIM Schemas request: method=%s url=%s",
+ request.method,
+ request.url,
+ )
+ schemas = _get_schemas()
+ return {
+ "schemas": ["urn:ietf:params:scim:api:messages:2.0:ListResponse"],
+ "totalResults": len(schemas),
+ "Resources": [s.model_dump() for s in schemas],
+ }
+
+
+@scim_router.get(
+ "/Schemas/{schema_id:path}",
+ status_code=200,
+ dependencies=[Depends(user_api_key_auth), Depends(set_scim_content_type)],
+)
+async def get_schema(
+ request: Request,
+ schema_id: str = Path(..., title="Schema URI"),
+):
+ """
+ Get a single Schema by its URI per RFC 7643 Section 7.
+ """
+ verbose_proxy_logger.debug("SCIM Schema request for id=%s", schema_id)
+ schemas = _get_schemas()
+ for s in schemas:
+ if s.id == schema_id:
+ return s.model_dump()
+ raise HTTPException(
+ status_code=404,
+ detail={"error": f"Schema not found: {schema_id}"},
+ )
+
+
@scim_router.get(
"/ServiceProviderConfig",
response_model=SCIMServiceProviderConfig,
diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py
index 63db2d72fe4..8e903180ac5 100644
--- a/litellm/proxy/management_endpoints/team_endpoints.py
+++ b/litellm/proxy/management_endpoints/team_endpoints.py
@@ -685,6 +685,7 @@ async def new_team( # noqa: PLR0915
- rpm_limit_type: Optional[Literal["guaranteed_throughput", "best_effort_throughput"]] - The type of RPM limit enforcement. Use "guaranteed_throughput" to raise an error if overallocating RPM, or "best_effort_throughput" for best effort enforcement.
- tpm_limit_type: Optional[Literal["guaranteed_throughput", "best_effort_throughput"]] - The type of TPM limit enforcement. Use "guaranteed_throughput" to raise an error if overallocating TPM, or "best_effort_throughput" for best effort enforcement.
- max_budget: Optional[float] - The maximum budget allocated to the team - all keys for this team_id will have at max this max_budget
+ - soft_budget: Optional[float] - The soft budget threshold for the team. If max_budget is set, soft_budget must be strictly lower than max_budget. Can be set independently if max_budget is not set.
- budget_duration: Optional[str] - The duration of the budget for the team. Doc [here](https://docs.litellm.ai/docs/proxy/team_budgets)
- models: Optional[list] - A list of models associated with the team - all keys for this team_id will have at most, these models. If empty, assumes all models are allowed.
- blocked: bool - Flag indicating if the team is blocked or not - will stop all calls from keys with this team_id.
@@ -760,6 +761,22 @@ async def new_team( # noqa: PLR0915
status_code=400,
detail={"error": f"team_member_budget cannot be negative. Received: {data.team_member_budget}"}
)
+ if data.soft_budget is not None and data.soft_budget < 0:
+ raise HTTPException(
+ status_code=400,
+ detail={"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"}
+ )
+
+ if data.soft_budget is not None:
+ if data.max_budget is not None:
+ # If max_budget is set, soft_budget must be strictly lower than max_budget
+ if data.soft_budget >= data.max_budget:
+ raise HTTPException(
+ status_code=400,
+ detail={
+ "error": f"soft_budget ({data.soft_budget}) must be strictly lower than max_budget ({data.max_budget})"
+ }
+ )
# Check if license is over limit
total_teams = await prisma_client.db.litellm_teamtable.count()
@@ -1226,6 +1243,7 @@ async def update_team( # noqa: PLR0915
- tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for this team - all keys with this team_id will have at max this TPM limit
- rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for this team - all keys associated with this team_id will have at max this RPM limit
- max_budget: Optional[float] - The maximum budget allocated to the team - all keys for this team_id will have at max this max_budget
+ - soft_budget: Optional[float] - The soft budget threshold for the team. If max_budget is set (either in the request or existing), soft_budget must be strictly lower than max_budget. Can be set independently if max_budget is not set.
- budget_duration: Optional[str] - The duration of the budget for the team. Doc [here](https://docs.litellm.ai/docs/proxy/team_budgets)
- models: Optional[list] - A list of models associated with the team - all keys for this team_id will have at most, these models. If empty, assumes all models are allowed.
- prompts: Optional[List[str]] - List of prompts that the team is allowed to use.
@@ -1302,6 +1320,11 @@ async def update_team( # noqa: PLR0915
status_code=400,
detail={"error": f"team_member_budget cannot be negative. Received: {data.team_member_budget}"}
)
+ if data.soft_budget is not None and data.soft_budget < 0:
+ raise HTTPException(
+ status_code=400,
+ detail={"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"}
+ )
existing_team_row = await prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": data.team_id}
@@ -1312,6 +1335,29 @@ async def update_team( # noqa: PLR0915
status_code=404,
detail={"error": f"Team not found, passed team_id={data.team_id}"},
)
+
+ if data.soft_budget is not None:
+ max_budget_to_check = data.max_budget if data.max_budget is not None else existing_team_row.max_budget
+ if max_budget_to_check is not None:
+ if data.soft_budget >= max_budget_to_check:
+ raise HTTPException(
+ status_code=400,
+ detail={
+ "error": f"soft_budget ({data.soft_budget}) must be strictly lower than max_budget ({max_budget_to_check})"
+ }
+ )
+
+ if data.max_budget is not None:
+ existing_soft_budget = getattr(existing_team_row, 'soft_budget', None)
+ soft_budget_to_check = data.soft_budget if data.soft_budget is not None else existing_soft_budget
+ if soft_budget_to_check is not None and isinstance(soft_budget_to_check, (int, float)):
+ if data.max_budget <= soft_budget_to_check:
+ raise HTTPException(
+ status_code=400,
+ detail={
+ "error": f"max_budget ({data.max_budget}) must be strictly greater than soft_budget ({soft_budget_to_check})"
+ }
+ )
if (
data.organization_id is not None and len(data.organization_id) > 0
diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py
index 4048b3731c1..2d248dc81f3 100644
--- a/litellm/proxy/management_endpoints/ui_sso.py
+++ b/litellm/proxy/management_endpoints/ui_sso.py
@@ -326,6 +326,7 @@ def generic_response_convertor(
jwt_handler: JWTHandler,
sso_jwt_handler: Optional[JWTHandler] = None,
role_mappings: Optional["RoleMappings"] = None,
+ team_mappings: Optional["TeamMappings"] = None,
) -> CustomOpenID:
generic_user_id_attribute_name = os.getenv(
"GENERIC_USER_ID_ATTRIBUTE", "preferred_username"
@@ -359,8 +360,20 @@ def generic_response_convertor(
team_ids = sso_jwt_handler.get_team_ids_from_jwt(cast(dict, response))
all_teams.extend(team_ids)
- team_ids = jwt_handler.get_team_ids_from_jwt(cast(dict, response))
- all_teams.extend(team_ids)
+ if team_mappings is not None and team_mappings.team_ids_jwt_field is not None:
+ team_ids_from_db_mapping: Optional[List[str]] = get_nested_value(
+ data=cast(dict, response),
+ key_path=team_mappings.team_ids_jwt_field,
+ default=[],
+ )
+ if team_ids_from_db_mapping:
+ all_teams.extend(team_ids_from_db_mapping)
+ verbose_proxy_logger.debug(
+ f"Loaded team_ids from DB team_mappings.team_ids_jwt_field='{team_mappings.team_ids_jwt_field}': {team_ids_from_db_mapping}"
+ )
+ else:
+ team_ids = jwt_handler.get_team_ids_from_jwt(cast(dict, response))
+ all_teams.extend(team_ids)
# Determine user role based on role_mappings if available
# Only apply role_mappings for GENERIC SSO provider
@@ -484,6 +497,43 @@ def _setup_generic_sso_env_vars(
)
+async def _setup_team_mappings() -> Optional["TeamMappings"]:
+ """Setup team mappings from SSO database settings."""
+ team_mappings: Optional["TeamMappings"] = None
+ try:
+ from litellm.proxy.utils import get_prisma_client_or_throw
+
+ prisma_client = get_prisma_client_or_throw(
+ "Prisma client is None, connect a database to your proxy"
+ )
+
+ sso_db_record = await prisma_client.db.litellm_ssoconfig.find_unique(
+ where={"id": "sso_config"}
+ )
+
+ if sso_db_record and sso_db_record.sso_settings:
+ sso_settings_dict = dict(sso_db_record.sso_settings)
+ team_mappings_data = sso_settings_dict.get("team_mappings")
+
+ if team_mappings_data:
+ from litellm.types.proxy.management_endpoints.ui_sso import TeamMappings
+ if isinstance(team_mappings_data, dict):
+ team_mappings = TeamMappings(**team_mappings_data)
+ elif isinstance(team_mappings_data, TeamMappings):
+ team_mappings = team_mappings_data
+
+ if team_mappings and team_mappings.team_ids_jwt_field:
+ verbose_proxy_logger.debug(
+ f"Loaded team_mappings with team_ids_jwt_field: '{team_mappings.team_ids_jwt_field}'"
+ )
+ except Exception as e:
+ verbose_proxy_logger.debug(
+ f"Could not load team_mappings from database: {e}. Continuing with config-based team mapping."
+ )
+
+ return team_mappings
+
+
async def _setup_role_mappings() -> Optional["RoleMappings"]:
"""Setup role mappings from SSO database settings."""
role_mappings: Optional["RoleMappings"] = None
@@ -494,7 +544,6 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]:
"Prisma client is None, connect a database to your proxy"
)
- # Get SSO config from dedicated table
sso_db_record = await prisma_client.db.litellm_ssoconfig.find_unique(
where={"id": "sso_config"}
)
@@ -515,7 +564,6 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]:
f"Loaded role_mappings for provider '{role_mappings.provider}'"
)
except Exception as e:
- # If we can't load role_mappings, continue with existing logic
verbose_proxy_logger.debug(
f"Could not load role_mappings from database: {e}. Continuing with existing role logic."
)
@@ -590,8 +638,8 @@ async def get_generic_sso_response(
userinfo_endpoint=generic_userinfo_endpoint,
)
- # Get role_mappings from SSO settings if available
role_mappings = await _setup_role_mappings()
+ team_mappings = await _setup_team_mappings()
def response_convertor(response, client):
nonlocal received_response # return for user debugging
@@ -601,6 +649,7 @@ async def get_generic_sso_response(
jwt_handler=jwt_handler,
sso_jwt_handler=sso_jwt_handler,
role_mappings=role_mappings,
+ team_mappings=team_mappings,
)
SSOProvider = create_provider(
diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py
index 2ff1183579f..f67dc5e2aaa 100644
--- a/litellm/proxy/openai_files_endpoints/common_utils.py
+++ b/litellm/proxy/openai_files_endpoints/common_utils.py
@@ -637,3 +637,127 @@ def _extract_model_param(request: "Request", request_body: dict) -> Optional[str
or request.query_params.get("model")
or request.headers.get("x-litellm-model")
)
+
+
+# ============================================================================
+# BATCH DATABASE OPERATIONS
+# ============================================================================
+
+
+async def get_batch_from_database(
+ batch_id: str,
+ unified_batch_id: Union[str, Literal[False]],
+ managed_files_obj,
+ prisma_client,
+ verbose_proxy_logger,
+):
+ """
+ Try to retrieve batch object from ManagedObjectTable for consistent state.
+
+ Args:
+ batch_id: The batch ID (may be unified/encoded)
+ unified_batch_id: Result from _is_base64_encoded_unified_file_id()
+ managed_files_obj: The managed_files proxy hook object
+ prisma_client: Prisma database client
+ verbose_proxy_logger: Logger instance
+
+ Returns:
+ Tuple of (db_batch_object, response_batch)
+ - db_batch_object: Raw database object (or None)
+ - response_batch: Parsed LiteLLMBatch object (or None)
+ """
+ import json
+ from litellm.types.utils import LiteLLMBatch
+
+ if managed_files_obj is None or not unified_batch_id:
+ return None, None
+
+ try:
+ if not prisma_client:
+ return None, None
+
+ db_batch_object = await prisma_client.db.litellm_managedobjecttable.find_first(
+ where={"unified_object_id": batch_id}
+ )
+
+ if not db_batch_object or not db_batch_object.file_object:
+ return None, None
+
+ # Parse the batch object from database
+ batch_data = json.loads(db_batch_object.file_object) if isinstance(db_batch_object.file_object, str) else db_batch_object.file_object
+ response = LiteLLMBatch(**batch_data)
+ response.id = batch_id
+
+ verbose_proxy_logger.debug(
+ f"Retrieved batch {batch_id} from ManagedObjectTable with status={response.status}"
+ )
+
+ return db_batch_object, response
+
+ except Exception as e:
+ verbose_proxy_logger.warning(
+ f"Failed to retrieve batch from ManagedObjectTable: {e}, falling back to provider"
+ )
+ return None, None
+
+
+async def update_batch_in_database(
+ batch_id: str,
+ unified_batch_id: Union[str, Literal[False]],
+ response,
+ managed_files_obj,
+ prisma_client,
+ verbose_proxy_logger,
+ db_batch_object=None,
+ operation: str = "update",
+):
+ """
+ Update batch status and object in ManagedObjectTable.
+
+ Args:
+ batch_id: The batch ID (unified/encoded)
+ unified_batch_id: Result from _is_base64_encoded_unified_file_id()
+ response: The batch response object with updated state
+ managed_files_obj: The managed_files proxy hook object
+ prisma_client: Prisma database client
+ verbose_proxy_logger: Logger instance
+ db_batch_object: Optional existing database object (for comparison)
+ operation: Description of operation ("update", "cancel", etc.)
+ """
+ import litellm.utils
+
+ if managed_files_obj is None or not unified_batch_id:
+ return
+
+ try:
+ if not prisma_client:
+ return
+
+ # Only update if status has changed (when db_batch_object is provided)
+ if db_batch_object and response.status == db_batch_object.status:
+ return
+
+ if db_batch_object:
+ verbose_proxy_logger.info(
+ f"Updating batch {batch_id} status from {db_batch_object.status} to {response.status}"
+ )
+ else:
+ verbose_proxy_logger.info(
+ f"Updating batch {batch_id} status to {response.status} after {operation}"
+ )
+
+ # Normalize status for database storage
+ db_status = response.status if response.status != "completed" else "complete"
+
+ await prisma_client.db.litellm_managedobjecttable.update(
+ where={"unified_object_id": batch_id},
+ data={
+ "status": db_status,
+ "file_object": response.model_dump_json(),
+ "updated_at": litellm.utils.get_utc_datetime(),
+ },
+ )
+ except Exception as e:
+ verbose_proxy_logger.error(
+ f"Failed to update batch status in ManagedObjectTable: {e}"
+ )
diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py
index da267eac981..ec6e9733344 100644
--- a/litellm/proxy/openai_files_endpoints/files_endpoints.py
+++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py
@@ -1101,6 +1101,7 @@ async def delete_file(
**data_without_file_id,
)
else:
+ data.pop("file_id", None)
response = await litellm.afile_delete(
custom_llm_provider=custom_llm_provider, file_id=file_id, **data # type: ignore
)
diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml
index e12e75b54ff..d87ae8b14ca 100644
--- a/litellm/proxy/proxy_config.yaml
+++ b/litellm/proxy/proxy_config.yaml
@@ -1,4 +1,14 @@
model_list:
+ - model_name: gpt-4o
+ litellm_params:
+ model: openai/gpt-4o
+ api_key: os.environ/OPENAI_API_KEY
+
+ - model_name: text-embedding-3-small
+ litellm_params:
+ model: openai/text-embedding-3-small
+ api_key: os.environ/OPENAI_API_KEY
+
- model_name: bedrock-claude-sonnet-3.5
litellm_params:
model: "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0"
@@ -22,4 +32,31 @@ model_list:
- model_name: bedrock-nova-premier
litellm_params:
model: "bedrock/us.amazon.nova-premier-v1:0"
- aws_region_name: "us-east-1"
\ No newline at end of file
+ aws_region_name: "us-east-1"
+
+# MCP Server Configuration
+mcp_servers:
+ # Wikipedia MCP - reliable and works without external deps
+ wikipedia:
+ transport: "stdio"
+ command: "uvx"
+ args: ["mcp-server-fetch"]
+ description: "Fetch web pages and Wikipedia content"
+ deepwiki:
+ transport: "http"
+ url: "https://mcp.deepwiki.com/mcp"
+
+# General Settings
+general_settings:
+ master_key: sk-1234
+ store_model_in_db: false
+
+# LiteLLM Settings
+litellm_settings:
+ # Enable MCP Semantic Tool Filter
+ mcp_semantic_tool_filter:
+ enabled: true
+ embedding_model: "text-embedding-3-small"
+ top_k: 5
+ similarity_threshold: 0.3
+
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index 053301d8ced..1a4db938474 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -47,6 +47,7 @@ from litellm.constants import (
DEFAULT_SLACK_ALERTING_THRESHOLD,
LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS,
LITELLM_SETTINGS_SAFE_DB_OVERRIDES,
+ LITELLM_UI_ALLOW_HEADERS,
)
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
@@ -239,6 +240,10 @@ from litellm.proxy._types import *
from litellm.proxy.agent_endpoints.a2a_endpoints import router as a2a_router
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.agent_endpoints.endpoints import router as agent_endpoints_router
+from litellm.proxy.agent_endpoints.model_list_helpers import (
+ append_agents_to_model_group,
+ append_agents_to_model_info,
+)
from litellm.proxy.analytics_endpoints.analytics_endpoints import (
router as analytics_router,
)
@@ -793,6 +798,21 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
redis_usage_cache=redis_usage_cache,
)
+ ## SEMANTIC TOOL FILTER ##
+ # Read litellm_settings from config for semantic filter initialization
+ try:
+ verbose_proxy_logger.debug("About to initialize semantic tool filter")
+ _config = proxy_config.get_config_state()
+ _litellm_settings = _config.get("litellm_settings", {})
+ verbose_proxy_logger.debug(f"litellm_settings keys = {list(_litellm_settings.keys())}")
+ await ProxyStartupEvent._initialize_semantic_tool_filter(
+ llm_router=llm_router,
+ litellm_settings=_litellm_settings,
+ )
+ verbose_proxy_logger.debug("After semantic tool filter initialization")
+ except Exception as e:
+ verbose_proxy_logger.error(f"Semantic filter init failed: {e}", exc_info=True)
+
## JWT AUTH ##
ProxyStartupEvent._initialize_jwt_auth(
general_settings=general_settings,
@@ -1195,6 +1215,7 @@ app.add_middleware(
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
+ expose_headers=LITELLM_UI_ALLOW_HEADERS,
)
app.add_middleware(PrometheusAuthMiddleware)
@@ -1843,6 +1864,7 @@ class ProxyConfig:
def __init__(self) -> None:
self.config: Dict[str, Any] = {}
+ self._last_semantic_filter_config: Optional[Dict[str, Any]] = None
def is_yaml(self, config_file_path: str) -> bool:
if not os.path.isfile(config_file_path):
@@ -3381,105 +3403,6 @@ class ProxyConfig:
decrypted_variables[k] = decrypted_value
return decrypted_variables
- async def _get_hierarchical_router_settings(
- self,
- user_api_key_dict: Optional["UserAPIKeyAuth"],
- prisma_client: Optional[PrismaClient],
- ) -> Optional[dict]:
- """
- Get router_settings in priority order: Key > Team > Global
-
- Returns:
- dict: Combined router_settings, or None if no settings found
- """
- if prisma_client is None:
- return None
-
- import json
-
- import yaml
-
- # 1. Try key-level router_settings
- if user_api_key_dict is not None:
- # Check if router_settings is available on the key object
- key_router_settings_value = getattr(
- user_api_key_dict, "router_settings", None
- )
- if key_router_settings_value is not None:
- key_router_settings = None
- if isinstance(key_router_settings_value, str):
- try:
- key_router_settings = yaml.safe_load(key_router_settings_value)
- except (yaml.YAMLError, json.JSONDecodeError):
- try:
- key_router_settings = json.loads(key_router_settings_value)
- except json.JSONDecodeError:
- pass
- elif isinstance(key_router_settings_value, dict):
- key_router_settings = key_router_settings_value
-
- # If key has router_settings (non-empty dict), use it
- if (
- key_router_settings is not None
- and isinstance(key_router_settings, dict)
- and key_router_settings
- ):
- return key_router_settings
-
- # 2. Try team-level router_settings
- if user_api_key_dict is not None and user_api_key_dict.team_id is not None:
- try:
- team_obj = await prisma_client.db.litellm_teamtable.find_unique(
- where={"team_id": user_api_key_dict.team_id}
- )
- if team_obj is not None:
- team_router_settings_value = getattr(
- team_obj, "router_settings", None
- )
- if team_router_settings_value is not None:
- team_router_settings = None
- if isinstance(team_router_settings_value, str):
- try:
- team_router_settings = yaml.safe_load(
- team_router_settings_value
- )
- except (yaml.YAMLError, json.JSONDecodeError):
- try:
- team_router_settings = json.loads(
- team_router_settings_value
- )
- except json.JSONDecodeError:
- pass
- elif isinstance(team_router_settings_value, dict):
- team_router_settings = team_router_settings_value
-
- # If team has router_settings (non-empty dict), use it
- if (
- team_router_settings is not None
- and isinstance(team_router_settings, dict)
- and team_router_settings
- ):
- return team_router_settings
- except Exception:
- # If team lookup fails, continue to global settings
- pass
-
- # 3. Try global router_settings
- try:
- db_router_settings = await prisma_client.db.litellm_config.find_first(
- where={"param_name": "router_settings"}
- )
- if (
- db_router_settings is not None
- and isinstance(db_router_settings.param_value, dict)
- and db_router_settings.param_value
- ):
- return db_router_settings.param_value
- except Exception:
- pass
-
- return None
-
async def _add_router_settings_from_db_config(
self,
config_data: dict,
@@ -4001,6 +3924,93 @@ class ProxyConfig:
prisma_client=prisma_client, proxy_config=self
)
+ if self._should_load_db_object(object_type="semantic_filter_settings"):
+ await self._init_semantic_filter_settings_in_db(
+ prisma_client=prisma_client
+ )
+
+ async def _init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient):
+ """
+ Initialize MCP semantic filter settings from database.
+ Called periodically (approximately every 10 seconds) by background task to hot-reload settings across all pods.
+ """
+ import json
+
+ import litellm
+ from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
+
+ try:
+ # Load litellm_settings from DB
+ config_record = await prisma_client.db.litellm_config.find_unique(
+ where={"param_name": "litellm_settings"}
+ )
+
+ if config_record is None or config_record.param_value is None:
+ return
+
+ litellm_settings = config_record.param_value
+ if isinstance(litellm_settings, str):
+ litellm_settings = json.loads(litellm_settings)
+
+ mcp_semantic_filter_config = litellm_settings.get(
+ "mcp_semantic_tool_filter", None
+ )
+
+ if mcp_semantic_filter_config is None:
+ return
+
+ # Check if settings have changed (compare with in-memory state)
+ if hasattr(self, "_last_semantic_filter_config"):
+ if self._last_semantic_filter_config == mcp_semantic_filter_config:
+ # If hook is missing or router isn't built yet, reinitialize anyway
+ active_hooks = (
+ litellm.logging_callback_manager.get_custom_loggers_for_type(
+ SemanticToolFilterHook
+ )
+ )
+ if active_hooks:
+ for active_hook in active_hooks:
+ if isinstance(active_hook, SemanticToolFilterHook):
+ if (
+ active_hook.filter is not None
+ and active_hook.filter.tool_router is not None
+ ):
+ verbose_proxy_logger.debug(
+ "Semantic filter settings unchanged, skipping reinitialization"
+ )
+ return
+ verbose_proxy_logger.info(
+ "Semantic filter settings unchanged, but hook is missing or uninitialized. Reinitializing."
+ )
+
+ # Remove old hooks using logging callback manager
+ litellm.logging_callback_manager.remove_callbacks_by_type(
+ litellm.callbacks, SemanticToolFilterHook
+ )
+
+ # Initialize new hook if enabled
+ if mcp_semantic_filter_config.get("enabled", False):
+ global llm_router
+ hook = await SemanticToolFilterHook.initialize_from_config(
+ config=mcp_semantic_filter_config,
+ llm_router=llm_router,
+ )
+ if hook:
+ litellm.logging_callback_manager.add_litellm_callback(hook)
+ verbose_proxy_logger.info(
+ "MCP Semantic Filter reinitialized from DB"
+ )
+ else:
+ verbose_proxy_logger.info("MCP Semantic Filter disabled")
+
+ # Store current config for comparison next time
+ self._last_semantic_filter_config = mcp_semantic_filter_config.copy()
+
+ except Exception as e:
+ verbose_proxy_logger.exception(
+ f"Error initializing semantic filter settings from DB: {e}"
+ )
+
async def _init_sso_settings_in_db(self, prisma_client: PrismaClient):
"""
Initialize SSO settings from database into the router on startup.
@@ -4011,8 +4021,9 @@ class ProxyConfig:
where={"id": "sso_config"}
)
if sso_settings is not None:
- # Capitalize all keys in sso_settings dictionary
sso_settings.sso_settings.pop("role_mappings", None)
+ sso_settings.sso_settings.pop("team_mappings", None)
+ sso_settings.sso_settings.pop("ui_access_mode", None)
uppercase_sso_settings = {
key.upper(): value
for key, value in sso_settings.sso_settings.items()
@@ -4841,6 +4852,34 @@ class ProxyStartupEvent:
llm_router=llm_router, redis_usage_cache=redis_usage_cache
)
+ @classmethod
+ async def _initialize_semantic_tool_filter(
+ cls,
+ llm_router: Optional[Router],
+ litellm_settings: Dict[str, Any],
+ ):
+ """Initialize MCP semantic tool filter if configured"""
+ from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
+
+ verbose_proxy_logger.info(
+ f"Initializing semantic tool filter: llm_router={llm_router is not None}, "
+ f"litellm_settings keys={list(litellm_settings.keys())}"
+ )
+
+ mcp_semantic_filter_config = litellm_settings.get("mcp_semantic_tool_filter", None)
+ verbose_proxy_logger.debug(f"Semantic filter config: {mcp_semantic_filter_config}")
+
+ hook = await SemanticToolFilterHook.initialize_from_config(
+ config=mcp_semantic_filter_config,
+ llm_router=llm_router,
+ )
+
+ if hook:
+ verbose_proxy_logger.debug("✅ Semantic tool filter hook registered")
+ litellm.logging_callback_manager.add_litellm_callback(hook)
+ else:
+ verbose_proxy_logger.warning("❌ Semantic tool filter hook not initialized")
+
@classmethod
def _initialize_jwt_auth(
cls,
@@ -8672,6 +8711,15 @@ async def model_info_v2(
)
verbose_proxy_logger.debug("all_models: %s", all_models)
+
+ # Append A2A agents to models list
+ all_models = await append_agents_to_model_info(
+ models=all_models,
+ user_api_key_dict=user_api_key_dict,
+ )
+
+ # Update total count to include agents
+ search_total_count = len(all_models)
return _paginate_models_response(
all_models=all_models,
@@ -9512,6 +9560,12 @@ async def model_group_info(
model_groups: List[ModelGroupInfoProxy] = _get_model_group_info(
llm_router=llm_router, all_models_str=all_models_str, model_group=model_group
)
+
+ # Append A2A agents to model groups
+ model_groups = await append_agents_to_model_group(
+ model_groups=model_groups,
+ user_api_key_dict=user_api_key_dict,
+ )
return {"data": model_groups}
diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py
index e2749eb8187..e941964644e 100644
--- a/litellm/proxy/route_llm_request.py
+++ b/litellm/proxy/route_llm_request.py
@@ -12,6 +12,11 @@ else:
LitellmRouter = Any
+def _is_a2a_agent_model(model_name: Any) -> bool:
+ """Check if the model name is for an A2A agent (a2a/ prefix)."""
+ return isinstance(model_name, str) and model_name.startswith("a2a/")
+
+
ROUTE_ENDPOINT_MAPPING = {
"acompletion": "/chat/completions",
"atext_completion": "/completions",
@@ -92,10 +97,18 @@ def add_shared_session_to_data(data: dict) -> None:
data: Dictionary to add the shared session to
"""
try:
+ from litellm._logging import verbose_proxy_logger
from litellm.proxy.proxy_server import shared_aiohttp_session
if shared_aiohttp_session is not None and not shared_aiohttp_session.closed:
data["shared_session"] = shared_aiohttp_session
+ verbose_proxy_logger.info(
+ f"SESSION REUSE: Attached shared aiohttp session to request (ID: {id(shared_aiohttp_session)})"
+ )
+ else:
+ verbose_proxy_logger.info(
+ "SESSION REUSE: No shared session available for this request"
+ )
except Exception:
# Silently continue without session reuse if import fails or session unavailable
pass
@@ -337,6 +350,15 @@ async def route_request(
except Exception:
# If router fails (e.g., model not found in router), fall back to direct call
return getattr(litellm, f"{route_type}")(**data)
+ elif _is_a2a_agent_model(data.get("model", "")):
+ from litellm.proxy.agent_endpoints.a2a_routing import (
+ route_a2a_agent_request,
+ )
+
+ result = route_a2a_agent_request(data, route_type)
+ if result is not None:
+ return result
+ # Fall through to raise exception below if result is None
elif user_model is not None:
return getattr(litellm, f"{route_type}")(**data)
diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma
index b118400b620..a6a573836b5 100644
--- a/litellm/proxy/schema.prisma
+++ b/litellm/proxy/schema.prisma
@@ -113,6 +113,7 @@ model LiteLLM_TeamTable {
members_with_roles Json @default("{}")
metadata Json @default("{}")
max_budget Float?
+ soft_budget Float?
spend Float @default(0.0)
models String[]
max_parallel_requests Int?
diff --git a/litellm/proxy/search_endpoints/search_tool_management.py b/litellm/proxy/search_endpoints/search_tool_management.py
index 21a4db75a01..4754316795e 100644
--- a/litellm/proxy/search_endpoints/search_tool_management.py
+++ b/litellm/proxy/search_endpoints/search_tool_management.py
@@ -48,7 +48,7 @@ def _convert_datetime_to_str(value: Union[datetime, str, None]) -> Union[str, No
)
async def list_search_tools():
"""
- List all search tools that are available in the database.
+ List all search tools that are available in the database and config file.
Example Request:
```bash
@@ -71,38 +71,100 @@ async def list_search_tools():
"description": "Perplexity search tool"
},
"created_at": "2023-11-09T12:34:56.789Z",
- "updated_at": "2023-11-09T12:34:56.789Z"
+ "updated_at": "2023-11-09T12:34:56.789Z",
+ "is_from_config": false
+ },
+ {
+ "search_tool_name": "config-search-tool",
+ "litellm_params": {
+ "search_provider": "tavily",
+ "api_key": "tvly-***"
+ },
+ "is_from_config": true
}
]
}
```
"""
- from litellm.proxy.proxy_server import prisma_client
+ from litellm.litellm_core_utils.litellm_logging import _get_masked_values
+ from litellm.proxy.proxy_server import prisma_client, proxy_config
if prisma_client is None:
raise HTTPException(status_code=500, detail="Prisma client not initialized")
try:
- search_tools = await SEARCH_TOOL_REGISTRY.get_all_search_tools_from_db(
+ search_tools_from_db = await SEARCH_TOOL_REGISTRY.get_all_search_tools_from_db(
prisma_client=prisma_client
)
+ db_tool_names = {
+ tool.get("search_tool_name") for tool in search_tools_from_db
+ }
+
search_tool_configs: List[SearchToolInfoResponse] = []
- for search_tool in search_tools:
+
+ config_search_tools = []
+
+ try:
+ config = await proxy_config.get_config()
+ parsed_tools = proxy_config.parse_search_tools(config)
+ if parsed_tools:
+ config_search_tools = parsed_tools
+ except Exception as e:
+ verbose_proxy_logger.debug(
+ f"Could not get config-defined search tools: {e}"
+ )
+
+ for search_tool in config_search_tools:
+ tool_name = search_tool.get("search_tool_name")
+ if tool_name:
+ litellm_params_dict = dict(search_tool.get("litellm_params", {}))
+ masked_litellm_params_dict = _get_masked_values(
+ litellm_params_dict,
+ unmasked_length=4,
+ number_of_asterisks=4,
+ )
+
+ search_tool_configs.append(
+ SearchToolInfoResponse(
+ search_tool_id=None,
+ search_tool_name=tool_name,
+ litellm_params=masked_litellm_params_dict,
+ search_tool_info=search_tool.get("search_tool_info"),
+ created_at=None,
+ updated_at=None,
+ is_from_config=True,
+ )
+ )
+
+ search_tool_configs = [
+ tool for tool in search_tool_configs
+ if tool.get("search_tool_name") not in db_tool_names
+ ]
+
+ for search_tool in search_tools_from_db:
+ litellm_params_dict = dict(search_tool.get("litellm_params", {}))
+ masked_litellm_params_dict = _get_masked_values(
+ litellm_params_dict,
+ unmasked_length=4,
+ number_of_asterisks=4,
+ )
+
search_tool_configs.append(
SearchToolInfoResponse(
search_tool_id=search_tool.get("search_tool_id"),
search_tool_name=search_tool.get("search_tool_name", ""),
- litellm_params=dict(search_tool.get("litellm_params", {})),
+ litellm_params=masked_litellm_params_dict,
search_tool_info=search_tool.get("search_tool_info"),
created_at=_convert_datetime_to_str(search_tool.get("created_at")),
updated_at=_convert_datetime_to_str(search_tool.get("updated_at")),
+ is_from_config=False,
)
)
return ListSearchToolsResponse(search_tools=search_tool_configs)
except Exception as e:
- verbose_proxy_logger.exception(f"Error getting search tools from db: {e}")
+ verbose_proxy_logger.exception(f"Error getting search tools: {e}")
raise HTTPException(status_code=500, detail=str(e))
@@ -382,6 +444,7 @@ async def get_search_tool_info(search_tool_id: str):
search_tool_info=result.get("search_tool_info"),
created_at=_convert_datetime_to_str(result.get("created_at")),
updated_at=_convert_datetime_to_str(result.get("updated_at")),
+ is_from_config=False, # This endpoint only returns DB tools
)
except HTTPException as e:
raise e
diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py
index 4a0268eeede..6d32310940b 100644
--- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py
+++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py
@@ -83,6 +83,11 @@ class UISettings(BaseModel):
description="List of page keys that internal users (non-admins) can see in the UI sidebar. If not set, all pages are visible based on role permissions.",
)
+ require_auth_for_public_ai_hub: bool = Field(
+ default=False,
+ description="If true, requires authentication for accessing the public AI Hub."
+ )
+
class UISettingsResponse(SettingsResponse):
"""Response model for UI settings"""
@@ -95,9 +100,44 @@ ALLOWED_UI_SETTINGS_FIELDS = {
"disable_model_add_for_internal_users",
"disable_team_admin_delete_team_user",
"enabled_ui_pages_internal_users",
+ "require_auth_for_public_ai_hub",
}
+class MCPSemanticFilterSettings(BaseModel):
+ """Configuration for MCP Semantic Tool Filter"""
+
+ enabled: bool = Field(
+ default=False,
+ description="Enable semantic filtering of MCP tools based on query relevance",
+ )
+
+ embedding_model: str = Field(
+ default="text-embedding-3-small",
+ description="Embedding model to use for semantic similarity (e.g., 'text-embedding-3-small', 'text-embedding-ada-002')",
+ )
+
+ top_k: int = Field(
+ default=10,
+ description="Number of most relevant tools to return",
+ ge=1,
+ le=100,
+ )
+
+ similarity_threshold: float = Field(
+ default=0.3,
+ description="Minimum similarity score for tool inclusion (0.0 to 1.0, where 1.0 = exact match)",
+ ge=0.0,
+ le=1.0,
+ )
+
+
+class MCPSemanticFilterSettingsResponse(SettingsResponse):
+ """Response model for MCP semantic filter settings"""
+
+ pass
+
+
@router.get(
"/get/allowed_ips",
tags=["Budget & Spend Tracking"],
@@ -325,7 +365,7 @@ async def update_default_team_member_budget(
async def _update_litellm_setting(
- settings: Union[DefaultInternalUserParams, DefaultTeamSSOParams],
+ settings: Union[DefaultInternalUserParams, DefaultTeamSSOParams, MCPSemanticFilterSettings],
settings_key: str,
in_memory_var: Any,
success_message: str,
@@ -449,7 +489,6 @@ async def get_sso_settings():
# Load settings from database
sso_settings_dict = dict(sso_db_record.sso_settings)
- # Extract role_mappings before removing it (it's a dict, not an env variable)
role_mappings_data = sso_settings_dict.pop("role_mappings", None)
role_mappings = None
if role_mappings_data:
@@ -460,6 +499,16 @@ async def get_sso_settings():
elif isinstance(role_mappings_data, RoleMappings):
role_mappings = role_mappings_data
+ team_mappings_data = sso_settings_dict.pop("team_mappings", None)
+ team_mappings = None
+ if team_mappings_data:
+ from litellm.types.proxy.management_endpoints.ui_sso import TeamMappings
+
+ if isinstance(team_mappings_data, dict):
+ team_mappings = TeamMappings(**team_mappings_data)
+ elif isinstance(team_mappings_data, TeamMappings):
+ team_mappings = team_mappings_data
+
decrypted_sso_settings_dict = proxy_config._decrypt_and_set_db_env_variables(
environment_variables=sso_settings_dict
)
@@ -495,6 +544,7 @@ async def get_sso_settings():
user_email=decrypted_sso_settings_dict.get("user_email"),
ui_access_mode=decrypted_sso_settings_dict.get("ui_access_mode"),
role_mappings=role_mappings,
+ team_mappings=team_mappings,
)
# Get the schema for UI display
@@ -759,6 +809,70 @@ async def update_ui_theme_settings(theme_config: UIThemeConfig):
}
+@router.get(
+ "/get/mcp_semantic_filter_settings",
+ tags=["Settings"],
+ dependencies=[Depends(user_api_key_auth)],
+ response_model=MCPSemanticFilterSettingsResponse,
+)
+async def get_mcp_semantic_filter_settings(
+ user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
+):
+ """
+ Get MCP semantic filter configuration.
+ Returns current settings for semantic tool filtering.
+ """
+ from litellm.proxy.proxy_server import prisma_client, proxy_config
+
+ if prisma_client is None:
+ raise HTTPException(
+ status_code=500,
+ detail={"error": "Database not connected. Please connect a database."},
+ )
+
+ config = await proxy_config.get_config()
+
+ return await _get_settings_with_schema(
+ settings_key="mcp_semantic_tool_filter",
+ settings_class=MCPSemanticFilterSettings,
+ config=config,
+ )
+
+
+@router.patch(
+ "/update/mcp_semantic_filter_settings",
+ tags=["Settings"],
+ dependencies=[Depends(user_api_key_auth)],
+)
+async def update_mcp_semantic_filter_settings(
+ settings: MCPSemanticFilterSettings,
+ user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
+):
+ """
+ Update MCP semantic filter settings in database.
+ Settings will be picked up by all pods within approximately 10 seconds via background polling.
+ """
+ result = await _update_litellm_setting(
+ settings=settings,
+ settings_key="mcp_semantic_tool_filter",
+ in_memory_var=None,
+ success_message="MCP Semantic Filter settings updated successfully. Changes will be applied across all pods within 10 seconds.",
+ )
+ try:
+ from litellm.proxy.proxy_server import prisma_client, proxy_config
+
+ if prisma_client is not None:
+ await proxy_config._init_semantic_filter_settings_in_db(
+ prisma_client=prisma_client
+ )
+ except Exception as e:
+ verbose_proxy_logger.warning(
+ f"Failed to reinitialize MCP semantic filter settings immediately: {e}"
+ )
+
+ return result
+
+
@router.get(
"/in_product_nudges",
tags=["UI Settings"],
diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py
index 6bbf0df74de..0aace65ff6b 100644
--- a/litellm/proxy/utils.py
+++ b/litellm/proxy/utils.py
@@ -1370,17 +1370,36 @@ class ProxyLogging:
],
user_info: CallInfo,
):
- if self.alerting is None:
- # do nothing if alerting is not switched on
+ # For soft_budget alerts with alert_emails set, allow email sending even if alerting is None
+ # This enables team-specific soft budget email alerts via metadata.soft_budget_alerting_emails
+ # Note: user_info is a CallInfo that can represent user/team/org level info. For team budgets,
+ # alert_emails is populated from team_object.metadata.soft_budget_alerting_emails (see auth_checks.py)
+ is_soft_budget_with_alert_emails = (
+ type == "soft_budget"
+ and user_info.alert_emails is not None
+ and len(user_info.alert_emails) > 0
+ )
+
+ if self.alerting is None and not is_soft_budget_with_alert_emails:
+ # do nothing if alerting is not switched on (unless it's a soft_budget alert with team-specific emails)
return
- if "slack" in self.alerting:
- await self.slack_alerting_instance.budget_alerts(
- type=type,
- user_info=user_info,
- )
+ if self.alerting is not None and "slack" in self.alerting:
+ if self.slack_alerting_instance is not None:
+ await self.slack_alerting_instance.budget_alerts(
+ type=type,
+ user_info=user_info,
+ )
- if "email" in self.alerting and self.email_logging_instance is not None:
+ # Call email_logging_instance if:
+ # 1. "email" is in alerting config, OR
+ # 2. It's a soft_budget alert with team-specific alert_emails (bypasses global alerting config)
+ should_send_email = (
+ (self.alerting is not None and "email" in self.alerting)
+ or is_soft_budget_with_alert_emails
+ )
+
+ if should_send_email and self.email_logging_instance is not None:
await self.email_logging_instance.budget_alerts(
type=type,
user_info=user_info,
@@ -2607,7 +2626,8 @@ class PrismaClient:
SELECT
v.*,
t.spend AS team_spend,
- t.max_budget AS team_max_budget,
+ t.max_budget AS team_max_budget,
+ t.soft_budget AS team_soft_budget,
t.tpm_limit AS team_tpm_limit,
t.rpm_limit AS team_rpm_limit,
t.models AS team_models,
diff --git a/litellm/proxy_auth/__init__.py b/litellm/proxy_auth/__init__.py
new file mode 100644
index 00000000000..27624a94fb9
--- /dev/null
+++ b/litellm/proxy_auth/__init__.py
@@ -0,0 +1,30 @@
+"""
+Proxy Authentication module for LiteLLM SDK.
+
+This module provides OAuth2/JWT token management for authenticating
+with LiteLLM Proxy or any OAuth2-protected endpoint.
+
+Usage:
+ from litellm.proxy_auth import AzureADCredential, ProxyAuthHandler
+
+ litellm.proxy_auth = ProxyAuthHandler(
+ credential=AzureADCredential(),
+ scope="api://my-proxy/.default"
+ )
+"""
+
+from .credentials import (
+ AccessToken,
+ TokenCredential,
+ AzureADCredential,
+ GenericOAuth2Credential,
+ ProxyAuthHandler,
+)
+
+__all__ = [
+ "AccessToken",
+ "TokenCredential",
+ "AzureADCredential",
+ "GenericOAuth2Credential",
+ "ProxyAuthHandler",
+]
diff --git a/litellm/proxy_auth/credentials.py b/litellm/proxy_auth/credentials.py
new file mode 100644
index 00000000000..103b0088d80
--- /dev/null
+++ b/litellm/proxy_auth/credentials.py
@@ -0,0 +1,240 @@
+"""
+Credential providers for proxy authentication.
+
+This module provides a provider-agnostic interface for obtaining OAuth2/JWT tokens.
+It follows the same TokenCredential protocol used by Azure SDK.
+"""
+
+import time
+from dataclasses import dataclass
+from typing import Any, Optional, Protocol, runtime_checkable
+
+
+@dataclass
+class AccessToken:
+ """
+ Represents an OAuth2 access token with expiration.
+
+ This matches the structure used by azure.core.credentials.AccessToken.
+
+ Attributes:
+ token: The access token string (typically a JWT).
+ expires_on: Unix timestamp when the token expires.
+ """
+
+ token: str
+ expires_on: int
+
+
+@runtime_checkable
+class TokenCredential(Protocol):
+ """
+ Protocol for credential providers.
+
+ This matches the azure.core.credentials.TokenCredential interface,
+ allowing any Azure SDK credential to be used directly.
+
+ Any class implementing get_token(scope) -> AccessToken can be used.
+ """
+
+ def get_token(self, scope: str) -> AccessToken:
+ """
+ Get an access token for the specified scope.
+
+ Args:
+ scope: The OAuth2 scope to request (e.g., "api://my-app/.default")
+
+ Returns:
+ AccessToken with the token string and expiration timestamp.
+ """
+ ...
+
+
+class AzureADCredential:
+ """
+ Wrapper for Azure Identity credentials.
+
+ This wraps any azure-identity credential (DefaultAzureCredential,
+ ClientSecretCredential, ManagedIdentityCredential, etc.) and converts
+ the token to our AccessToken format.
+
+ If no credential is provided, it will use DefaultAzureCredential
+ which tries multiple authentication methods automatically.
+
+ Example:
+ # Use default credential chain (env vars, managed identity, CLI, etc.)
+ cred = AzureADCredential()
+
+ # Or provide a specific credential
+ from azure.identity import ClientSecretCredential
+ azure_cred = ClientSecretCredential(tenant_id, client_id, client_secret)
+ cred = AzureADCredential(credential=azure_cred)
+ """
+
+ def __init__(self, credential: Optional[Any] = None):
+ """
+ Initialize with an optional Azure credential.
+
+ Args:
+ credential: An azure-identity credential object. If None,
+ DefaultAzureCredential will be used on first token request.
+ """
+ self._credential: Any = credential
+ self._initialized = credential is not None
+
+ def get_token(self, scope: str) -> AccessToken:
+ """
+ Get an access token from Azure AD.
+
+ Args:
+ scope: The OAuth2 scope (e.g., "api://my-app/.default")
+
+ Returns:
+ AccessToken with the JWT and expiration.
+
+ Raises:
+ ImportError: If azure-identity is not installed.
+ """
+ if not self._initialized:
+ try:
+ from azure.identity import DefaultAzureCredential
+
+ self._credential = DefaultAzureCredential()
+ self._initialized = True
+ except ImportError:
+ raise ImportError(
+ "azure-identity is required for AzureADCredential. "
+ "Install it with: pip install azure-identity"
+ )
+
+ result = self._credential.get_token(scope)
+ return AccessToken(token=result.token, expires_on=result.expires_on)
+
+
+class GenericOAuth2Credential:
+ """
+ Generic OAuth2 client credentials flow.
+
+ This works with any OAuth2 provider (Okta, Auth0, Keycloak, etc.)
+ that supports the client_credentials grant type.
+
+ Example:
+ cred = GenericOAuth2Credential(
+ client_id="my-client-id",
+ client_secret="my-client-secret",
+ token_url="https://my-idp.com/oauth2/token"
+ )
+ """
+
+ def __init__(self, client_id: str, client_secret: str, token_url: str):
+ """
+ Initialize OAuth2 client credentials.
+
+ Args:
+ client_id: OAuth2 client ID
+ client_secret: OAuth2 client secret
+ token_url: Token endpoint URL (e.g., "https://idp.com/oauth2/token")
+ """
+ self.client_id = client_id
+ self.client_secret = client_secret
+ self.token_url = token_url
+ self._cached_token: Optional[AccessToken] = None
+
+ def get_token(self, scope: str) -> AccessToken:
+ """
+ Get an access token using OAuth2 client credentials flow.
+
+ Tokens are cached and reused until they expire (with 60s buffer).
+
+ Args:
+ scope: The OAuth2 scope to request
+
+ Returns:
+ AccessToken with the token and expiration.
+ """
+ # Return cached token if still valid (with 60s buffer)
+ if self._cached_token and self._cached_token.expires_on > time.time() + 60:
+ return self._cached_token
+
+ import httpx
+
+ response = httpx.post(
+ self.token_url,
+ data={
+ "grant_type": "client_credentials",
+ "client_id": self.client_id,
+ "client_secret": self.client_secret,
+ "scope": scope,
+ },
+ )
+ response.raise_for_status()
+ data = response.json()
+
+ self._cached_token = AccessToken(
+ token=data["access_token"],
+ expires_on=int(time.time()) + data.get("expires_in", 3600),
+ )
+ return self._cached_token
+
+
+class ProxyAuthHandler:
+ """
+ Manages OAuth2/JWT token lifecycle for proxy authentication.
+
+ This handler:
+ - Obtains tokens from the configured credential provider
+ - Caches tokens to avoid unnecessary requests
+ - Automatically refreshes tokens before they expire (60s buffer)
+ - Generates Authorization headers for HTTP requests
+
+ Set this as litellm.proxy_auth to automatically inject auth headers
+ into all requests to your LiteLLM Proxy.
+
+ Example:
+ import litellm
+ from litellm.proxy_auth import AzureADCredential, ProxyAuthHandler
+
+ litellm.proxy_auth = ProxyAuthHandler(
+ credential=AzureADCredential(),
+ scope="api://my-litellm-proxy/.default"
+ )
+ litellm.api_base = "https://my-proxy.example.com"
+
+ # Auth headers are now automatically injected
+ response = litellm.completion(model="gpt-4", messages=[...])
+ """
+
+ def __init__(self, credential: TokenCredential, scope: str):
+ """
+ Initialize the proxy auth handler.
+
+ Args:
+ credential: A TokenCredential implementation (AzureADCredential,
+ GenericOAuth2Credential, or any custom implementation)
+ scope: The OAuth2 scope to request tokens for
+ """
+ self.credential = credential
+ self.scope = scope
+ self._cached_token: Optional[AccessToken] = None
+
+ def get_token(self) -> AccessToken:
+ """
+ Get a valid access token, refreshing if necessary.
+
+ Returns:
+ AccessToken that is valid for at least 60 more seconds.
+ """
+ # Refresh if no token or token expires within 60 seconds
+ if not self._cached_token or self._cached_token.expires_on <= time.time() + 60:
+ self._cached_token = self.credential.get_token(self.scope)
+ return self._cached_token
+
+ def get_auth_headers(self) -> dict:
+ """
+ Get HTTP headers for authentication.
+
+ Returns:
+ Dict with Authorization header containing Bearer token.
+ """
+ token = self.get_token()
+ return {"Authorization": f"Bearer {token.token}"}
diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py
index 0a78fb7b72a..01b83067650 100644
--- a/litellm/realtime_api/main.py
+++ b/litellm/realtime_api/main.py
@@ -3,8 +3,8 @@
from typing import Any, Optional, cast
import litellm
-from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
+from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.secret_managers.main import get_secret_str
@@ -16,12 +16,16 @@ from litellm.utils import ProviderConfigManager
from ..litellm_core_utils.get_litellm_params import get_litellm_params
from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from ..llms.azure.realtime.handler import AzureOpenAIRealtime
-from ..llms.openai.realtime.handler import OpenAIRealtime
-from ..utils import client as wrapper_client
+from ..llms.bedrock.realtime.handler import BedrockRealtime
from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
+from ..llms.openai.realtime.handler import OpenAIRealtime
+from ..llms.xai.realtime.handler import XAIRealtime
+from ..utils import client as wrapper_client
azure_realtime = AzureOpenAIRealtime()
openai_realtime = OpenAIRealtime()
+bedrock_realtime = BedrockRealtime()
+xai_realtime = XAIRealtime()
base_llm_http_handler = BaseLLMHTTPHandler()
@@ -153,6 +157,63 @@ async def _arealtime(
timeout=timeout,
query_params=query_params,
)
+ elif _custom_llm_provider == "bedrock":
+ # Extract AWS parameters from kwargs
+ aws_region_name = kwargs.get("aws_region_name")
+ aws_access_key_id = kwargs.get("aws_access_key_id")
+ aws_secret_access_key = kwargs.get("aws_secret_access_key")
+ aws_session_token = kwargs.get("aws_session_token")
+ aws_role_name = kwargs.get("aws_role_name")
+ aws_session_name = kwargs.get("aws_session_name")
+ aws_profile_name = kwargs.get("aws_profile_name")
+ aws_web_identity_token = kwargs.get("aws_web_identity_token")
+ aws_sts_endpoint = kwargs.get("aws_sts_endpoint")
+ aws_bedrock_runtime_endpoint = kwargs.get("aws_bedrock_runtime_endpoint")
+ aws_external_id = kwargs.get("aws_external_id")
+
+ await bedrock_realtime.async_realtime(
+ model=model,
+ websocket=websocket,
+ logging_obj=litellm_logging_obj,
+ api_base=dynamic_api_base or api_base,
+ api_key=dynamic_api_key or api_key,
+ timeout=timeout,
+ aws_region_name=aws_region_name,
+ aws_access_key_id=aws_access_key_id,
+ aws_secret_access_key=aws_secret_access_key,
+ aws_session_token=aws_session_token,
+ aws_role_name=aws_role_name,
+ aws_session_name=aws_session_name,
+ aws_profile_name=aws_profile_name,
+ aws_web_identity_token=aws_web_identity_token,
+ aws_sts_endpoint=aws_sts_endpoint,
+ aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
+ aws_external_id=aws_external_id,
+ )
+ elif _custom_llm_provider == "xai":
+ api_base = (
+ dynamic_api_base
+ or litellm_params.api_base
+ or get_secret_str("XAI_API_BASE")
+ or "https://api.x.ai/v1"
+ )
+ # set API KEY
+ api_key = (
+ dynamic_api_key
+ or litellm.api_key
+ or get_secret_str("XAI_API_KEY")
+ )
+
+ await xai_realtime.async_realtime(
+ model=model,
+ websocket=websocket,
+ logging_obj=litellm_logging_obj,
+ api_base=api_base,
+ api_key=api_key,
+ client=None,
+ timeout=timeout,
+ query_params=query_params,
+ )
else:
raise ValueError(f"Unsupported model: {model}")
@@ -195,6 +256,10 @@ async def _realtime_health_check(
url = openai_realtime._construct_url(
api_base=api_base or "https://api.openai.com/", query_params={"model": model}
)
+ elif custom_llm_provider == "xai":
+ url = xai_realtime._construct_url(
+ api_base=api_base or "https://api.x.ai/v1", query_params={"model": model}
+ )
else:
raise ValueError(f"Unsupported model: {model}")
ssl_context = get_shared_realtime_ssl_context()
diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py
index 74cc87713da..df298f7c448 100644
--- a/litellm/responses/litellm_completion_transformation/transformation.py
+++ b/litellm/responses/litellm_completion_transformation/transformation.py
@@ -1413,6 +1413,7 @@ class LiteLLMCompletionResponsesConfig:
),
user=getattr(chat_completion_response, "user", None),
)
+ responses_api_response._hidden_params = getattr(chat_completion_response, "_hidden_params", {})
return responses_api_response
@staticmethod
diff --git a/litellm/responses/main.py b/litellm/responses/main.py
index b2c2493c812..8f524690be1 100644
--- a/litellm/responses/main.py
+++ b/litellm/responses/main.py
@@ -24,6 +24,7 @@ from litellm.constants import request_timeout
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.prompt_templates.common_utils import (
update_responses_input_with_model_file_ids,
+ update_responses_tools_with_model_file_ids,
)
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
@@ -595,13 +596,32 @@ def responses(
litellm_params.api_base = dynamic_api_base
#########################################################
- # Update input with provider-specific file IDs if managed files are used
+ # Update input and tools with provider-specific file IDs if managed files are used
#########################################################
+ model_file_id_mapping = kwargs.get("model_file_id_mapping")
+ model_info_id = kwargs.get("model_info", {}).get("id") if isinstance(kwargs.get("model_info"), dict) else None
+
input = cast(
Union[str, ResponseInputParam],
- update_responses_input_with_model_file_ids(input=input),
+ update_responses_input_with_model_file_ids(
+ input=input,
+ model_id=model_info_id,
+ model_file_id_mapping=model_file_id_mapping,
+ ),
)
local_vars["input"] = input
+
+ # Update tools with provider-specific file IDs if needed
+ if tools:
+ tools = cast(
+ Optional[Iterable[ToolParam]],
+ update_responses_tools_with_model_file_ids(
+ tools=cast(Optional[List[Dict[str, Any]]], tools),
+ model_id=model_info_id,
+ model_file_id_mapping=model_file_id_mapping,
+ ),
+ )
+ local_vars["tools"] = tools
#########################################################
# Native MCP Responses API
diff --git a/litellm/router.py b/litellm/router.py
index ed480d6468a..374f80db361 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -117,6 +117,9 @@ from litellm.router_utils.pre_call_checks.prompt_caching_deployment_check import
from litellm.router_utils.pre_call_checks.responses_api_deployment_check import (
ResponsesApiDeploymentCheck,
)
+from litellm.router_utils.pre_call_checks.model_rate_limit_check import (
+ ModelRateLimitingCheck,
+)
from litellm.router_utils.router_callbacks.track_deployment_metrics import (
increment_deployment_failures_for_current_minute,
increment_deployment_successes_for_current_minute,
@@ -224,6 +227,7 @@ class Router:
redis_host: Optional[str] = None,
redis_port: Optional[int] = None,
redis_password: Optional[str] = None,
+ redis_db: Optional[int] = None,
cache_responses: Optional[bool] = False,
cache_kwargs: dict = {}, # additional kwargs to pass to RedisCache (see caching.py)
caching_groups: Optional[
@@ -410,6 +414,12 @@ class Router:
if redis_password is not None:
cache_config["password"] = redis_password
+ if redis_db is not None:
+ verbose_router_logger.warning(
+ "Deprecated 'redis_db' argument used. Please remove 'redis_db' from your config/database and use 'cache_kwargs' instead."
+ )
+ cache_config["db"] = str(redis_db)
+
# Add additional key-value pairs from cache_kwargs
cache_config.update(cache_kwargs)
redis_cache = self._create_redis_cache(cache_config)
@@ -1187,6 +1197,8 @@ class Router:
)
elif pre_call_check == "responses_api_deployment_check":
_callback = ResponsesApiDeploymentCheck()
+ elif pre_call_check == "enforce_model_rate_limits":
+ _callback = ModelRateLimitingCheck(dual_cache=self.cache)
if _callback is not None:
if self.optional_callbacks is None:
self.optional_callbacks = []
@@ -3657,7 +3669,7 @@ class Router:
)
raise e
- async def _acreate_file(
+ async def _acreate_file( # noqa: PLR0915
self,
model: str,
**kwargs,
@@ -3718,7 +3730,8 @@ class Router:
)
kwargs_copy["file"] = file
-
+ if "gcs_bucket_name" in data: # TODO: Remove this once we have a better way to handle GCS bucket name: Problem is that we need to pass the gcs_bucket_name to the router for the create_file call but it doesn't show up there
+ kwargs_copy.setdefault("litellm_metadata", {})["gcs_bucket_name"] = data["gcs_bucket_name"]
response = litellm.acreate_file(
**{
**data,
diff --git a/litellm/router_utils/pre_call_checks/model_rate_limit_check.py b/litellm/router_utils/pre_call_checks/model_rate_limit_check.py
new file mode 100644
index 00000000000..e5be61690ba
--- /dev/null
+++ b/litellm/router_utils/pre_call_checks/model_rate_limit_check.py
@@ -0,0 +1,373 @@
+"""
+Enforce TPM/RPM rate limits set on model deployments.
+
+This pre-call check ensures that model-level TPM/RPM limits are enforced
+across all requests, regardless of routing strategy.
+
+When enabled via `enforce_model_rate_limits: true` in litellm_settings,
+requests that exceed the configured TPM/RPM limits will receive a 429 error.
+"""
+
+from typing import TYPE_CHECKING, Any, Dict, Optional, Union
+
+import httpx
+
+import litellm
+from litellm._logging import verbose_router_logger
+from litellm.caching.dual_cache import DualCache
+from litellm.integrations.custom_logger import CustomLogger
+from litellm.types.router import RouterErrors
+from litellm.types.utils import StandardLoggingPayload
+from litellm.utils import get_utc_datetime
+
+if TYPE_CHECKING:
+ from opentelemetry.trace import Span as _Span
+
+ Span = Union[_Span, Any]
+else:
+ Span = Any
+
+
+class RoutingArgs:
+ ttl: int = 60 # 1min (RPM/TPM expire key)
+
+
+class ModelRateLimitingCheck(CustomLogger):
+ """
+ Pre-call check that enforces TPM/RPM limits on model deployments.
+
+ This check runs before each request and raises a RateLimitError
+ if the deployment has exceeded its configured TPM or RPM limits.
+
+ Unlike the usage-based-routing strategy which uses limits for routing decisions,
+ this check actively enforces those limits across ALL routing strategies.
+ """
+
+ def __init__(self, dual_cache: DualCache):
+ self.dual_cache = dual_cache
+
+ def _get_deployment_limits(
+ self, deployment: Dict
+ ) -> tuple[Optional[int], Optional[int]]:
+ """
+ Extract TPM and RPM limits from a deployment configuration.
+
+ Checks in order:
+ 1. Top-level 'tpm'/'rpm' fields
+ 2. litellm_params.tpm/rpm
+ 3. model_info.tpm/rpm
+
+ Returns:
+ Tuple of (tpm_limit, rpm_limit)
+ """
+ # Check top-level
+ tpm = deployment.get("tpm")
+ rpm = deployment.get("rpm")
+
+ # Check litellm_params
+ if tpm is None:
+ tpm = deployment.get("litellm_params", {}).get("tpm")
+ if rpm is None:
+ rpm = deployment.get("litellm_params", {}).get("rpm")
+
+ # Check model_info
+ if tpm is None:
+ tpm = deployment.get("model_info", {}).get("tpm")
+ if rpm is None:
+ rpm = deployment.get("model_info", {}).get("rpm")
+
+ return tpm, rpm
+
+ def _get_cache_keys(self, deployment: Dict, current_minute: str) -> tuple[str, str]:
+ """Get the cache keys for TPM and RPM tracking."""
+ model_id = deployment.get("model_info", {}).get("id")
+ deployment_name = deployment.get("litellm_params", {}).get("model")
+
+ tpm_key = f"{model_id}:{deployment_name}:tpm:{current_minute}"
+ rpm_key = f"{model_id}:{deployment_name}:rpm:{current_minute}"
+
+ return tpm_key, rpm_key
+
+ def pre_call_check(self, deployment: Dict) -> Optional[Dict]:
+ """
+ Synchronous pre-call check for model rate limits.
+
+ Raises RateLimitError if deployment exceeds TPM/RPM limits.
+ """
+ try:
+ tpm_limit, rpm_limit = self._get_deployment_limits(deployment)
+
+ # If no limits are set, allow the request
+ if tpm_limit is None and rpm_limit is None:
+ return deployment
+
+ dt = get_utc_datetime()
+ current_minute = dt.strftime("%H-%M")
+ tpm_key, rpm_key = self._get_cache_keys(deployment, current_minute)
+
+ model_id = deployment.get("model_info", {}).get("id")
+ model_name = deployment.get("litellm_params", {}).get("model")
+ model_group = deployment.get("model_name", "")
+
+ # Check TPM limit
+ if tpm_limit is not None:
+ # First check local cache
+ current_tpm = self.dual_cache.get_cache(key=tpm_key, local_only=True)
+ if current_tpm is not None and current_tpm >= tpm_limit:
+ raise litellm.RateLimitError(
+ message=f"Model rate limit exceeded. TPM limit={tpm_limit}, current usage={current_tpm}",
+ llm_provider="",
+ model=model_name,
+ response=httpx.Response(
+ status_code=429,
+ content=f"{RouterErrors.user_defined_ratelimit_error.value} tpm limit={tpm_limit}. current usage={current_tpm}. id={model_id}, model_group={model_group}",
+ headers={"retry-after": str(60)},
+ request=httpx.Request(
+ method="model_rate_limit_check",
+ url="https://github.com/BerriAI/litellm",
+ ),
+ ),
+ )
+
+ # Check RPM limit
+ if rpm_limit is not None:
+ # First check local cache
+ current_rpm = self.dual_cache.get_cache(key=rpm_key, local_only=True)
+ if current_rpm >= rpm_limit:
+ raise litellm.RateLimitError(
+ message=f"Model rate limit exceeded. RPM limit={rpm_limit}, current usage={current_rpm}",
+ llm_provider="",
+ model=model_name,
+ response=httpx.Response(
+ status_code=429,
+ content=f"{RouterErrors.user_defined_ratelimit_error.value} rpm limit={rpm_limit}. current usage={current_rpm}. id={model_id}, model_group={model_group}",
+ headers={"retry-after": str(60)},
+ request=httpx.Request(
+ method="model_rate_limit_check",
+ url="https://github.com/BerriAI/litellm",
+ ),
+ ),
+ )
+
+ # Check redis cache and increment
+ current_rpm = self.dual_cache.increment_cache(
+ key=rpm_key, value=1, ttl=RoutingArgs.ttl
+ )
+ if current_rpm is not None and current_rpm > rpm_limit:
+ raise litellm.RateLimitError(
+ message=f"Model rate limit exceeded. RPM limit={rpm_limit}, current usage={current_rpm}",
+ llm_provider="",
+ model=model_name,
+ response=httpx.Response(
+ status_code=429,
+ content=f"{RouterErrors.user_defined_ratelimit_error.value} rpm limit={rpm_limit}. current usage={current_rpm}. id={model_id}, model_group={model_group}",
+ headers={"retry-after": str(60)},
+ request=httpx.Request(
+ method="model_rate_limit_check",
+ url="https://github.com/BerriAI/litellm",
+ ),
+ ),
+ )
+
+ return deployment
+
+ except litellm.RateLimitError:
+ raise
+ except Exception as e:
+ verbose_router_logger.debug(
+ f"Error in ModelRateLimitingCheck.pre_call_check: {str(e)}"
+ )
+ # Don't fail the request if rate limit check fails
+ return deployment
+
+ async def async_pre_call_check(
+ self, deployment: Dict, parent_otel_span: Optional[Span] = None
+ ) -> Optional[Dict]:
+ """
+ Async pre-call check for model rate limits.
+
+ Raises RateLimitError if deployment exceeds TPM/RPM limits.
+ """
+ try:
+ tpm_limit, rpm_limit = self._get_deployment_limits(deployment)
+
+ # If no limits are set, allow the request
+ if tpm_limit is None and rpm_limit is None:
+ return deployment
+
+ dt = get_utc_datetime()
+ current_minute = dt.strftime("%H-%M")
+ tpm_key, rpm_key = self._get_cache_keys(deployment, current_minute)
+
+ model_id = deployment.get("model_info", {}).get("id")
+ model_name = deployment.get("litellm_params", {}).get("model")
+ model_group = deployment.get("model_name", "")
+
+ # Check TPM limit
+ if tpm_limit is not None:
+ # First check local cache
+ current_tpm = await self.dual_cache.async_get_cache(
+ key=tpm_key, local_only=True
+ )
+ if current_tpm is not None and current_tpm >= tpm_limit:
+ raise litellm.RateLimitError(
+ message=f"Model rate limit exceeded. TPM limit={tpm_limit}, current usage={current_tpm}",
+ llm_provider="",
+ model=model_name,
+ response=httpx.Response(
+ status_code=429,
+ content=f"{RouterErrors.user_defined_ratelimit_error.value} tpm limit={tpm_limit}. current usage={current_tpm}. id={model_id}, model_group={model_group}",
+ headers={"retry-after": str(60)},
+ request=httpx.Request(
+ method="model_rate_limit_check",
+ url="https://github.com/BerriAI/litellm",
+ ),
+ ),
+ num_retries=0, # Don't retry - return 429 immediately
+ )
+
+ # Check RPM limit
+ if rpm_limit is not None:
+ # First check local cache
+ current_rpm = await self.dual_cache.async_get_cache(
+ key=rpm_key, local_only=True
+ )
+ if current_rpm is not None and current_rpm >= rpm_limit:
+ raise litellm.RateLimitError(
+ message=f"Model rate limit exceeded. RPM limit={rpm_limit}, current usage={current_rpm}",
+ llm_provider="",
+ model=model_name,
+ response=httpx.Response(
+ status_code=429,
+ content=f"{RouterErrors.user_defined_ratelimit_error.value} rpm limit={rpm_limit}. current usage={current_rpm}. id={model_id}, model_group={model_group}",
+ headers={"retry-after": str(60)},
+ request=httpx.Request(
+ method="model_rate_limit_check",
+ url="https://github.com/BerriAI/litellm",
+ ),
+ ),
+ num_retries=0, # Don't retry - return 429 immediately
+ )
+
+ # Check redis cache and increment
+ current_rpm = await self.dual_cache.async_increment_cache(
+ key=rpm_key,
+ value=1,
+ ttl=RoutingArgs.ttl,
+ parent_otel_span=parent_otel_span,
+ )
+ if current_rpm is not None and current_rpm > rpm_limit:
+ raise litellm.RateLimitError(
+ message=f"Model rate limit exceeded. RPM limit={rpm_limit}, current usage={current_rpm}",
+ llm_provider="",
+ model=model_name,
+ response=httpx.Response(
+ status_code=429,
+ content=f"{RouterErrors.user_defined_ratelimit_error.value} rpm limit={rpm_limit}. current usage={current_rpm}. id={model_id}, model_group={model_group}",
+ headers={"retry-after": str(60)},
+ request=httpx.Request(
+ method="model_rate_limit_check",
+ url="https://github.com/BerriAI/litellm",
+ ),
+ ),
+ num_retries=0, # Don't retry - return 429 immediately
+ )
+
+ return deployment
+
+ except litellm.RateLimitError:
+ raise
+ except Exception as e:
+ verbose_router_logger.debug(
+ f"Error in ModelRateLimitingCheck.async_pre_call_check: {str(e)}"
+ )
+ # Don't fail the request if rate limit check fails
+ return deployment
+
+ async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
+ """
+ Track TPM usage after successful request.
+
+ This updates the TPM counter with the actual tokens used.
+ Always tracks tokens - the pre-call check handles enforcement.
+ """
+ try:
+ standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get(
+ "standard_logging_object"
+ )
+ if standard_logging_object is None:
+ return
+
+ model_id = standard_logging_object.get("model_id")
+ if model_id is None:
+ return
+
+ total_tokens = standard_logging_object.get("total_tokens", 0)
+ model = standard_logging_object.get("hidden_params", {}).get(
+ "litellm_model_name"
+ )
+
+ verbose_router_logger.debug(
+ f"[TPM TRACKING] model_id={model_id}, total_tokens={total_tokens}, model={model}"
+ )
+
+ if not model or not total_tokens:
+ return
+
+ dt = get_utc_datetime()
+ current_minute = dt.strftime("%H-%M")
+ tpm_key = f"{model_id}:{model}:tpm:{current_minute}"
+
+ verbose_router_logger.debug(
+ f"[TPM TRACKING] Incrementing {tpm_key} by {total_tokens}"
+ )
+
+ await self.dual_cache.async_increment_cache(
+ key=tpm_key,
+ value=total_tokens,
+ ttl=RoutingArgs.ttl,
+ )
+
+ except Exception as e:
+ verbose_router_logger.debug(
+ f"Error in ModelRateLimitingCheck.async_log_success_event: {str(e)}"
+ )
+
+ def log_success_event(self, kwargs, response_obj, start_time, end_time):
+ """
+ Sync version of tracking TPM usage after successful request.
+ Always tracks tokens - the pre-call check handles enforcement.
+ """
+ try:
+ standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get(
+ "standard_logging_object"
+ )
+ if standard_logging_object is None:
+ return
+
+ model_id = standard_logging_object.get("model_id")
+ if model_id is None:
+ return
+
+ total_tokens = standard_logging_object.get("total_tokens", 0)
+ model = standard_logging_object.get("hidden_params", {}).get(
+ "litellm_model_name"
+ )
+
+ if not model or not total_tokens:
+ return
+
+ dt = get_utc_datetime()
+ current_minute = dt.strftime("%H-%M")
+ tpm_key = f"{model_id}:{model}:tpm:{current_minute}"
+
+ self.dual_cache.increment_cache(
+ key=tpm_key,
+ value=total_tokens,
+ ttl=RoutingArgs.ttl,
+ )
+
+ except Exception as e:
+ verbose_router_logger.debug(
+ f"Error in ModelRateLimitingCheck.log_success_event: {str(e)}"
+ )
diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py
index ca22049720e..74ccb34ca6e 100644
--- a/litellm/types/guardrails.py
+++ b/litellm/types/guardrails.py
@@ -14,14 +14,14 @@ from litellm.types.proxy.guardrails.guardrail_hooks.grayswan import (
from litellm.types.proxy.guardrails.guardrail_hooks.ibm import (
IBMGuardrailsBaseConfigModel,
)
-from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import (
- ToolPermissionGuardrailConfigModel,
+from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
+ ContentFilterCategoryConfig,
)
from litellm.types.proxy.guardrails.guardrail_hooks.qualifire import (
QualifireGuardrailConfigModel,
)
-from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
- ContentFilterCategoryConfig,
+from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import (
+ ToolPermissionGuardrailConfigModel,
)
"""
@@ -68,6 +68,7 @@ class SupportedGuardrailIntegrations(Enum):
PROMPT_SECURITY = "prompt_security"
GENERIC_GUARDRAIL_API = "generic_guardrail_api"
QUALIFIRE = "qualifire"
+ CUSTOM_CODE = "custom_code"
class Role(Enum):
@@ -296,13 +297,7 @@ class PresidioConfigModel(PresidioPresidioConfigModelUserInterface):
pii_entities_config: Optional[Dict[Union[PiiEntityType, str], PiiAction]] = Field(
default=None, description="Configuration for PII entity types and actions"
)
- presidio_filter_scope: Literal["input", "output", "both"] = Field(
- default="both",
- description=(
- "Where to apply Presidio checks: 'input' runs on user → model traffic, "
- "'output' runs on model → user traffic, and 'both' applies to both."
- ),
- )
+
presidio_score_thresholds: Optional[Dict[Union[PiiEntityType, str], float]] = Field(
default=None,
description=(
@@ -656,6 +651,12 @@ class BaseLitellmParams(
description="Additional provider-specific parameters for generic guardrail APIs",
)
+ # Custom code guardrail params
+ custom_code: Optional[str] = Field(
+ default=None,
+ description="Python-like code containing the apply_guardrail function for custom guardrail logic",
+ )
+
model_config = ConfigDict(extra="allow", protected_namespaces=())
diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py
index 62e775d4faa..fedf419efd6 100644
--- a/litellm/types/llms/anthropic.py
+++ b/litellm/types/llms/anthropic.py
@@ -613,7 +613,7 @@ ANTHROPIC_API_ONLY_HEADERS = { # fails if calling anthropic on vertex ai / bedr
class AnthropicThinkingParam(TypedDict, total=False):
- type: Literal["enabled"]
+ type: Literal["enabled", "adaptive"]
budget_tokens: int
@@ -633,6 +633,7 @@ class ANTHROPIC_BETA_HEADER_VALUES(str, Enum):
WEB_FETCH_2025_09_10 = "web-fetch-2025-09-10"
WEB_SEARCH_2025_03_05 = "web-search-2025-03-05"
CONTEXT_MANAGEMENT_2025_06_27 = "context-management-2025-06-27"
+ COMPACT_2026_01_12 = "compact-2026-01-12"
STRUCTURED_OUTPUT_2025_09_25 = "structured-outputs-2025-11-13"
ADVANCED_TOOL_USE_2025_11_20 = "advanced-tool-use-2025-11-20"
diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py
index a85aaafe23d..998c60ab60d 100644
--- a/litellm/types/llms/bedrock.py
+++ b/litellm/types/llms/bedrock.py
@@ -8,6 +8,7 @@ from .openai import ChatCompletionToolCallChunk
class CachePointBlock(TypedDict, total=False):
type: Literal["default"]
+ ttl: str
class SystemContentBlock(TypedDict, total=False):
@@ -397,6 +398,7 @@ class CohereEmbeddingRequest(TypedDict, total=False):
input_type: Required[COHERE_EMBEDDING_INPUT_TYPES]
truncate: Literal["NONE", "START", "END"]
embedding_types: Literal["float", "int8", "uint8", "binary", "ubinary"]
+ output_dimension: int
class CohereEmbeddingRequestWithModel(CohereEmbeddingRequest):
@@ -960,6 +962,7 @@ class BedrockGetBatchResponse(TypedDict, total=False):
timeoutDurationInHours: Optional[int]
clientRequestToken: Optional[str]
+
class BedrockToolBlock(TypedDict, total=False):
toolSpec: Optional[ToolSpecBlock]
systemTool: Optional[SystemToolBlock] # For Nova grounding
diff --git a/litellm/types/proxy/management_endpoints/scim_v2.py b/litellm/types/proxy/management_endpoints/scim_v2.py
index bff9f0b876c..c4d95d99ed4 100644
--- a/litellm/types/proxy/management_endpoints/scim_v2.py
+++ b/litellm/types/proxy/management_endpoints/scim_v2.py
@@ -1,7 +1,7 @@
from typing import Any, Dict, List, Literal, Optional, Union
from fastapi import HTTPException
-from pydantic import BaseModel, EmailStr, field_validator
+from pydantic import BaseModel, ConfigDict, EmailStr, field_validator
class LiteLLM_UserScimMetadata(BaseModel):
@@ -112,3 +112,67 @@ class SCIMServiceProviderConfig(BaseModel):
etag: SCIMFeature = SCIMFeature(supported=False)
authenticationSchemes: Optional[List[Dict[str, Any]]] = None
meta: Optional[Dict[str, Any]] = None
+
+
+# SCIM ResourceType Models (RFC 7643 Section 6)
+class SCIMSchemaExtension(BaseModel):
+ model_config = ConfigDict(populate_by_name=True)
+
+ schema_: str # aliased to "schema" in serialization
+ required: bool
+
+ def model_dump(self, **kwargs):
+ d = super().model_dump(**kwargs)
+ d["schema"] = d.pop("schema_")
+ return d
+
+
+class SCIMResourceType(BaseModel):
+ model_config = ConfigDict(populate_by_name=True)
+
+ schemas: List[str] = [
+ "urn:ietf:params:scim:schemas:core:2.0:ResourceType"
+ ]
+ id: str
+ name: str
+ description: Optional[str] = None
+ endpoint: str
+ schema_: str # "schema" is a reserved name in Pydantic context
+
+ schemaExtensions: Optional[List[SCIMSchemaExtension]] = None
+ meta: Optional[Dict[str, Any]] = None
+
+ def model_dump(self, **kwargs):
+ d = super().model_dump(**kwargs)
+ d["schema"] = d.pop("schema_")
+ if d.get("schemaExtensions") is None:
+ d.pop("schemaExtensions", None)
+ return d
+
+
+# SCIM Schema Models (RFC 7643 Section 7)
+class SCIMSchemaAttribute(BaseModel):
+ name: str
+ type: str
+ multiValued: bool = False
+ description: Optional[str] = None
+ required: bool = False
+ mutability: str = "readWrite"
+ returned: str = "default"
+ uniqueness: str = "none"
+ subAttributes: Optional[List["SCIMSchemaAttribute"]] = None
+
+ def model_dump(self, **kwargs):
+ d = super().model_dump(**kwargs)
+ if d.get("subAttributes") is None:
+ d.pop("subAttributes", None)
+ return d
+
+
+class SCIMSchema(BaseModel):
+ schemas: List[str] = ["urn:ietf:params:scim:schemas:core:2.0:Schema"]
+ id: str
+ name: str
+ description: Optional[str] = None
+ attributes: List[SCIMSchemaAttribute] = []
+ meta: Optional[Dict[str, Any]] = None
diff --git a/litellm/types/proxy/management_endpoints/ui_sso.py b/litellm/types/proxy/management_endpoints/ui_sso.py
index c9d998f6a92..6743c4a5b9b 100644
--- a/litellm/types/proxy/management_endpoints/ui_sso.py
+++ b/litellm/types/proxy/management_endpoints/ui_sso.py
@@ -86,6 +86,20 @@ class RoleMappings(LiteLLMPydanticObjectBase):
)
+class TeamMappings(LiteLLMPydanticObjectBase):
+ """
+ Configuration for mapping SSO JWT fields to team IDs.
+
+ This allows configuring team_ids_jwt_field via the database instead of
+ requiring config file changes and restarts.
+ """
+
+ team_ids_jwt_field: Optional[str] = Field(
+ default=None,
+ description="The field name in the SSO/JWT token that contains the team IDs array (e.g., 'groups', 'teams'). Supports dot notation for nested fields.",
+ )
+
+
class SSOConfig(LiteLLMPydanticObjectBase):
"""
Configuration for SSO environment variables and settings
@@ -159,6 +173,12 @@ class SSOConfig(LiteLLMPydanticObjectBase):
description="Configuration for mapping SSO groups to LiteLLM roles based on group claims in the SSO token",
)
+ # Team Mappings
+ team_mappings: Optional[TeamMappings] = Field(
+ default=None,
+ description="Configuration for mapping SSO JWT fields to team IDs. Takes precedence over config file settings.",
+ )
+
class DefaultTeamSSOParams(LiteLLMPydanticObjectBase):
"""
diff --git a/litellm/types/router.py b/litellm/types/router.py
index f31c6df3005..f78789c9772 100644
--- a/litellm/types/router.py
+++ b/litellm/types/router.py
@@ -95,18 +95,16 @@ class ModelInfo(BaseModel):
id: Optional[
str
] # Allow id to be optional on input, but it will always be present as a str in the model instance
- db_model: bool = (
- False # used for proxy - to separate models which are stored in the db vs. config.
- )
+ db_model: bool = False # used for proxy - to separate models which are stored in the db vs. config.
updated_at: Optional[datetime.datetime] = None
updated_by: Optional[str] = None
created_at: Optional[datetime.datetime] = None
created_by: Optional[str] = None
- base_model: Optional[str] = (
- None # specify if the base model is azure/gpt-3.5-turbo etc for accurate cost tracking
- )
+ base_model: Optional[
+ str
+ ] = None # specify if the base model is azure/gpt-3.5-turbo etc for accurate cost tracking
tier: Optional[Literal["free", "paid"]] = None
"""
@@ -172,12 +170,12 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
custom_llm_provider: Optional[str] = None
tpm: Optional[int] = None
rpm: Optional[int] = None
- timeout: Optional[Union[float, str, httpx.Timeout]] = (
- None # if str, pass in as os.environ/
- )
- stream_timeout: Optional[Union[float, str]] = (
- None # timeout when making stream=True calls, if str, pass in as os.environ/
- )
+ timeout: Optional[
+ Union[float, str, httpx.Timeout]
+ ] = None # if str, pass in as os.environ/
+ stream_timeout: Optional[
+ Union[float, str]
+ ] = None # timeout when making stream=True calls, if str, pass in as os.environ/
max_retries: Optional[int] = None
organization: Optional[str] = None # for openai orgs
configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None
@@ -276,9 +274,9 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
if max_retries is not None and isinstance(max_retries, str):
max_retries = int(max_retries) # cast to int
# We need to keep max_retries in args since it's a parameter of GenericLiteLLMParams
- args["max_retries"] = (
- max_retries # Put max_retries back in args after popping it
- )
+ args[
+ "max_retries"
+ ] = max_retries # Put max_retries back in args after popping it
super().__init__(**args, **params)
def __contains__(self, key):
@@ -805,6 +803,7 @@ OptionalPreCallChecks = List[
"router_budget_limiting",
"responses_api_deployment_check",
"forward_client_headers_by_model_group",
+ "enforce_model_rate_limits",
]
]
diff --git a/litellm/types/search.py b/litellm/types/search.py
index 661a2feda33..b0ce0636aed 100644
--- a/litellm/types/search.py
+++ b/litellm/types/search.py
@@ -60,6 +60,7 @@ class SearchToolInfoResponse(TypedDict, total=False):
search_tool_info: Optional[dict]
created_at: Optional[str]
updated_at: Optional[str]
+ is_from_config: Optional[bool] # True if this tool is defined in config file, False if from DB
class ListSearchToolsResponse(TypedDict):
diff --git a/litellm/types/utils.py b/litellm/types/utils.py
index 6c330d0f83c..e1f780ffcc3 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -2129,6 +2129,7 @@ class ImageObject(OpenAIImage):
b64_json: The base64-encoded JSON of the generated image, if response_format is b64_json.
url: The URL of the generated image, if response_format is url (default).
revised_prompt: The prompt that was used to generate the image, if there was any revision to the prompt.
+ provider_specific_fields: Provider-specific fields not part of OpenAI spec.
https://platform.openai.com/docs/api-reference/images/object
"""
@@ -2136,9 +2137,12 @@ class ImageObject(OpenAIImage):
b64_json: Optional[str] = None
url: Optional[str] = None
revised_prompt: Optional[str] = None
+ provider_specific_fields: Optional[Dict[str, Any]] = None
- def __init__(self, b64_json=None, url=None, revised_prompt=None, **kwargs):
+ def __init__(self, b64_json=None, url=None, revised_prompt=None, provider_specific_fields=None, **kwargs):
super().__init__(b64_json=b64_json, url=url, revised_prompt=revised_prompt) # type: ignore
+ if provider_specific_fields:
+ self.provider_specific_fields = provider_specific_fields
def __contains__(self, key):
# Define custom behavior for the 'in' operator
@@ -3025,6 +3029,7 @@ class LlmProviders(str, Enum):
MISTRAL = "mistral"
MILVUS = "milvus"
GROQ = "groq"
+ A2A = "a2a"
GIGACHAT = "gigachat"
NVIDIA_NIM = "nvidia_nim"
CEREBRAS = "cerebras"
diff --git a/litellm/utils.py b/litellm/utils.py
index 7c4eec7ba32..7109f7aa881 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -199,6 +199,8 @@ from litellm.types.utils import (
all_litellm_params,
)
+_CALL_TYPE_ENUM_MAP: dict = {ct.value: ct for ct in CallTypes}
+
# +-----------------------------------------------+
# | |
# | Give Feedback / Get Help |
@@ -1451,6 +1453,10 @@ def client(original_function): # noqa: PLR0915
logging_obj, kwargs = function_setup(
original_function.__name__, rules_obj, start_time, *args, **kwargs
)
+
+ # Type assertion: logging_obj is guaranteed to be non-None after function_setup
+ assert logging_obj is not None, "logging_obj should not be None after function_setup"
+
## LOAD CREDENTIALS
load_credentials_from_list(kwargs)
kwargs["litellm_logging_obj"] = logging_obj
@@ -1746,6 +1752,7 @@ def client(original_function): # noqa: PLR0915
print_args_passed_to_litellm(original_function, args, kwargs)
start_time = datetime.datetime.now()
result = None
+ _update_response_metadata = getattr(sys.modules[__name__], "update_response_metadata")
logging_obj: Optional[LiteLLMLoggingObject] = kwargs.get(
"litellm_logging_obj", None
)
@@ -1768,6 +1775,9 @@ def client(original_function): # noqa: PLR0915
logging_obj, kwargs = function_setup(
original_function.__name__, rules_obj, start_time, *args, **kwargs
)
+
+ # Type assertion: logging_obj is guaranteed to be non-None after function_setup
+ assert logging_obj is not None, "logging_obj should not be None after function_setup"
modified_kwargs = await async_pre_call_deployment_hook(kwargs, call_type)
if modified_kwargs is not None:
@@ -1786,9 +1796,10 @@ def client(original_function): # noqa: PLR0915
)
# [OPTIONAL] CHECK CACHE
- print_verbose(
- f"ASYNC kwargs[caching]: {kwargs.get('caching', False)}; litellm.cache: {litellm.cache}; kwargs.get('cache'): {kwargs.get('cache', None)}"
- )
+ if _is_debugging_on():
+ print_verbose(
+ f"ASYNC kwargs[caching]: {kwargs.get('caching', False)}; litellm.cache: {litellm.cache}; kwargs.get('cache'): {kwargs.get('cache', None)}"
+ )
_caching_handler_response: "Optional[CachingHandlerResponse]" = (
await _llm_caching_handler._async_get_cache(
model=model or "",
@@ -1864,10 +1875,7 @@ def client(original_function): # noqa: PLR0915
chunks, messages=kwargs.get("messages", None)
)
else:
- update_response_metadata = getattr(
- sys.modules[__name__], "update_response_metadata"
- )
- update_response_metadata(
+ _update_response_metadata(
result=result,
logging_obj=logging_obj,
model=model,
@@ -1887,11 +1895,12 @@ def client(original_function): # noqa: PLR0915
rules_obj=rules_obj,
)
# Only run if call_type is a valid value in CallTypes
- if call_type in [ct.value for ct in CallTypes]:
+ _call_type_enum = _CALL_TYPE_ENUM_MAP.get(call_type)
+ if _call_type_enum is not None:
result = await async_post_call_success_deployment_hook(
request_data=kwargs,
response=result,
- call_type=CallTypes(call_type),
+ call_type=_call_type_enum,
)
## Add response to cache
@@ -1931,10 +1940,7 @@ def client(original_function): # noqa: PLR0915
end_time=end_time,
)
- update_response_metadata = getattr(
- sys.modules[__name__], "update_response_metadata"
- )
- update_response_metadata(
+ _update_response_metadata(
result=result,
logging_obj=logging_obj,
model=model,
@@ -2644,7 +2650,14 @@ def get_supported_regions(
model=model, custom_llm_provider=custom_llm_provider
)
- supported_regions = model_info.get("supported_regions", None)
+ # Get the key used in model_cost to look up supported_regions
+ # since ModelInfoBase doesn't include this field
+ model_key = model_info.get("key")
+ if model_key is None:
+ return None
+
+ model_cost_data = litellm.model_cost.get(model_key, {})
+ supported_regions = model_cost_data.get("supported_regions", None)
if supported_regions is None:
return None
@@ -3242,7 +3255,7 @@ def get_optional_params_embeddings( # noqa: PLR0915
object = litellm.AmazonTitanMultimodalEmbeddingG1Config()
elif "amazon.titan-embed-text-v2:0" in model:
object = litellm.AmazonTitanV2Config()
- elif "cohere.embed-multilingual-v3" in model:
+ elif "cohere.embed-multilingual-v3" in model or "cohere.embed-v4" in model:
object = litellm.BedrockCohereEmbeddingConfig()
elif "twelvelabs" in model or "marengo" in model:
object = litellm.TwelveLabsMarengoEmbeddingConfig()
@@ -7793,6 +7806,7 @@ class ProviderConfigManager:
# Simple provider mappings (no model parameter needed)
LlmProviders.DEEPSEEK: (lambda: litellm.DeepSeekChatConfig(), False),
LlmProviders.GROQ: (lambda: litellm.GroqChatConfig(), False),
+ LlmProviders.A2A: (lambda: litellm.A2AConfig(), False),
LlmProviders.BYTEZ: (lambda: litellm.BytezChatConfig(), False),
LlmProviders.DATABRICKS: (lambda: litellm.DatabricksConfig(), False),
LlmProviders.XAI: (lambda: litellm.XAIChatConfig(), False),
diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json
index 0f84bba941d..0da47634a94 100644
--- a/model_prices_and_context_window.json
+++ b/model_prices_and_context_window.json
@@ -744,12 +744,13 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_streaming": true
},
"anthropic.claude-3-5-sonnet-20240620-v1:0": {
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
@@ -758,14 +759,22 @@
"supports_pdf_input": true,
"supports_response_schema": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "input_cost_per_token_above_200k_tokens": 6e-06,
+ "output_cost_per_token_above_200k_tokens": 3e-05,
+ "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
+ "cache_read_input_token_cost_above_200k_tokens": 6e-07,
+ "cache_creation_input_token_cost_above_1hr": 7.5e-06,
+ "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.5e-05,
+ "cache_creation_input_token_cost": 3.75e-06,
+ "cache_read_input_token_cost": 3e-07
},
"anthropic.claude-3-5-sonnet-20241022-v2:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
@@ -777,7 +786,13 @@
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "input_cost_per_token_above_200k_tokens": 6e-06,
+ "output_cost_per_token_above_200k_tokens": 3e-05,
+ "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
+ "cache_read_input_token_cost_above_200k_tokens": 6e-07,
+ "cache_creation_input_token_cost_above_1hr": 7.5e-06,
+ "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.5e-05
},
"anthropic.claude-3-7-sonnet-20240620-v1:0": {
"cache_creation_input_token_cost": 4.5e-06,
@@ -948,6 +963,306 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 159
},
+ "anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "global.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "global.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "us.anthropic.claude-opus-4-6-v1:0": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "us.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "eu.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "eu.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "apac.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "apac.anthropic.claude-opus-4-6-v1": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "bedrock_converse",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
"anthropic.claude-sonnet-4-20250514-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
@@ -1429,6 +1744,33 @@
"supports_tool_choice": true,
"supports_vision": true
},
+ "azure_ai/claude-opus-4-6": {
+ "input_cost_per_token": 5e-06,
+ "output_cost_per_token": 2.5e-05,
+ "litellm_provider": "azure_ai",
+ "max_input_tokens": 200000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_1hr": 1e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 159
+ },
"azure_ai/claude-opus-4-1": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_creation_input_token_cost_above_1hr": 3e-05,
@@ -6715,13 +7057,13 @@
"supports_tool_choice": true
},
"cerebras/gpt-oss-120b": {
- "input_cost_per_token": 2.5e-07,
+ "input_cost_per_token": 3.5e-07,
"litellm_provider": "cerebras",
"max_input_tokens": 131072,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 6.9e-07,
+ "output_cost_per_token": 7.5e-07,
"source": "https://www.cerebras.ai/blog/openai-gpt-oss-120b-runs-fastest-on-cerebras",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
@@ -6739,6 +7081,7 @@
"output_cost_per_token": 8e-07,
"source": "https://inference-docs.cerebras.ai/support/pricing",
"supports_function_calling": true,
+ "supports_reasoning": true,
"supports_tool_choice": true
},
"cerebras/zai-glm-4.6": {
@@ -7439,6 +7782,130 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 159
},
+ "claude-opus-4-6": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_creation_input_token_cost_above_1hr": 1e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "anthropic",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "us/claude-opus-4-6": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_creation_input_token_cost_above_1hr": 1.1e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "anthropic",
+ "max_input_tokens": 200000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "claude-opus-4-6-20260205": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_creation_input_token_cost_above_1hr": 1e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "anthropic",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "us/claude-opus-4-6-20260205": {
+ "cache_creation_input_token_cost": 6.875e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
+ "cache_creation_input_token_cost_above_1hr": 1.1e-05,
+ "cache_read_input_token_cost": 5.5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
+ "input_cost_per_token": 5.5e-06,
+ "input_cost_per_token_above_200k_tokens": 1.1e-05,
+ "litellm_provider": "anthropic",
+ "max_input_tokens": 200000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.75e-05,
+ "output_cost_per_token_above_200k_tokens": 4.125e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
+ },
"claude-sonnet-4-20250514": {
"deprecation_date": "2026-05-14",
"cache_creation_input_token_cost": 3.75e-06,
@@ -10559,6 +11026,32 @@
"/v1/audio/transcriptions"
]
},
+ "elevenlabs/eleven_v3": {
+ "input_cost_per_character": 0.00018,
+ "litellm_provider": "elevenlabs",
+ "metadata": {
+ "calculation": "$0.18/1000 characters (Scale plan pricing, 1 credit per character)",
+ "notes": "ElevenLabs Eleven v3 - most expressive TTS model with 70+ languages and audio tags support"
+ },
+ "mode": "audio_speech",
+ "source": "https://elevenlabs.io/pricing",
+ "supported_endpoints": [
+ "/v1/audio/speech"
+ ]
+ },
+ "elevenlabs/eleven_multilingual_v2": {
+ "input_cost_per_character": 0.00018,
+ "litellm_provider": "elevenlabs",
+ "metadata": {
+ "calculation": "$0.18/1000 characters (Scale plan pricing, 1 credit per character)",
+ "notes": "ElevenLabs Eleven Multilingual v2 - default TTS model with 29 languages support"
+ },
+ "mode": "audio_speech",
+ "source": "https://elevenlabs.io/pricing",
+ "supported_endpoints": [
+ "/v1/audio/speech"
+ ]
+ },
"embed-english-light-v2.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "cohere",
@@ -12835,6 +13328,40 @@
"supports_vision": true,
"supports_web_search": true
},
+ "deep-research-pro-preview-12-2025": {
+ "input_cost_per_image": 0.0011,
+ "input_cost_per_token": 2e-06,
+ "input_cost_per_token_batches": 1e-06,
+ "litellm_provider": "vertex_ai-language-models",
+ "max_input_tokens": 65536,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "image_generation",
+ "output_cost_per_image": 0.134,
+ "output_cost_per_image_token": 0.00012,
+ "output_cost_per_token": 1.2e-05,
+ "output_cost_per_token_batches": 6e-06,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/completions",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "image"
+ ],
+ "supports_function_calling": false,
+ "supports_prompt_caching": true,
+ "supports_response_schema": true,
+ "supports_system_messages": true,
+ "supports_vision": true,
+ "supports_web_search": true
+ },
"gemini-2.5-flash-lite": {
"cache_read_input_token_cost": 1e-08,
"input_cost_per_audio_token": 3e-07,
@@ -13289,7 +13816,8 @@
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_native_streaming": true
},
"vertex_ai/gemini-3-pro-preview": {
"cache_read_input_token_cost": 2e-07,
@@ -13337,7 +13865,8 @@
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_native_streaming": true
},
"vertex_ai/gemini-3-flash-preview": {
"cache_read_input_token_cost": 5e-08,
@@ -13380,7 +13909,8 @@
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_native_streaming": true
},
"gemini-2.5-pro-exp-03-25": {
"cache_read_input_token_cost": 1.25e-07,
@@ -14747,6 +15277,42 @@
"supports_vision": true,
"supports_web_search": true
},
+ "gemini/deep-research-pro-preview-12-2025": {
+ "input_cost_per_image": 0.0011,
+ "input_cost_per_token": 2e-06,
+ "input_cost_per_token_batches": 1e-06,
+ "litellm_provider": "gemini",
+ "max_input_tokens": 65536,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "image_generation",
+ "output_cost_per_image": 0.134,
+ "output_cost_per_image_token": 0.00012,
+ "output_cost_per_token": 1.2e-05,
+ "rpm": 1000,
+ "tpm": 4000000,
+ "output_cost_per_token_batches": 6e-06,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/completions",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "image"
+ ],
+ "supports_function_calling": false,
+ "supports_prompt_caching": true,
+ "supports_response_schema": true,
+ "supports_system_messages": true,
+ "supports_vision": true,
+ "supports_web_search": true
+ },
"gemini/gemini-2.5-flash-lite": {
"cache_read_input_token_cost": 1e-08,
"input_cost_per_audio_token": 3e-07,
@@ -15331,6 +15897,7 @@
"supports_url_context": true,
"supports_vision": true,
"supports_web_search": true,
+ "supports_native_streaming": true,
"tpm": 800000
},
"gemini-3-flash-preview": {
@@ -15376,7 +15943,8 @@
"supports_tool_choice": true,
"supports_url_context": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_native_streaming": true
},
"gemini/gemini-2.5-pro-exp-03-25": {
"cache_read_input_token_cost": 0.0,
@@ -21473,6 +22041,20 @@
"supports_tool_choice": true,
"supports_web_search": true
},
+ "moonshot/kimi-k2.5": {
+ "cache_read_input_token_cost": 1e-07,
+ "input_cost_per_token": 6e-07,
+ "litellm_provider": "moonshot",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 3e-06,
+ "source": "https://platform.moonshot.ai/docs/pricing/chat",
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_vision": true
+ },
"moonshot/kimi-latest": {
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 2e-06,
@@ -24314,6 +24896,31 @@
"supports_tool_choice": true,
"supports_function_calling": true
},
+ "openrouter/qwen/qwen3-235b-a22b-2507": {
+ "input_cost_per_token": 7.1e-08,
+ "litellm_provider": "openrouter",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 1e-07,
+ "source": "https://openrouter.ai/qwen/qwen3-235b-a22b-2507",
+ "supports_function_calling": true,
+ "supports_tool_choice": true
+ },
+ "openrouter/qwen/qwen3-235b-a22b-thinking-2507": {
+ "input_cost_per_token": 1.1e-07,
+ "litellm_provider": "openrouter",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 262144,
+ "max_tokens": 262144,
+ "mode": "chat",
+ "output_cost_per_token": 6e-07,
+ "source": "https://openrouter.ai/qwen/qwen3-235b-a22b-thinking-2507",
+ "supports_function_calling": true,
+ "supports_reasoning": true,
+ "supports_tool_choice": true
+ },
"openrouter/switchpoint/router": {
"input_cost_per_token": 8.5e-07,
"litellm_provider": "openrouter",
@@ -24390,21 +24997,21 @@
"supports_tool_choice": true
},
"openrouter/xiaomi/mimo-v2-flash": {
- "input_cost_per_token": 9e-08,
- "output_cost_per_token": 2.9e-07,
- "cache_creation_input_token_cost": 0.0,
- "cache_read_input_token_cost": 0.0,
- "litellm_provider": "openrouter",
- "max_input_tokens": 262144,
- "max_output_tokens": 16384,
- "max_tokens": 16384,
- "mode": "chat",
- "supports_function_calling": true,
- "supports_tool_choice": true,
- "supports_reasoning": true,
- "supports_vision": false,
- "supports_prompt_caching": false
- },
+ "input_cost_per_token": 9e-08,
+ "output_cost_per_token": 2.9e-07,
+ "cache_creation_input_token_cost": 0.0,
+ "cache_read_input_token_cost": 0.0,
+ "litellm_provider": "openrouter",
+ "max_input_tokens": 262144,
+ "max_output_tokens": 16384,
+ "max_tokens": 16384,
+ "mode": "chat",
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_reasoning": true,
+ "supports_vision": false,
+ "supports_prompt_caching": false
+ },
"openrouter/z-ai/glm-4.7": {
"input_cost_per_token": 4e-07,
"output_cost_per_token": 1.5e-06,
@@ -26319,13 +26926,13 @@
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"mode": "image_edit",
- "output_cost_per_image": 0.40
+ "output_cost_per_image": 0.4
},
"stability.stable-creative-upscale-v1:0": {
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"mode": "image_edit",
- "output_cost_per_image": 0.60
+ "output_cost_per_image": 0.6
},
"stability.stable-fast-upscale-v1:0": {
"litellm_provider": "bedrock",
@@ -27084,6 +27691,34 @@
"supports_reasoning": true,
"supports_tool_choice": true
},
+ "together_ai/zai-org/GLM-4.7": {
+ "input_cost_per_token": 4.5e-07,
+ "litellm_provider": "together_ai",
+ "max_input_tokens": 200000,
+ "max_output_tokens": 200000,
+ "max_tokens": 200000,
+ "mode": "chat",
+ "output_cost_per_token": 2e-06,
+ "source": "https://www.together.ai/models/glm-4-7",
+ "supports_function_calling": true,
+ "supports_parallel_function_calling": true,
+ "supports_reasoning": true,
+ "supports_tool_choice": true
+ },
+ "together_ai/moonshotai/Kimi-K2.5": {
+ "input_cost_per_token": 5e-07,
+ "litellm_provider": "together_ai",
+ "max_input_tokens": 256000,
+ "max_output_tokens": 256000,
+ "max_tokens": 256000,
+ "mode": "chat",
+ "output_cost_per_token": 2.8e-06,
+ "source": "https://www.together.ai/models/kimi-k2-5",
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_reasoning": true
+ },
"together_ai/moonshotai/Kimi-K2-Instruct-0905": {
"input_cost_per_token": 1e-06,
"litellm_provider": "together_ai",
@@ -27800,7 +28435,9 @@
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
- "output_cost_per_token": 3e-07
+ "output_cost_per_token": 3e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/alibaba/qwen3-coder": {
"input_cost_per_token": 4e-07,
@@ -27809,7 +28446,9 @@
"max_output_tokens": 66536,
"max_tokens": 66536,
"mode": "chat",
- "output_cost_per_token": 1.6e-06
+ "output_cost_per_token": 1.6e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/amazon/nova-lite": {
"input_cost_per_token": 6e-08,
@@ -27818,7 +28457,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 2.4e-07
+ "output_cost_per_token": 2.4e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/amazon/nova-micro": {
"input_cost_per_token": 3.5e-08,
@@ -27827,7 +28469,9 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 1.4e-07
+ "output_cost_per_token": 1.4e-07,
+ "supports_function_calling": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/amazon/nova-pro": {
"input_cost_per_token": 8e-07,
@@ -27836,7 +28480,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 3.2e-06
+ "output_cost_per_token": 3.2e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/amazon/titan-embed-text-v2": {
"input_cost_per_token": 2e-08,
@@ -27856,7 +28503,11 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 1.25e-06
+ "output_cost_per_token": 1.25e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-3-opus": {
"cache_creation_input_token_cost": 1.875e-05,
@@ -27867,7 +28518,11 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 7.5e-05
+ "output_cost_per_token": 7.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-3.5-haiku": {
"cache_creation_input_token_cost": 1e-06,
@@ -27878,7 +28533,11 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 4e-06
+ "output_cost_per_token": 4e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-3.5-sonnet": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -27889,7 +28548,11 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-3.7-sonnet": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -27900,7 +28563,11 @@
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-4-opus": {
"cache_creation_input_token_cost": 1.875e-05,
@@ -27911,7 +28578,11 @@
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
- "output_cost_per_token": 7.5e-05
+ "output_cost_per_token": 7.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/anthropic/claude-4-sonnet": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -27922,7 +28593,9 @@
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/cohere/command-a": {
"input_cost_per_token": 2.5e-06,
@@ -27931,7 +28604,9 @@
"max_output_tokens": 8000,
"max_tokens": 8000,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/cohere/command-r": {
"input_cost_per_token": 1.5e-07,
@@ -27940,7 +28615,9 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 6e-07
+ "output_cost_per_token": 6e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/cohere/command-r-plus": {
"input_cost_per_token": 2.5e-06,
@@ -27949,7 +28626,9 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/cohere/embed-v4.0": {
"input_cost_per_token": 1.2e-07,
@@ -27967,7 +28646,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 2.19e-06
+ "output_cost_per_token": 2.19e-06,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/deepseek/deepseek-r1-distill-llama-70b": {
"input_cost_per_token": 7.5e-07,
@@ -27976,7 +28656,10 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 9.9e-07
+ "output_cost_per_token": 9.9e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/deepseek/deepseek-v3": {
"input_cost_per_token": 9e-07,
@@ -27985,7 +28668,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 9e-07
+ "output_cost_per_token": 9e-07,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/google/gemini-2.0-flash": {
"deprecation_date": "2026-03-31",
@@ -27995,7 +28679,11 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 6e-07
+ "output_cost_per_token": 6e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/google/gemini-2.0-flash-lite": {
"deprecation_date": "2026-03-31",
@@ -28005,7 +28693,11 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 3e-07
+ "output_cost_per_token": 3e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/google/gemini-2.5-flash": {
"input_cost_per_token": 3e-07,
@@ -28014,7 +28706,11 @@
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
- "output_cost_per_token": 2.5e-06
+ "output_cost_per_token": 2.5e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/google/gemini-2.5-pro": {
"input_cost_per_token": 2.5e-06,
@@ -28023,7 +28719,11 @@
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/google/gemini-embedding-001": {
"input_cost_per_token": 1.5e-07,
@@ -28041,7 +28741,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 2e-07
+ "output_cost_per_token": 2e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/google/text-embedding-005": {
"input_cost_per_token": 2.5e-08,
@@ -28077,7 +28780,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 7.9e-07
+ "output_cost_per_token": 7.9e-07,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-3-8b": {
"input_cost_per_token": 5e-08,
@@ -28086,7 +28790,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 8e-08
+ "output_cost_per_token": 8e-08,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-3.1-70b": {
"input_cost_per_token": 7.2e-07,
@@ -28095,7 +28800,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 7.2e-07
+ "output_cost_per_token": 7.2e-07,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-3.1-8b": {
"input_cost_per_token": 5e-08,
@@ -28104,7 +28810,9 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 8e-08
+ "output_cost_per_token": 8e-08,
+ "supports_function_calling": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/meta/llama-3.2-11b": {
"input_cost_per_token": 1.6e-07,
@@ -28113,7 +28821,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 1.6e-07
+ "output_cost_per_token": 1.6e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-3.2-1b": {
"input_cost_per_token": 1e-07,
@@ -28131,7 +28842,9 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 1.5e-07
+ "output_cost_per_token": 1.5e-07,
+ "supports_function_calling": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/meta/llama-3.2-90b": {
"input_cost_per_token": 7.2e-07,
@@ -28140,7 +28853,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 7.2e-07
+ "output_cost_per_token": 7.2e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-3.3-70b": {
"input_cost_per_token": 7.2e-07,
@@ -28149,7 +28865,9 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 7.2e-07
+ "output_cost_per_token": 7.2e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-4-maverick": {
"input_cost_per_token": 2e-07,
@@ -28158,7 +28876,8 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 6e-07
+ "output_cost_per_token": 6e-07,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/meta/llama-4-scout": {
"input_cost_per_token": 1e-07,
@@ -28167,7 +28886,10 @@
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
- "output_cost_per_token": 3e-07
+ "output_cost_per_token": 3e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/mistral/codestral": {
"input_cost_per_token": 3e-07,
@@ -28176,7 +28898,9 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 9e-07
+ "output_cost_per_token": 9e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/mistral/codestral-embed": {
"input_cost_per_token": 1.5e-07,
@@ -28194,7 +28918,10 @@
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
- "output_cost_per_token": 2.8e-07
+ "output_cost_per_token": 2.8e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/mistral/magistral-medium": {
"input_cost_per_token": 2e-06,
@@ -28203,7 +28930,10 @@
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
- "output_cost_per_token": 5e-06
+ "output_cost_per_token": 5e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/mistral/magistral-small": {
"input_cost_per_token": 5e-07,
@@ -28212,7 +28942,8 @@
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
- "output_cost_per_token": 1.5e-06
+ "output_cost_per_token": 1.5e-06,
+ "supports_function_calling": true
},
"vercel_ai_gateway/mistral/ministral-3b": {
"input_cost_per_token": 4e-08,
@@ -28221,7 +28952,9 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 4e-08
+ "output_cost_per_token": 4e-08,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/mistral/ministral-8b": {
"input_cost_per_token": 1e-07,
@@ -28230,7 +28963,10 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 1e-07
+ "output_cost_per_token": 1e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/mistral/mistral-embed": {
"input_cost_per_token": 1e-07,
@@ -28248,7 +28984,9 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 6e-06
+ "output_cost_per_token": 6e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/mistral/mistral-saba-24b": {
"input_cost_per_token": 7.9e-07,
@@ -28266,7 +29004,10 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 3e-07
+ "output_cost_per_token": 3e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/mistral/mixtral-8x22b-instruct": {
"input_cost_per_token": 1.2e-06,
@@ -28275,7 +29016,8 @@
"max_output_tokens": 2048,
"max_tokens": 2048,
"mode": "chat",
- "output_cost_per_token": 1.2e-06
+ "output_cost_per_token": 1.2e-06,
+ "supports_function_calling": true
},
"vercel_ai_gateway/mistral/pixtral-12b": {
"input_cost_per_token": 1.5e-07,
@@ -28284,7 +29026,11 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 1.5e-07
+ "output_cost_per_token": 1.5e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/mistral/pixtral-large": {
"input_cost_per_token": 2e-06,
@@ -28293,7 +29039,11 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 6e-06
+ "output_cost_per_token": 6e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/moonshotai/kimi-k2": {
"input_cost_per_token": 5.5e-07,
@@ -28302,7 +29052,9 @@
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
- "output_cost_per_token": 2.2e-06
+ "output_cost_per_token": 2.2e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/morph/morph-v3-fast": {
"input_cost_per_token": 8e-07,
@@ -28329,7 +29081,9 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 1.5e-06
+ "output_cost_per_token": 1.5e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/openai/gpt-3.5-turbo-instruct": {
"input_cost_per_token": 1.5e-06,
@@ -28347,7 +29101,10 @@
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
- "output_cost_per_token": 3e-05
+ "output_cost_per_token": 3e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/openai/gpt-4.1": {
"cache_creation_input_token_cost": 0.0,
@@ -28358,7 +29115,11 @@
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 8e-06
+ "output_cost_per_token": 8e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/gpt-4.1-mini": {
"cache_creation_input_token_cost": 0.0,
@@ -28369,7 +29130,11 @@
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 1.6e-06
+ "output_cost_per_token": 1.6e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/gpt-4.1-nano": {
"cache_creation_input_token_cost": 0.0,
@@ -28380,7 +29145,11 @@
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 4e-07
+ "output_cost_per_token": 4e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/gpt-4o": {
"cache_creation_input_token_cost": 0.0,
@@ -28391,7 +29160,11 @@
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/gpt-4o-mini": {
"cache_creation_input_token_cost": 0.0,
@@ -28402,7 +29175,11 @@
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
- "output_cost_per_token": 6e-07
+ "output_cost_per_token": 6e-07,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/o1": {
"cache_creation_input_token_cost": 0.0,
@@ -28413,7 +29190,11 @@
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
- "output_cost_per_token": 6e-05
+ "output_cost_per_token": 6e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/o3": {
"cache_creation_input_token_cost": 0.0,
@@ -28424,7 +29205,11 @@
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
- "output_cost_per_token": 8e-06
+ "output_cost_per_token": 8e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/o3-mini": {
"cache_creation_input_token_cost": 0.0,
@@ -28435,7 +29220,10 @@
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
- "output_cost_per_token": 4.4e-06
+ "output_cost_per_token": 4.4e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/o4-mini": {
"cache_creation_input_token_cost": 0.0,
@@ -28446,7 +29234,11 @@
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
- "output_cost_per_token": 4.4e-06
+ "output_cost_per_token": 4.4e-06,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true,
+ "supports_response_schema": true
},
"vercel_ai_gateway/openai/text-embedding-3-large": {
"input_cost_per_token": 1.3e-07,
@@ -28518,7 +29310,10 @@
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/vercel/v0-1.5-md": {
"input_cost_per_token": 3e-06,
@@ -28527,7 +29322,10 @@
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-2": {
"input_cost_per_token": 2e-06,
@@ -28536,7 +29334,9 @@
"max_output_tokens": 4000,
"max_tokens": 4000,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-2-vision": {
"input_cost_per_token": 2e-06,
@@ -28545,7 +29345,10 @@
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
- "output_cost_per_token": 1e-05
+ "output_cost_per_token": 1e-05,
+ "supports_vision": true,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-3": {
"input_cost_per_token": 3e-06,
@@ -28554,7 +29357,9 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-3-fast": {
"input_cost_per_token": 5e-06,
@@ -28563,7 +29368,8 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 2.5e-05
+ "output_cost_per_token": 2.5e-05,
+ "supports_function_calling": true
},
"vercel_ai_gateway/xai/grok-3-mini": {
"input_cost_per_token": 3e-07,
@@ -28572,7 +29378,9 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 5e-07
+ "output_cost_per_token": 5e-07,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-3-mini-fast": {
"input_cost_per_token": 6e-07,
@@ -28581,7 +29389,9 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 4e-06
+ "output_cost_per_token": 4e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/xai/grok-4": {
"input_cost_per_token": 3e-06,
@@ -28590,7 +29400,9 @@
"max_output_tokens": 256000,
"max_tokens": 256000,
"mode": "chat",
- "output_cost_per_token": 1.5e-05
+ "output_cost_per_token": 1.5e-05,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/zai/glm-4.5": {
"input_cost_per_token": 6e-07,
@@ -28599,7 +29411,9 @@
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
- "output_cost_per_token": 2.2e-06
+ "output_cost_per_token": 2.2e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/zai/glm-4.5-air": {
"input_cost_per_token": 2e-07,
@@ -28608,7 +29422,9 @@
"max_output_tokens": 96000,
"max_tokens": 96000,
"mode": "chat",
- "output_cost_per_token": 1.1e-06
+ "output_cost_per_token": 1.1e-06,
+ "supports_function_calling": true,
+ "supports_tool_choice": true
},
"vercel_ai_gateway/zai/glm-4.6": {
"litellm_provider": "vercel_ai_gateway",
@@ -28676,7 +29492,9 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_streaming": true,
+ "supports_vision": true
},
"vertex_ai/claude-3-5-sonnet": {
"input_cost_per_token": 3e-06,
@@ -28947,7 +29765,38 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_streaming": true
+ },
+ "vertex_ai/claude-opus-4-6": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_200k_tokens": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_200k_tokens": 1e-05,
+ "litellm_provider": "vertex_ai-anthropic_models",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "output_cost_per_token_above_200k_tokens": 3.75e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "tool_use_system_prompt_tokens": 346
},
"vertex_ai/claude-sonnet-4-5": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -28999,7 +29848,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_native_streaming": true
},
"vertex_ai/claude-opus-4@20250514": {
"cache_creation_input_token_cost": 1.875e-05,
@@ -29281,6 +30131,21 @@
"output_cost_per_token_batches": 6e-06,
"source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
},
+ "vertex_ai/deep-research-pro-preview-12-2025": {
+ "input_cost_per_image": 0.0011,
+ "input_cost_per_token": 2e-06,
+ "input_cost_per_token_batches": 1e-06,
+ "litellm_provider": "vertex_ai-language-models",
+ "max_input_tokens": 65536,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "image_generation",
+ "output_cost_per_image": 0.134,
+ "output_cost_per_image_token": 0.00012,
+ "output_cost_per_token": 1.2e-05,
+ "output_cost_per_token_batches": 6e-06,
+ "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
+ },
"vertex_ai/imagegeneration@006": {
"litellm_provider": "vertex_ai-image-models",
"mode": "image_generation",
@@ -29770,6 +30635,9 @@
"mode": "chat",
"output_cost_per_token": 1e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
+ "supported_regions": [
+ "global"
+ ],
"supports_function_calling": true,
"supports_tool_choice": true
},
@@ -29782,6 +30650,9 @@
"mode": "chat",
"output_cost_per_token": 4e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
+ "supported_regions": [
+ "global"
+ ],
"supports_function_calling": true,
"supports_tool_choice": true
},
@@ -29794,6 +30665,9 @@
"mode": "chat",
"output_cost_per_token": 1.2e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
+ "supported_regions": [
+ "global"
+ ],
"supports_function_calling": true,
"supports_tool_choice": true
},
@@ -29806,6 +30680,9 @@
"mode": "chat",
"output_cost_per_token": 1.2e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
+ "supported_regions": [
+ "global"
+ ],
"supports_function_calling": true,
"supports_tool_choice": true
},
@@ -34754,4 +35631,4 @@
"output_cost_per_token": 0,
"supports_reasoning": true
}
-}
\ No newline at end of file
+}
diff --git a/poetry.lock b/poetry.lock
index 537367c5aa0..b37fd863431 100644
--- a/poetry.lock
+++ b/poetry.lock
@@ -1,4 +1,4 @@
-# This file is automatically @generated by Poetry 2.2.0 and should not be changed by hand.
+# This file is automatically @generated by Poetry 2.3.2 and should not be changed by hand.
[[package]]
name = "a2a-sdk"
@@ -398,6 +398,7 @@ files = [
{file = "azure_core-1.36.0-py3-none-any.whl", hash = "sha256:fee9923a3a753e94a259563429f3644aaf05c486d45b1215d098115102d91d3b"},
{file = "azure_core-1.36.0.tar.gz", hash = "sha256:22e5605e6d0bf1d229726af56d9e92bc37b6e726b141a18be0b4d424131741b7"},
]
+markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
requests = ">=2.21.0"
@@ -418,6 +419,7 @@ files = [
{file = "azure_identity-1.25.1-py3-none-any.whl", hash = "sha256:e9edd720af03dff020223cd269fa3a61e8f345ea75443858273bcb44844ab651"},
{file = "azure_identity-1.25.1.tar.gz", hash = "sha256:87ca8328883de6036443e1c37b40e8dc8fb74898240f61071e09d2e369361456"},
]
+markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
azure-core = ">=1.31.0"
@@ -718,11 +720,23 @@ files = [
{file = "cffi-2.0.0-cp39-cp39-win_amd64.whl", hash = "sha256:b882b3df248017dba09d6b16defe9b5c407fe32fc7c65a9c69798e6175601be9"},
{file = "cffi-2.0.0.tar.gz", hash = "sha256:44d1b5909021139fe36001ae048dbdde8214afa20200eda0f64c068cac5d5529"},
]
-markers = {main = "platform_python_implementation != \"PyPy\" or extra == \"proxy\"", dev = "platform_python_implementation != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\""}
+markers = {main = "(platform_python_implementation != \"PyPy\" or extra == \"proxy\") and (python_version >= \"3.10\" or extra == \"proxy\" or extra == \"extra-proxy\") and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\")", dev = "platform_python_implementation != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\""}
[package.dependencies]
pycparser = {version = "*", markers = "implementation_name != \"PyPy\""}
+[[package]]
+name = "chardet"
+version = "5.2.0"
+description = "Universal encoding detector for Python 3"
+optional = false
+python-versions = ">=3.7"
+groups = ["dev"]
+files = [
+ {file = "chardet-5.2.0-py3-none-any.whl", hash = "sha256:e1cf59446890a00105fe7b7912492ea04b6e6f06d4b742b2c788469e34c82970"},
+ {file = "chardet-5.2.0.tar.gz", hash = "sha256:1b3b6ff479a8c414bc3fa2c0852995695c4a026dcd6d0633b2dd092ca39c1cf7"},
+]
+
[[package]]
name = "charset-normalizer"
version = "3.4.4"
@@ -1137,7 +1151,6 @@ description = "cryptography is a package which provides cryptographic recipes an
optional = false
python-versions = ">=3.7"
groups = ["main", "dev", "proxy-dev"]
-markers = "python_version == \"3.9\""
files = [
{file = "cryptography-43.0.3-cp37-abi3-macosx_10_9_universal2.whl", hash = "sha256:bf7a1932ac4176486eab36a19ed4c0492da5d97123f1406cf15e41b05e787d2e"},
{file = "cryptography-43.0.3-cp37-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:63efa177ff54aec6e1c0aefaa1a241232dcd37413835a9b674b6e3f0ae2bfd3e"},
@@ -1167,6 +1180,7 @@ files = [
{file = "cryptography-43.0.3-pp39-pypy39_pp73-win_amd64.whl", hash = "sha256:2ce6fae5bdad59577b44e4dfed356944fbf1d925269114c28be377692643b4ff"},
{file = "cryptography-43.0.3.tar.gz", hash = "sha256:315b9001266a492a6ff443b61238f956b214dbec9910a081ba5b6646a055a805"},
]
+markers = {main = "python_version == \"3.9\" and (extra == \"proxy\" or extra == \"extra-proxy\")", dev = "python_version == \"3.9\"", proxy-dev = "python_version == \"3.9\""}
[package.dependencies]
cffi = {version = ">=1.12", markers = "platform_python_implementation != \"PyPy\""}
@@ -1188,7 +1202,6 @@ description = "cryptography is a package which provides cryptographic recipes an
optional = false
python-versions = "!=3.9.0,!=3.9.1,>=3.8"
groups = ["main", "dev", "proxy-dev"]
-markers = "python_version >= \"3.10\""
files = [
{file = "cryptography-46.0.3-cp311-abi3-macosx_10_9_universal2.whl", hash = "sha256:109d4ddfadf17e8e7779c39f9b18111a09efb969a301a31e987416a0191ed93a"},
{file = "cryptography-46.0.3-cp311-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:09859af8466b69bc3c27bdf4f5d84a665e0f7ab5088412e9e2ec49758eca5cbc"},
@@ -1245,6 +1258,7 @@ files = [
{file = "cryptography-46.0.3-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:6b5063083824e5509fdba180721d55909ffacccc8adbec85268b48439423d78c"},
{file = "cryptography-46.0.3.tar.gz", hash = "sha256:a8b17438104fed022ce745b362294d9ce35b4c2e45c1d958ad4a4b019285f4a1"},
]
+markers = {main = "python_version >= \"3.10\" and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\")", dev = "python_version >= \"3.10\"", proxy-dev = "python_version >= \"3.10\""}
[package.dependencies]
cffi = {version = ">=2.0.0", markers = "python_full_version >= \"3.9.0\" and platform_python_implementation != \"PyPy\""}
@@ -1300,6 +1314,27 @@ dev = ["autoflake", "black", "build", "databricks-connect", "httpx", "ipython",
notebook = ["ipython (>=8,<10)", "ipywidgets (>=8,<9)"]
openai = ["httpx", "langchain-openai ; python_version > \"3.7\"", "openai"]
+[[package]]
+name = "diff-cover"
+version = "9.7.2"
+description = "Run coverage and linting reports on diffs"
+optional = false
+python-versions = ">=3.9"
+groups = ["dev"]
+files = [
+ {file = "diff_cover-9.7.2-py3-none-any.whl", hash = "sha256:cd6498620c747c2493a6c83c14362c32868bfd91cd8d0dd093f136070ec4ffc5"},
+ {file = "diff_cover-9.7.2.tar.gz", hash = "sha256:872c820d2ecbf79c61d52c7dc70419015e0ab9289589566c791dd270fc0c6e3b"},
+]
+
+[package.dependencies]
+chardet = ">=3.0.0"
+Jinja2 = ">=2.7.1"
+pluggy = ">=0.13.1,<2"
+Pygments = ">=2.19.1,<3.0.0"
+
+[package.extras]
+toml = ["tomli (>=1.2.1)"]
+
[[package]]
name = "diskcache"
version = "5.6.3"
@@ -2242,11 +2277,11 @@ files = [
]
[package.dependencies]
-google-api-core = {version = ">=1.34.1,<2.0.dev0 || >=2.11.dev0,<3.0.0dev", extras = ["grpc"]}
-google-auth = ">=2.14.1,<2.24.0 || >2.24.0,<2.25.0 || >2.25.0,<3.0.0dev"
-grpc-google-iam-v1 = ">=0.12.4,<1.0.0dev"
-proto-plus = ">=1.22.3,<2.0.0dev"
-protobuf = ">=3.20.2,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<6.0.0dev"
+google-api-core = {version = ">=1.34.1,<2.0.dev0 || >=2.11.dev0,<3.0.0.dev0", extras = ["grpc"]}
+google-auth = ">=2.14.1,<2.24.0 || >2.24.0,<2.25.0 || >2.25.0,<3.0.0.dev0"
+grpc-google-iam-v1 = ">=0.12.4,<1.0.0.dev0"
+proto-plus = ">=1.22.3,<2.0.0.dev0"
+protobuf = ">=3.20.2,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<6.0.0.dev0"
[[package]]
name = "google-cloud-resource-manager"
@@ -3085,7 +3120,7 @@ version = "3.1.6"
description = "A very fast and expressive template engine."
optional = false
python-versions = ">=3.7"
-groups = ["main", "proxy-dev"]
+groups = ["main", "dev", "proxy-dev"]
files = [
{file = "jinja2-3.1.6-py3-none-any.whl", hash = "sha256:85ece4451f492d0c13c5dd7c13a64681a86afae63a5f347908daf103ce6d2f67"},
{file = "jinja2-3.1.6.tar.gz", hash = "sha256:0137fb05990d35f1275a587e9aee6d56da821fc83491a0fb838183be43f66d6d"},
@@ -3249,7 +3284,7 @@ files = [
[package.dependencies]
attrs = ">=22.2.0"
-jsonschema-specifications = ">=2023.03.6"
+jsonschema-specifications = ">=2023.3.6"
referencing = ">=0.28.4"
rpds-py = ">=0.7.1"
@@ -3426,15 +3461,15 @@ files = [
[[package]]
name = "litellm-proxy-extras"
-version = "0.4.29"
+version = "0.4.30"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
optional = true
python-versions = "!=2.7.*,!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,!=3.7.*,>=3.8"
groups = ["main"]
markers = "extra == \"proxy\""
files = [
- {file = "litellm_proxy_extras-0.4.29-py3-none-any.whl", hash = "sha256:c36c1b69675c61acccc6b61dd610eb37daeb72c6fd819461cefb5b0cc7e0550f"},
- {file = "litellm_proxy_extras-0.4.29.tar.gz", hash = "sha256:1a8266911e0546f1e17e6714ca20b72e9fef47c1683f9c16399cf2d1786437a0"},
+ {file = "litellm_proxy_extras-0.4.30-py3-none-any.whl", hash = "sha256:0b7df68f0968eb817462b847eaee81bba23d935adb2e84d2e342a77711887051"},
+ {file = "litellm_proxy_extras-0.4.30.tar.gz", hash = "sha256:5d32f8dc3d37d36fb15ab6995fea706dd8a453ff7f12e70b47cba35e5368da10"},
]
[[package]]
@@ -3515,7 +3550,7 @@ version = "3.0.3"
description = "Safely add untrusted strings to HTML/XML markup."
optional = false
python-versions = ">=3.9"
-groups = ["main", "proxy-dev"]
+groups = ["main", "dev", "proxy-dev"]
files = [
{file = "markupsafe-3.0.3-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:2f981d352f04553a7171b8e44369f2af4055f888dfb147d55e42d29e29e74559"},
{file = "markupsafe-3.0.3-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:e1c1493fb6e50ab01d20a22826e57520f1284df32f2d8601fdd90b6304601419"},
@@ -3913,6 +3948,7 @@ files = [
{file = "msal-1.34.0-py3-none-any.whl", hash = "sha256:f669b1644e4950115da7a176441b0e13ec2975c29528d8b9e81316023676d6e1"},
{file = "msal-1.34.0.tar.gz", hash = "sha256:76ba83b716ea5a6d75b0279c0ac353a0e05b820ca1f6682c0eb7f45190c43c2f"},
]
+markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
cryptography = ">=2.5,<49"
@@ -3933,6 +3969,7 @@ files = [
{file = "msal_extensions-1.3.1-py3-none-any.whl", hash = "sha256:96d3de4d034504e969ac5e85bae8106c8373b5c6568e4c8fa7af2eca9dbe6bca"},
{file = "msal_extensions-1.3.1.tar.gz", hash = "sha256:c5b0fd10f65ef62b5f1d62f4251d51cbcaf003fcedae8c91b040a488614be1a4"},
]
+markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
msal = ">=1.29,<2"
@@ -4183,6 +4220,7 @@ files = [
{file = "nodeenv-1.9.1-py2.py3-none-any.whl", hash = "sha256:ba11c9782d29c27c70ffbdda2d7415098754709be8a7056d79a737cd901155c9"},
{file = "nodeenv-1.9.1.tar.gz", hash = "sha256:6ec12890a2dab7946721edbfbcd91f3319c6ccc9aec47be7c7e6b7011ee6645f"},
]
+markers = {main = "extra == \"extra-proxy\""}
[[package]]
name = "numpy"
@@ -4390,7 +4428,7 @@ files = [
{file = "opentelemetry_api-1.39.1-py3-none-any.whl", hash = "sha256:2edd8463432a7f8443edce90972169b195e7d6a05500cd29e6d13898187c9950"},
{file = "opentelemetry_api-1.39.1.tar.gz", hash = "sha256:fbde8c80e1b937a2c61f20347e91c0c18a1940cecf012d62e65a7caf08967c9c"},
]
-markers = {main = "python_version >= \"3.10\""}
+markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""}
[package.dependencies]
importlib-metadata = ">=6.0,<8.8.0"
@@ -4505,7 +4543,7 @@ files = [
{file = "opentelemetry_sdk-1.39.1-py3-none-any.whl", hash = "sha256:4d5482c478513ecb0a5d938dcc61394e647066e0cc2676bee9f3af3f3f45f01c"},
{file = "opentelemetry_sdk-1.39.1.tar.gz", hash = "sha256:cf4d4563caf7bff906c9f7967e2be22d0d6b349b908be0d90fb21c8e9c995cc6"},
]
-markers = {main = "python_version >= \"3.10\""}
+markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""}
[package.dependencies]
opentelemetry-api = "1.39.1"
@@ -4523,7 +4561,7 @@ files = [
{file = "opentelemetry_semantic_conventions-0.60b1-py3-none-any.whl", hash = "sha256:9fa8c8b0c110da289809292b0591220d3a7b53c1526a23021e977d68597893fb"},
{file = "opentelemetry_semantic_conventions-0.60b1.tar.gz", hash = "sha256:87c228b5a0669b748c76d76df6c364c369c28f1c465e50f661e39737e84bc953"},
]
-markers = {main = "python_version >= \"3.10\""}
+markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""}
[package.dependencies]
opentelemetry-api = "1.39.1"
@@ -5000,6 +5038,7 @@ files = [
{file = "prisma-0.11.0-py3-none-any.whl", hash = "sha256:22bb869e59a2968b99f3483bb417717273ffbc569fd1e9ceed95e5614cbaf53a"},
{file = "prisma-0.11.0.tar.gz", hash = "sha256:3f2f2fd2361e1ec5ff655f2a04c7860c2f2a5bc4c91f78ca9c5c6349735bf693"},
]
+markers = {main = "extra == \"extra-proxy\""}
[package.dependencies]
click = ">=7.1.2"
@@ -5316,7 +5355,7 @@ files = [
{file = "pycparser-2.23-py3-none-any.whl", hash = "sha256:e5c6e8d3fbad53479cab09ac03729e0a9faf2bee3db8208a550daf5af81a5934"},
{file = "pycparser-2.23.tar.gz", hash = "sha256:78816d4f24add8f10a06d6f05b4d424ad9e96cfebf68a4ddc99c65c0720d00c2"},
]
-markers = {main = "implementation_name != \"PyPy\" and (platform_python_implementation != \"PyPy\" or extra == \"proxy\")", dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\"", proxy-dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\""}
+markers = {main = "implementation_name != \"PyPy\" and (platform_python_implementation != \"PyPy\" or extra == \"proxy\") and (python_version >= \"3.10\" or extra == \"proxy\" or extra == \"extra-proxy\") and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\")", dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\"", proxy-dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\""}
[[package]]
name = "pydantic"
@@ -5516,14 +5555,14 @@ files = [
name = "pygments"
version = "2.19.2"
description = "Pygments is a syntax highlighting package written in Python."
-optional = true
+optional = false
python-versions = ">=3.8"
-groups = ["main"]
-markers = "extra == \"utils\" or extra == \"proxy\""
+groups = ["main", "dev"]
files = [
{file = "pygments-2.19.2-py3-none-any.whl", hash = "sha256:86540386c03d588bb81d44bc3928634ff26449851e99741617ecb9037ee5ec0b"},
{file = "pygments-2.19.2.tar.gz", hash = "sha256:636cb2477cec7f8952536970bc533bc43743542f70392ae026374600add5b887"},
]
+markers = {main = "extra == \"utils\" or extra == \"proxy\""}
[package.extras]
windows-terminal = ["colorama (>=0.4.6)"]
@@ -5539,6 +5578,7 @@ files = [
{file = "PyJWT-2.10.1-py3-none-any.whl", hash = "sha256:dcdd193e30abefd5debf142f9adfcdd2b58004e644f25406ffaebd50bd98dacb"},
{file = "pyjwt-2.10.1.tar.gz", hash = "sha256:3cc5772eb20009233caf06e9d8a0577824723b44e6648ee0a2aedb6cf9381953"},
]
+markers = {main = "extra == \"extra-proxy\" or extra == \"proxy\""}
[package.dependencies]
cryptography = {version = ">=3.4.0", optional = true, markers = "extra == \"crypto\""}
@@ -6601,29 +6641,29 @@ pyasn1 = ">=0.1.3"
[[package]]
name = "ruff"
-version = "0.1.15"
+version = "0.2.2"
description = "An extremely fast Python linter and code formatter, written in Rust."
optional = false
python-versions = ">=3.7"
groups = ["dev"]
files = [
- {file = "ruff-0.1.15-py3-none-macosx_10_12_x86_64.macosx_11_0_arm64.macosx_10_12_universal2.whl", hash = "sha256:5fe8d54df166ecc24106db7dd6a68d44852d14eb0729ea4672bb4d96c320b7df"},
- {file = "ruff-0.1.15-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:6f0bfbb53c4b4de117ac4d6ddfd33aa5fc31beeaa21d23c45c6dd249faf9126f"},
- {file = "ruff-0.1.15-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:e0d432aec35bfc0d800d4f70eba26e23a352386be3a6cf157083d18f6f5881c8"},
- {file = "ruff-0.1.15-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:9405fa9ac0e97f35aaddf185a1be194a589424b8713e3b97b762336ec79ff807"},
- {file = "ruff-0.1.15-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:c66ec24fe36841636e814b8f90f572a8c0cb0e54d8b5c2d0e300d28a0d7bffec"},
- {file = "ruff-0.1.15-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:6f8ad828f01e8dd32cc58bc28375150171d198491fc901f6f98d2a39ba8e3ff5"},
- {file = "ruff-0.1.15-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:86811954eec63e9ea162af0ffa9f8d09088bab51b7438e8b6488b9401863c25e"},
- {file = "ruff-0.1.15-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:fd4025ac5e87d9b80e1f300207eb2fd099ff8200fa2320d7dc066a3f4622dc6b"},
- {file = "ruff-0.1.15-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:b17b93c02cdb6aeb696effecea1095ac93f3884a49a554a9afa76bb125c114c1"},
- {file = "ruff-0.1.15-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:ddb87643be40f034e97e97f5bc2ef7ce39de20e34608f3f829db727a93fb82c5"},
- {file = "ruff-0.1.15-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:abf4822129ed3a5ce54383d5f0e964e7fef74a41e48eb1dfad404151efc130a2"},
- {file = "ruff-0.1.15-py3-none-musllinux_1_2_i686.whl", hash = "sha256:6c629cf64bacfd136c07c78ac10a54578ec9d1bd2a9d395efbee0935868bf852"},
- {file = "ruff-0.1.15-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:1bab866aafb53da39c2cadfb8e1c4550ac5340bb40300083eb8967ba25481447"},
- {file = "ruff-0.1.15-py3-none-win32.whl", hash = "sha256:2417e1cb6e2068389b07e6fa74c306b2810fe3ee3476d5b8a96616633f40d14f"},
- {file = "ruff-0.1.15-py3-none-win_amd64.whl", hash = "sha256:3837ac73d869efc4182d9036b1405ef4c73d9b1f88da2413875e34e0d6919587"},
- {file = "ruff-0.1.15-py3-none-win_arm64.whl", hash = "sha256:9a933dfb1c14ec7a33cceb1e49ec4a16b51ce3c20fd42663198746efc0427360"},
- {file = "ruff-0.1.15.tar.gz", hash = "sha256:f6dfa8c1b21c913c326919056c390966648b680966febcb796cc9d1aaab8564e"},
+ {file = "ruff-0.2.2-py3-none-macosx_10_12_x86_64.macosx_11_0_arm64.macosx_10_12_universal2.whl", hash = "sha256:0a9efb032855ffb3c21f6405751d5e147b0c6b631e3ca3f6b20f917572b97eb6"},
+ {file = "ruff-0.2.2-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:d450b7fbff85913f866a5384d8912710936e2b96da74541c82c1b458472ddb39"},
+ {file = "ruff-0.2.2-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:ecd46e3106850a5c26aee114e562c329f9a1fbe9e4821b008c4404f64ff9ce73"},
+ {file = "ruff-0.2.2-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:5e22676a5b875bd72acd3d11d5fa9075d3a5f53b877fe7b4793e4673499318ba"},
+ {file = "ruff-0.2.2-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:1695700d1e25a99d28f7a1636d85bafcc5030bba9d0578c0781ba1790dbcf51c"},
+ {file = "ruff-0.2.2-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:b0c232af3d0bd8f521806223723456ffebf8e323bd1e4e82b0befb20ba18388e"},
+ {file = "ruff-0.2.2-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:f63d96494eeec2fc70d909393bcd76c69f35334cdbd9e20d089fb3f0640216ca"},
+ {file = "ruff-0.2.2-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:6a61ea0ff048e06de273b2e45bd72629f470f5da8f71daf09fe481278b175001"},
+ {file = "ruff-0.2.2-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:5e1439c8f407e4f356470e54cdecdca1bd5439a0673792dbe34a2b0a551a2fe3"},
+ {file = "ruff-0.2.2-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:940de32dc8853eba0f67f7198b3e79bc6ba95c2edbfdfac2144c8235114d6726"},
+ {file = "ruff-0.2.2-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:0c126da55c38dd917621552ab430213bdb3273bb10ddb67bc4b761989210eb6e"},
+ {file = "ruff-0.2.2-py3-none-musllinux_1_2_i686.whl", hash = "sha256:3b65494f7e4bed2e74110dac1f0d17dc8e1f42faaa784e7c58a98e335ec83d7e"},
+ {file = "ruff-0.2.2-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:1ec49be4fe6ddac0503833f3ed8930528e26d1e60ad35c2446da372d16651ce9"},
+ {file = "ruff-0.2.2-py3-none-win32.whl", hash = "sha256:d920499b576f6c68295bc04e7b17b6544d9d05f196bb3aac4358792ef6f34325"},
+ {file = "ruff-0.2.2-py3-none-win_amd64.whl", hash = "sha256:cc9a91ae137d687f43a44c900e5d95e9617cb37d4c989e462980ba27039d239d"},
+ {file = "ruff-0.2.2-py3-none-win_arm64.whl", hash = "sha256:c9d15fc41e6054bfc7200478720570078f0b41c9ae4f010bcc16bd6f4d1aacdd"},
+ {file = "ruff-0.2.2.tar.gz", hash = "sha256:e62ed7f36b3068a30ba39193a14274cd706bc486fad521276458022f7bccb31d"},
]
[[package]]
@@ -6640,10 +6680,10 @@ files = [
]
[package.dependencies]
-botocore = ">=1.37.4,<2.0a.0"
+botocore = ">=1.37.4,<2.0a0"
[package.extras]
-crt = ["botocore[crt] (>=1.37.4,<2.0a.0)"]
+crt = ["botocore[crt] (>=1.37.4,<2.0a0)"]
[[package]]
name = "scikit-learn"
@@ -6876,9 +6916,9 @@ tornado = ">=6.4.2,<7"
urllib3 = ">=1.26,<3"
[package.extras]
-all = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)", "cohere (>=5.9.4,<6.00)", "dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\"", "google-cloud-aiplatform (>=1.45.0,<2)", "ipykernel (>=6.25.0,<7)", "llama-cpp-python (>=0.2.28,<0.2.86) ; python_version < \"3.13\"", "mistralai (>=0.0.12,<0.1.0)", "mypy (>=1.7.1,<2)", "ollama (>=0.1.7)", "pillow (>=10.2.0,<11.0.0) ; python_version < \"3.13\"", "pinecone[asyncio] (>=7.0.0,<8.0.0)", "psycopg[binary] (>=3.1.0,<4)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "qdrant-client (>=1.11.1,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "sentence-transformers (>=5.0.0) ; python_version < \"3.13\"", "tokenizers (>=0.19) ; python_version < \"3.13\"", "torch (>=2.6.0) ; python_version < \"3.13\"", "torchvision (>=0.17.0) ; python_version < \"3.13\"", "transformers (>=4.36.2) ; python_version < \"3.13\"", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"]
+all = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)", "cohere (>=5.9.4,<6.0)", "dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\"", "google-cloud-aiplatform (>=1.45.0,<2)", "ipykernel (>=6.25.0,<7)", "llama-cpp-python (>=0.2.28,<0.2.86) ; python_version < \"3.13\"", "mistralai (>=0.0.12,<0.1.0)", "mypy (>=1.7.1,<2)", "ollama (>=0.1.7)", "pillow (>=10.2.0,<11.0.0) ; python_version < \"3.13\"", "pinecone[asyncio] (>=7.0.0,<8.0.0)", "psycopg[binary] (>=3.1.0,<4)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "qdrant-client (>=1.11.1,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "sentence-transformers (>=5.0.0) ; python_version < \"3.13\"", "tokenizers (>=0.19) ; python_version < \"3.13\"", "torch (>=2.6.0) ; python_version < \"3.13\"", "torchvision (>=0.17.0) ; python_version < \"3.13\"", "transformers (>=4.36.2) ; python_version < \"3.13\"", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"]
bedrock = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)"]
-cohere = ["cohere (>=5.9.4,<6.00)"]
+cohere = ["cohere (>=5.9.4,<6.0)"]
dev = ["dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "ipykernel (>=6.25.0,<7)", "mypy (>=1.7.1,<2)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"]
docs = ["pydoc-markdown (>=4.8.2) ; python_version < \"3.12\""]
fastembed = ["fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\""]
@@ -7722,6 +7762,7 @@ files = [
{file = "tomlkit-0.13.3-py3-none-any.whl", hash = "sha256:c89c649d79ee40629a9fda55f8ace8c6a1b42deb912b2a8fd8d942ddadb606b0"},
{file = "tomlkit-0.13.3.tar.gz", hash = "sha256:430cf247ee57df2b94ee3fbe588e71d362a941ebb545dec29b53961d61add2a1"},
]
+markers = {main = "extra == \"extra-proxy\""}
[[package]]
name = "tornado"
@@ -8490,4 +8531,4 @@ utils = ["numpydoc"]
[metadata]
lock-version = "2.1"
python-versions = ">=3.9,<4.0"
-content-hash = "95fd27dc139d0e52e70093220c50582f16c78e5977ec77f4297f50a30df964c6"
+content-hash = "e5447e14dd37e324ac07a8fc6286d27e9a0d355ed93ebb24fc11e3f5df12fd3e"
diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json
index 0738c6e4e09..fd17b5309e8 100644
--- a/provider_endpoints_support.json
+++ b/provider_endpoints_support.json
@@ -32,6 +32,23 @@
}
},
"providers": {
+ "a2a": {
+ "display_name": "A2A (Agent-to-Agent) (`a2a`)",
+ "url": "https://docs.litellm.ai/docs/providers/a2a",
+ "endpoints": {
+ "chat_completions": true,
+ "messages": false,
+ "responses": false,
+ "embeddings": false,
+ "image_generations": false,
+ "audio_transcriptions": false,
+ "audio_speech": false,
+ "moderations": false,
+ "batches": false,
+ "rerank": false,
+ "a2a": false
+ }
+ },
"abliteration": {
"display_name": "Abliteration (`abliteration`)",
"url": "https://docs.litellm.ai/docs/providers/abliteration",
@@ -2166,7 +2183,8 @@
"batches": false,
"rerank": false,
"a2a": true,
- "interactions": true
+ "interactions": true,
+ "realtime": true
}
},
"xinference": {
diff --git a/pyproject.toml b/pyproject.toml
index 450dadac930..fe76d8e15df 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm"
-version = "1.81.6"
+version = "1.81.8"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI"]
license = "MIT"
@@ -61,9 +61,9 @@ boto3 = { version = "1.40.76", optional = true }
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
mcp = {version = ">=1.25.0,<2.0.0", optional = true, python = ">=3.10"}
a2a-sdk = {version = "^0.3.22", optional = true, python = ">=3.10"}
-litellm-proxy-extras = {version = "0.4.29", optional = true}
+litellm-proxy-extras = {version = "0.4.31", optional = true}
rich = {version = "13.7.1", optional = true}
-litellm-enterprise = {version = "0.1.27", optional = true}
+litellm-enterprise = {version = "0.1.31", optional = true}
diskcache = {version = "^5.6.1", optional = true}
polars = {version = "^1.31.0", optional = true, python = ">=3.10"}
semantic-router = {version = ">=0.1.12", optional = true, python = ">=3.9,<3.14"}
@@ -139,6 +139,7 @@ litellm = 'litellm:run_server'
litellm-proxy = 'litellm.proxy.client.cli:cli'
[tool.poetry.group.dev.dependencies]
+diff-cover = "^9.0"
flake8 = "^6.1.0"
black = "^23.12.0"
mypy = "^1.0"
@@ -149,7 +150,7 @@ pytest-retry = "^1.6.3"
requests-mock = "^1.12.1"
responses = "^0.25.7"
respx = "^0.22.0"
-ruff = "^0.1.0"
+ruff = "^0.2.1"
types-requests = "*"
types-setuptools = "*"
types-redis = "*"
@@ -174,7 +175,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
-version = "1.81.6"
+version = "1.81.8"
version_files = [
"pyproject.toml:^version"
]
diff --git a/requirements.txt b/requirements.txt
index 0b7cf4992e8..8b69d4ac85c 100644
--- a/requirements.txt
+++ b/requirements.txt
@@ -50,7 +50,7 @@ sentry_sdk==2.21.0 # for sentry error handling
detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests
cryptography==44.0.1
tzdata==2025.1 # IANA time zone database
-litellm-proxy-extras==0.4.29 # for proxy extras - e.g. prisma migrations
+litellm-proxy-extras==0.4.31 # for proxy extras - e.g. prisma migrations
llm-sandbox==0.3.31 # for skill execution in sandbox
### LITELLM PACKAGE DEPENDENCIES
python-dotenv==1.0.1 # for env
@@ -73,4 +73,4 @@ pypdf>=6.6.2 # for PDF text extraction in RAG ingestion
########################
# LITELLM ENTERPRISE DEPENDENCIES
########################
-litellm-enterprise==0.1.28
+litellm-enterprise==0.1.31
diff --git a/schema.prisma b/schema.prisma
index b118400b620..240e0dfea48 100644
--- a/schema.prisma
+++ b/schema.prisma
@@ -113,6 +113,7 @@ model LiteLLM_TeamTable {
members_with_roles Json @default("{}")
metadata Json @default("{}")
max_budget Float?
+ soft_budget Float?
spend Float @default(0.0)
models String[]
max_parallel_requests Int?
@@ -129,6 +130,7 @@ model LiteLLM_TeamTable {
team_member_permissions String[] @default([])
policies String[] @default([])
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
+ allow_team_guardrail_config Boolean @default(false) // if true, team admin can configure guardrails for this team
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id])
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
@@ -160,6 +162,7 @@ model LiteLLM_DeletedTeamTable {
team_member_permissions String[] @default([])
policies String[] @default([])
model_id Int? // id for LiteLLM_ModelTable -> stores team-level model aliases
+ allow_team_guardrail_config Boolean @default(false)
// Original timestamps from team creation/updates
created_at DateTime? @map("created_at")
diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py
index d5640f4256c..71e7798b09e 100644
--- a/tests/code_coverage_tests/recursive_detector.py
+++ b/tests/code_coverage_tests/recursive_detector.py
@@ -40,6 +40,8 @@ IGNORE_FUNCTIONS = [
"filter_exceptions_from_params", # max depth set (default 20) to prevent infinite recursion.
"__getattr__", # lazy loading pattern in litellm/__init__.py with proper caching to prevent infinite recursion.
"_validate_inheritance_chain", # max depth set (default 100) to prevent infinite recursion in policy inheritance validation.
+ "_basic_json_schema_validate", # max depth set.
+ "extract_text_from_a2a_message", # max depth set (default 10) to prevent infinite recursion in A2A message parsing.
]
diff --git a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py
index 0a57d046c72..c39454728a8 100644
--- a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py
+++ b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py
@@ -1,4 +1,3 @@
-import io
import os
import sys
@@ -10,13 +9,10 @@ from datetime import datetime, timedelta, timezone
from unittest.mock import MagicMock, call, patch
import pytest
-from prometheus_client import REGISTRY, CollectorRegistry
+from prometheus_client import REGISTRY
import litellm
-from litellm import completion
from litellm._logging import verbose_logger
-from litellm._uuid import uuid
-from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.types.utils import (
StandardLoggingHiddenParams,
StandardLoggingMetadata,
@@ -37,7 +33,6 @@ from litellm.proxy._types import UserAPIKeyAuth
verbose_logger.setLevel(logging.DEBUG)
litellm.set_verbose = True
-import time
@pytest.fixture
@@ -293,7 +288,6 @@ async def test_increment_remaining_budget_metrics(prometheus_logger):
) as mock_get_team, patch(
"litellm.proxy.auth.auth_checks.get_key_object"
) as mock_get_key:
-
mock_get_team.return_value = MagicMock(budget_reset_at=future_reset_time_team)
mock_get_key.return_value = MagicMock(budget_reset_at=future_reset_time_key)
@@ -648,25 +642,16 @@ async def test_async_log_failure_event(prometheus_logger):
)
# litellm_llm_api_failed_requests_metric incremented
- """
- Expected metrics
- end_user_id,
- user_api_key,
- user_api_key_alias,
- model,
- user_api_team,
- user_api_team_alias,
- user_id,
- """
+ # Labels: end_user, api_key_hash, api_key_alias, model, team, team_alias, user, model_id
prometheus_logger.litellm_llm_api_failed_requests_metric.labels.assert_called_once_with(
- None,
+ None, # end_user_id
"test_hash",
"test_alias",
"gpt-3.5-turbo",
"test_team",
"test_team_alias",
"test_user",
- "model-123",
+ "model-123", # model_id from standard_logging_payload
)
prometheus_logger.litellm_llm_api_failed_requests_metric.labels().inc.assert_called_once()
@@ -678,38 +663,54 @@ async def test_async_log_failure_event(prometheus_logger):
api_provider="openai",
)
- # deployment failure responses incremented
- prometheus_logger.litellm_deployment_failure_responses.labels.assert_called_once_with(
- litellm_model_name="gpt-3.5-turbo",
- model_id="model-123",
- api_base="https://api.openai.com",
- api_provider="openai",
- exception_status="None",
- exception_class="Exception",
- requested_model="openai-gpt", # passed in standard logging payload
- hashed_api_key="test_hash",
- api_key_alias="test_alias",
- team="test_team",
- team_alias="test_team_alias",
- client_ip="127.0.0.1", # from standard logging payload
- user_agent=None,
+ # deployment failure responses incremented - verify key labels are populated
+ prometheus_logger.litellm_deployment_failure_responses.labels.assert_called_once()
+ actual_failure_labels = (
+ prometheus_logger.litellm_deployment_failure_responses.labels.call_args.kwargs
)
+ expected_failure_labels = {
+ "litellm_model_name": "gpt-3.5-turbo",
+ "model_id": "model-123",
+ "api_base": "https://api.openai.com",
+ "api_provider": "openai",
+ "exception_class": "Exception",
+ "requested_model": "openai-gpt",
+ "hashed_api_key": "test_hash",
+ "api_key_alias": "test_alias",
+ "team": "test_team",
+ "team_alias": "test_team_alias",
+ }
+ for key, expected_val in expected_failure_labels.items():
+ assert key in actual_failure_labels, f"Missing label {key}"
+ assert (
+ actual_failure_labels[key] == expected_val
+ ), f"Label {key}: expected {expected_val!r}, got {actual_failure_labels[key]!r}"
+ assert actual_failure_labels.get("exception_status") in ("None", None)
+ assert actual_failure_labels.get("client_ip") == "127.0.0.1"
prometheus_logger.litellm_deployment_failure_responses.labels().inc.assert_called_once()
- # deployment total requests incremented
- prometheus_logger.litellm_deployment_total_requests.labels.assert_called_once_with(
- litellm_model_name="gpt-3.5-turbo",
- model_id="model-123",
- api_base="https://api.openai.com",
- api_provider="openai",
- requested_model="openai-gpt", # passed in standard logging payload
- hashed_api_key="test_hash",
- api_key_alias="test_alias",
- team="test_team",
- team_alias="test_team_alias",
- client_ip="127.0.0.1", # from standard logging payload
- user_agent=None,
+ # deployment total requests incremented - verify key labels are populated
+ prometheus_logger.litellm_deployment_total_requests.labels.assert_called_once()
+ actual_total_labels = (
+ prometheus_logger.litellm_deployment_total_requests.labels.call_args.kwargs
)
+ expected_total_labels = {
+ "litellm_model_name": "gpt-3.5-turbo",
+ "model_id": "model-123",
+ "api_base": "https://api.openai.com",
+ "api_provider": "openai",
+ "requested_model": "openai-gpt",
+ "hashed_api_key": "test_hash",
+ "api_key_alias": "test_alias",
+ "team": "test_team",
+ "team_alias": "test_team_alias",
+ }
+ for key, expected_val in expected_total_labels.items():
+ assert key in actual_total_labels, f"Missing label {key}"
+ assert (
+ actual_total_labels[key] == expected_val
+ ), f"Label {key}: expected {expected_val!r}, got {actual_total_labels[key]!r}"
+ assert actual_total_labels.get("client_ip") == "127.0.0.1"
prometheus_logger.litellm_deployment_total_requests.labels().inc.assert_called_once()
@@ -1095,7 +1096,7 @@ def test_increment_deployment_cooled_down(prometheus_logger):
import inspect
method_sig = inspect.signature(prometheus_logger.increment_deployment_cooled_down)
- expected_label_count = len([p for p in method_sig.parameters.keys() if p != 'self'])
+ expected_label_count = len([p for p in method_sig.parameters.keys() if p != "self"])
mock_chain = MagicMock()
@@ -1103,11 +1104,15 @@ def test_increment_deployment_cooled_down(prometheus_logger):
"""Validate label count matches metric definition"""
total = len(label_values) + len(label_kwargs)
if total != expected_label_count:
- raise ValueError(f"Incorrect label count: expected {expected_label_count}, got {total}")
+ raise ValueError(
+ f"Incorrect label count: expected {expected_label_count}, got {total}"
+ )
return mock_chain
prometheus_logger.litellm_deployment_cooled_down = MagicMock()
- prometheus_logger.litellm_deployment_cooled_down.labels = MagicMock(side_effect=validating_labels)
+ prometheus_logger.litellm_deployment_cooled_down.labels = MagicMock(
+ side_effect=validating_labels
+ )
prometheus_logger.increment_deployment_cooled_down(
litellm_model_name="gpt-3.5-turbo",
@@ -1179,8 +1184,12 @@ def test_get_custom_labels_from_top_level_metadata(monkeypatch):
metadata = {
"requester_ip_address": "10.48.203.20", # Top-level field
"user_api_key_alias": "TestAlias", # Top-level field
- "requester_metadata": {"nested_field": "nested_value"}, # Nested dict (excluded)
- "user_api_key_auth_metadata": {"another_nested": "value"}, # Nested dict (excluded)
+ "requester_metadata": {
+ "nested_field": "nested_value"
+ }, # Nested dict (excluded)
+ "user_api_key_auth_metadata": {
+ "another_nested": "value"
+ }, # Nested dict (excluded)
}
result = get_custom_labels_from_metadata(metadata)
assert result == {
@@ -1217,7 +1226,9 @@ def test_get_custom_labels_from_top_level_and_nested_metadata(monkeypatch):
}
-async def test_async_log_success_event_with_top_level_metadata(prometheus_logger, monkeypatch):
+async def test_async_log_success_event_with_top_level_metadata(
+ prometheus_logger, monkeypatch
+):
"""
Test that async_log_success_event correctly extracts custom labels from top-level metadata
fields like requester_ip_address, not just from nested dictionaries.
@@ -1231,7 +1242,9 @@ async def test_async_log_success_event_with_top_level_metadata(prometheus_logger
standard_logging_object = create_standard_logging_payload()
standard_logging_object["metadata"]["requester_ip_address"] = "10.48.203.20"
standard_logging_object["metadata"]["requester_metadata"] = {} # Empty nested dict
- standard_logging_object["metadata"]["user_api_key_auth_metadata"] = {} # Empty nested dict
+ standard_logging_object["metadata"][
+ "user_api_key_auth_metadata"
+ ] = {} # Empty nested dict
kwargs = {
"model": "gpt-3.5-turbo",
@@ -1273,7 +1286,9 @@ async def test_async_log_success_event_with_top_level_metadata(prometheus_logger
prometheus_logger.litellm_remaining_user_budget_metric = create_mock_metric()
prometheus_logger.litellm_user_max_budget_metric = create_mock_metric()
prometheus_logger.litellm_user_budget_remaining_hours_metric = create_mock_metric()
- prometheus_logger.litellm_remaining_api_key_requests_for_model = create_mock_metric()
+ prometheus_logger.litellm_remaining_api_key_requests_for_model = (
+ create_mock_metric()
+ )
prometheus_logger.litellm_remaining_api_key_tokens_for_model = create_mock_metric()
prometheus_logger.litellm_llm_api_time_to_first_token_metric = create_mock_metric()
prometheus_logger.litellm_llm_api_latency_metric = create_mock_metric()
@@ -1302,7 +1317,7 @@ async def test_async_log_success_event_with_top_level_metadata(prometheus_logger
# This confirms that the custom label extraction logic ran without errors
assert prometheus_logger.litellm_requests_metric.labels.called
assert prometheus_logger.litellm_spend_metric.labels.called
-
+
# Verify that the labels() method was called with some arguments (either positional or keyword)
# This ensures the custom label extraction happened and didn't cause a "Incorrect label names" error
call_args = prometheus_logger.litellm_requests_metric.labels.call_args
@@ -1494,7 +1509,6 @@ async def test_initialize_remaining_budget_metrics(prometheus_logger):
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams"
) as mock_get_teams:
-
# Create mock team data with proper datetime objects for budget_reset_at
future_reset = datetime.now() + timedelta(hours=24) # Reset 24 hours from now
mock_teams = [
@@ -1592,21 +1606,22 @@ async def test_initialize_remaining_budget_metrics_exception_handling(
) as mock_get_teams, patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper"
) as mock_list_keys:
-
# Make get_paginated_teams raise an exception
mock_get_teams.side_effect = Exception("Database error")
mock_list_keys.side_effect = Exception("Key listing error")
-
+
# Mock prisma_client structure to raise an exception for user budget metrics
# The code accesses prisma_client.db.litellm_usertable.find_many and count
mock_usertable = MagicMock()
- mock_usertable.find_many = MagicMock(side_effect=Exception("User database error"))
+ mock_usertable.find_many = MagicMock(
+ side_effect=Exception("User database error")
+ )
mock_usertable.count = MagicMock(side_effect=Exception("User count error"))
-
+
# Mock litellm_teamtable to raise an exception for team count metrics
mock_teamtable = MagicMock()
mock_teamtable.count = MagicMock(side_effect=Exception("Team count error"))
-
+
mock_db = MagicMock()
mock_db.litellm_usertable = mock_usertable
mock_db.litellm_teamtable = mock_teamtable
@@ -1661,7 +1676,6 @@ async def test_initialize_api_key_budget_metrics(prometheus_logger):
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper"
) as mock_list_keys:
-
# Create mock key data with proper datetime objects for budget_reset_at
future_reset = datetime.now() + timedelta(hours=24) # Reset 24 hours from now
key1 = UserAPIKeyAuth(
@@ -1916,7 +1930,6 @@ def test_prometheus_label_factory_with_custom_tags(monkeypatch):
Test that prometheus_label_factory correctly handles custom tags
"""
from litellm.integrations.prometheus import (
- get_custom_labels_from_tags,
prometheus_label_factory,
)
from litellm.types.integrations.prometheus import UserAPIKeyLabelValues
@@ -1954,7 +1967,6 @@ def test_prometheus_label_factory_with_no_custom_tags(monkeypatch):
Test that prometheus_label_factory works when no custom tags are configured
"""
from litellm.integrations.prometheus import (
- get_custom_labels_from_tags,
prometheus_label_factory,
)
from litellm.types.integrations.prometheus import UserAPIKeyLabelValues
@@ -2179,9 +2191,7 @@ async def test_prometheus_token_metrics_with_prometheus_config():
All three metrics should be properly incremented when making a successful completion request.
"""
- from prometheus_client import CollectorRegistry, Counter
- import litellm
from litellm.types.integrations.prometheus import PrometheusMetricsConfig
# Clear registry before test
diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py
index 3fd19cfa18f..946c5ad1729 100644
--- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py
+++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py
@@ -587,6 +587,499 @@ def test_update_responses_input_with_multiple_file_ids():
assert updated_input[0]["content"][1]["text"] == "Compare these files"
+def test_update_responses_input_with_model_file_id_mapping():
+ """
+ Test that update_responses_input_with_model_file_ids correctly uses
+ model_file_id_mapping to map managed file IDs to provider-specific file IDs.
+ """
+ from litellm.litellm_core_utils.prompt_templates.common_utils import (
+ update_responses_input_with_model_file_ids,
+ )
+
+ # Managed file ID (unified)
+ managed_file_id = "litellm_proxy_file_123"
+
+ # Model file ID mapping
+ model_file_id_mapping = {
+ managed_file_id: {
+ "model_id_1": "openai_file_abc",
+ "model_id_2": "azure_file_xyz",
+ }
+ }
+
+ input_data = [
+ {
+ "role": "user",
+ "content": [
+ {
+ "type": "input_file",
+ "file_id": managed_file_id,
+ },
+ {
+ "type": "input_text",
+ "text": "Analyze this file",
+ },
+ ],
+ }
+ ]
+
+ # Update input with model_id_1 mapping
+ updated_input = update_responses_input_with_model_file_ids(
+ input=input_data,
+ model_id="model_id_1",
+ model_file_id_mapping=model_file_id_mapping,
+ )
+
+ # Verify the file_id was mapped to the correct provider-specific file ID
+ assert updated_input[0]["content"][0]["file_id"] == "openai_file_abc"
+
+ # Test with different model_id
+ updated_input_2 = update_responses_input_with_model_file_ids(
+ input=input_data,
+ model_id="model_id_2",
+ model_file_id_mapping=model_file_id_mapping,
+ )
+
+ assert updated_input_2[0]["content"][0]["file_id"] == "azure_file_xyz"
+
+
+def test_update_responses_tools_with_model_file_id_mapping():
+ """
+ Test that update_responses_tools_with_model_file_ids correctly maps
+ file IDs in code_interpreter tools with container.file_ids.
+
+ This is a regression test for the issue where managed file IDs in
+ tools.container.file_ids were not being replaced with provider-specific
+ file IDs, causing "string too long" errors from OpenAI.
+ """
+ from litellm.litellm_core_utils.prompt_templates.common_utils import (
+ update_responses_tools_with_model_file_ids,
+ )
+
+ # Managed file IDs
+ managed_file_id_1 = "litellm_proxy_file_123"
+ managed_file_id_2 = "litellm_proxy_file_456"
+
+ # Model file ID mapping
+ model_file_id_mapping = {
+ managed_file_id_1: {
+ "model_id_1": "openai_file_abc",
+ },
+ managed_file_id_2: {
+ "model_id_1": "openai_file_def",
+ },
+ }
+
+ tools = [
+ {
+ "type": "code_interpreter",
+ "container": {
+ "type": "auto",
+ "file_ids": [managed_file_id_1, managed_file_id_2],
+ },
+ }
+ ]
+
+ # Update tools with model mapping
+ updated_tools = update_responses_tools_with_model_file_ids(
+ tools=tools,
+ model_id="model_id_1",
+ model_file_id_mapping=model_file_id_mapping,
+ )
+
+ # Verify the file IDs were mapped to provider-specific file IDs
+ assert updated_tools[0]["type"] == "code_interpreter"
+ assert updated_tools[0]["container"]["file_ids"] == ["openai_file_abc", "openai_file_def"]
+
+
+def test_update_responses_tools_without_mapping():
+ """
+ Test that update_responses_tools_with_model_file_ids keeps file IDs
+ unchanged when no mapping is provided.
+ """
+ from litellm.litellm_core_utils.prompt_templates.common_utils import (
+ update_responses_tools_with_model_file_ids,
+ )
+
+ regular_file_id = "file-abc123"
+
+ tools = [
+ {
+ "type": "code_interpreter",
+ "container": {
+ "type": "auto",
+ "file_ids": [regular_file_id],
+ },
+ }
+ ]
+
+ # Update tools without mapping
+ updated_tools = update_responses_tools_with_model_file_ids(
+ tools=tools,
+ model_id=None,
+ model_file_id_mapping=None,
+ )
+
+ # Verify the file ID was kept unchanged
+ assert updated_tools[0]["container"]["file_ids"] == [regular_file_id]
+
+
+def test_update_responses_tools_with_mixed_file_ids():
+ """
+ Test that update_responses_tools_with_model_file_ids correctly handles
+ a mix of managed and regular file IDs.
+ """
+ from litellm.litellm_core_utils.prompt_templates.common_utils import (
+ update_responses_tools_with_model_file_ids,
+ )
+
+ managed_file_id = "litellm_proxy_file_123"
+ regular_file_id = "file-abc123"
+
+ model_file_id_mapping = {
+ managed_file_id: {
+ "model_id_1": "openai_file_abc",
+ },
+ }
+
+ tools = [
+ {
+ "type": "code_interpreter",
+ "container": {
+ "type": "auto",
+ "file_ids": [managed_file_id, regular_file_id],
+ },
+ }
+ ]
+
+ # Update tools
+ updated_tools = update_responses_tools_with_model_file_ids(
+ tools=tools,
+ model_id="model_id_1",
+ model_file_id_mapping=model_file_id_mapping,
+ )
+
+ # Verify managed file ID was mapped and regular file ID was kept
+ assert updated_tools[0]["container"]["file_ids"] == ["openai_file_abc", regular_file_id]
+
+
+def test_get_file_ids_from_responses_tools():
+ """
+ Test that get_file_ids_from_responses_tools correctly extracts
+ file IDs from the tools parameter.
+ """
+ proxy_managed_files = _PROXY_LiteLLMManagedFiles(
+ DualCache(), prisma_client=MagicMock()
+ )
+
+ tools = [
+ {
+ "type": "code_interpreter",
+ "container": {
+ "type": "auto",
+ "file_ids": ["file-123", "file-456"],
+ },
+ }
+ ]
+
+ file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools)
+
+ assert file_ids == ["file-123", "file-456"]
+
+
+def test_get_file_ids_from_responses_tools_multiple_tools():
+ """
+ Test that get_file_ids_from_responses_tools handles multiple tools.
+ """
+ proxy_managed_files = _PROXY_LiteLLMManagedFiles(
+ DualCache(), prisma_client=MagicMock()
+ )
+
+ tools = [
+ {
+ "type": "code_interpreter",
+ "container": {
+ "type": "auto",
+ "file_ids": ["file-123"],
+ },
+ },
+ {
+ "type": "file_search",
+ },
+ {
+ "type": "code_interpreter",
+ "container": {
+ "type": "auto",
+ "file_ids": ["file-456", "file-789"],
+ },
+ },
+ ]
+
+ file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools)
+
+ # Should extract file IDs only from code_interpreter tools
+ assert file_ids == ["file-123", "file-456", "file-789"]
+
+
+def test_get_file_ids_from_responses_tools_empty():
+ """
+ Test that get_file_ids_from_responses_tools handles empty or None tools.
+ """
+ proxy_managed_files = _PROXY_LiteLLMManagedFiles(
+ DualCache(), prisma_client=MagicMock()
+ )
+
+ # Test with None
+ file_ids = proxy_managed_files.get_file_ids_from_responses_tools(None)
+ assert file_ids == []
+
+ # Test with empty list
+ file_ids = proxy_managed_files.get_file_ids_from_responses_tools([])
+ assert file_ids == []
+
+ # Test with tools without file_ids
+ tools = [{"type": "file_search"}]
+ file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools)
+ assert file_ids == []
+
+
+@pytest.mark.asyncio
+async def test_check_file_ids_access_with_unified_file_ids():
+ """
+ Test that check_file_ids_access validates user access to managed file IDs.
+ """
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ # Create a unified file ID
+ unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My"
+ regular_file_id = "file-abc123"
+
+ # Mock the access check to return True
+ prisma_client = AsyncMock()
+ internal_usage_cache = MagicMock()
+
+ proxy_managed_files = _PROXY_LiteLLMManagedFiles(
+ internal_usage_cache=internal_usage_cache,
+ prisma_client=prisma_client,
+ )
+
+ # Mock can_user_call_unified_file_id to return True
+ proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True)
+
+ user_api_key_dict = UserAPIKeyAuth(
+ user_id="test_user_123",
+ parent_otel_span=MagicMock(),
+ )
+
+ # Should not raise an exception for accessible files
+ await proxy_managed_files.check_file_ids_access(
+ [unified_file_id, regular_file_id],
+ user_api_key_dict,
+ )
+
+ # Verify can_user_call_unified_file_id was called for the unified file ID
+ proxy_managed_files.can_user_call_unified_file_id.assert_called_once_with(
+ unified_file_id, user_api_key_dict
+ )
+
+
+@pytest.mark.asyncio
+async def test_check_file_ids_access_denied():
+ """
+ Test that check_file_ids_access raises HTTPException when user doesn't have access.
+ """
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My"
+
+ prisma_client = AsyncMock()
+ internal_usage_cache = MagicMock()
+
+ proxy_managed_files = _PROXY_LiteLLMManagedFiles(
+ internal_usage_cache=internal_usage_cache,
+ prisma_client=prisma_client,
+ )
+
+ # Mock can_user_call_unified_file_id to return False (access denied)
+ proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=False)
+
+ user_api_key_dict = UserAPIKeyAuth(
+ user_id="test_user_123",
+ parent_otel_span=MagicMock(),
+ )
+
+ # Should raise HTTPException with 403 status code
+ with pytest.raises(HTTPException) as exc_info:
+ await proxy_managed_files.check_file_ids_access(
+ [unified_file_id],
+ user_api_key_dict,
+ )
+
+ assert exc_info.value.status_code == 403
+ assert "does not have access to the file" in exc_info.value.detail
+
+
+@pytest.mark.asyncio
+async def test_check_file_ids_access_with_regular_files_only():
+ """
+ Test that check_file_ids_access doesn't check access for regular (non-unified) file IDs.
+ """
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ regular_file_id_1 = "file-abc123"
+ regular_file_id_2 = "file-xyz789"
+
+ prisma_client = AsyncMock()
+ internal_usage_cache = MagicMock()
+
+ proxy_managed_files = _PROXY_LiteLLMManagedFiles(
+ internal_usage_cache=internal_usage_cache,
+ prisma_client=prisma_client,
+ )
+
+ # Mock can_user_call_unified_file_id (should not be called for regular files)
+ proxy_managed_files.can_user_call_unified_file_id = AsyncMock()
+
+ user_api_key_dict = UserAPIKeyAuth(
+ user_id="test_user_123",
+ parent_otel_span=MagicMock(),
+ )
+
+ # Should not raise exception and should not call can_user_call_unified_file_id
+ await proxy_managed_files.check_file_ids_access(
+ [regular_file_id_1, regular_file_id_2],
+ user_api_key_dict,
+ )
+
+ # Verify can_user_call_unified_file_id was NOT called
+ proxy_managed_files.can_user_call_unified_file_id.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_completion_with_file_access_check():
+ """
+ Test that completion call type checks file access before processing.
+ """
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My"
+
+ prisma_client = AsyncMock()
+ prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
+
+ internal_usage_cache = MagicMock()
+ internal_usage_cache.async_get_cache = AsyncMock(return_value=None)
+
+ proxy_managed_files = _PROXY_LiteLLMManagedFiles(
+ internal_usage_cache=internal_usage_cache,
+ prisma_client=prisma_client,
+ )
+
+ # Mock the get_model_file_id_mapping to return empty dict
+ proxy_managed_files.get_model_file_id_mapping = AsyncMock(return_value={})
+
+ # Mock access check to allow access
+ proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True)
+
+ user_api_key_dict = UserAPIKeyAuth(
+ user_id="test_user_123",
+ parent_otel_span=MagicMock(),
+ )
+
+ data = {
+ "messages": [
+ {
+ "role": "user",
+ "content": [
+ {"type": "text", "text": "What's in this file?"},
+ {
+ "type": "file",
+ "file": {"file_id": unified_file_id},
+ },
+ ],
+ }
+ ],
+ "model": "gpt-4",
+ }
+
+ # Should not raise exception
+ result = await proxy_managed_files.async_pre_call_hook(
+ user_api_key_dict=user_api_key_dict,
+ cache=DualCache(),
+ data=data,
+ call_type="acompletion",
+ )
+
+ # Verify access check was called
+ proxy_managed_files.can_user_call_unified_file_id.assert_called_once()
+
+
+@pytest.mark.asyncio
+async def test_responses_with_file_access_check():
+ """
+ Test that responses API checks file access for files in both input and tools.
+ """
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ unified_file_id_1 = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My"
+ unified_file_id_2 = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsNzc3Nzc3Nzc7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1YWVo7bGxtX291dHB1dF9maWxlX21vZGVsX2lkLG1vZGVsXzEyMw"
+
+ prisma_client = AsyncMock()
+ prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
+
+ internal_usage_cache = MagicMock()
+ internal_usage_cache.async_get_cache = AsyncMock(return_value=None)
+
+ proxy_managed_files = _PROXY_LiteLLMManagedFiles(
+ internal_usage_cache=internal_usage_cache,
+ prisma_client=prisma_client,
+ )
+
+ # Mock the get_model_file_id_mapping to return empty dict
+ proxy_managed_files.get_model_file_id_mapping = AsyncMock(return_value={})
+
+ # Mock access check to allow access
+ proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True)
+
+ user_api_key_dict = UserAPIKeyAuth(
+ user_id="test_user_123",
+ parent_otel_span=MagicMock(),
+ )
+
+ data = {
+ "input": [
+ {
+ "role": "user",
+ "content": [
+ {"type": "input_text", "text": "Analyze this"},
+ {"type": "input_file", "file_id": unified_file_id_1},
+ ],
+ }
+ ],
+ "tools": [
+ {
+ "type": "code_interpreter",
+ "container": {
+ "type": "auto",
+ "file_ids": [unified_file_id_2],
+ },
+ }
+ ],
+ "model": "gpt-4",
+ }
+
+ # Should not raise exception
+ result = await proxy_managed_files.async_pre_call_hook(
+ user_api_key_dict=user_api_key_dict,
+ cache=DualCache(),
+ data=data,
+ call_type="aresponses",
+ )
+
+ # Verify access check was called for both file IDs
+ assert proxy_managed_files.can_user_call_unified_file_id.call_count == 2
+
+
@pytest.mark.asyncio
async def test_store_unified_file_id_with_none_file_object():
"""
diff --git a/tests/litellm/test_proxy_auth.py b/tests/litellm/test_proxy_auth.py
new file mode 100644
index 00000000000..1d73e143e10
--- /dev/null
+++ b/tests/litellm/test_proxy_auth.py
@@ -0,0 +1,204 @@
+"""
+Unit tests for litellm.proxy_auth module.
+
+Tests the OAuth2/JWT token management for LiteLLM Proxy authentication.
+"""
+
+import time
+from unittest.mock import Mock, patch
+
+import pytest
+
+from litellm.proxy_auth import (
+ AccessToken,
+ AzureADCredential,
+ GenericOAuth2Credential,
+ ProxyAuthHandler,
+)
+
+
+class TestAccessToken:
+ """Tests for AccessToken dataclass."""
+
+ def test_access_token_creation(self):
+ """Test AccessToken can be created with required fields."""
+ token = AccessToken(token="test-token", expires_on=1234567890)
+ assert token.token == "test-token"
+ assert token.expires_on == 1234567890
+
+ def test_access_token_equality(self):
+ """Test AccessToken equality comparison."""
+ token1 = AccessToken(token="test", expires_on=123)
+ token2 = AccessToken(token="test", expires_on=123)
+ assert token1 == token2
+
+
+class MockCredential:
+ """Mock credential for testing."""
+
+ def __init__(self, expires_in_seconds: int = 3600):
+ self.call_count = 0
+ self.expires_in = expires_in_seconds
+
+ def get_token(self, scope: str) -> AccessToken:
+ self.call_count += 1
+ return AccessToken(
+ token=f"mock-token-{self.call_count}",
+ expires_on=int(time.time()) + self.expires_in,
+ )
+
+
+class TestProxyAuthHandler:
+ """Tests for ProxyAuthHandler."""
+
+ def test_get_auth_headers_returns_bearer_token(self):
+ """Test that get_auth_headers returns correct Authorization header."""
+ cred = MockCredential()
+ handler = ProxyAuthHandler(credential=cred, scope="test-scope")
+
+ headers = handler.get_auth_headers()
+
+ assert "Authorization" in headers
+ assert headers["Authorization"].startswith("Bearer ")
+ assert "mock-token-1" in headers["Authorization"]
+
+ def test_token_caching(self):
+ """Test that tokens are cached and not re-requested."""
+ cred = MockCredential(expires_in_seconds=3600) # Long expiry
+ handler = ProxyAuthHandler(credential=cred, scope="test-scope")
+
+ # Multiple calls should only request token once
+ handler.get_auth_headers()
+ handler.get_auth_headers()
+ handler.get_auth_headers()
+
+ assert cred.call_count == 1
+
+ def test_token_refresh_when_about_to_expire(self):
+ """Test that tokens are refreshed when about to expire (within 60s buffer)."""
+ cred = MockCredential(expires_in_seconds=30) # Expires in 30s (< 60s buffer)
+ handler = ProxyAuthHandler(credential=cred, scope="test-scope")
+
+ # First call gets token
+ handler.get_auth_headers()
+ # Second call should refresh because token expires within 60s buffer
+ handler.get_auth_headers()
+
+ assert cred.call_count == 2
+
+ def test_get_token_method(self):
+ """Test the get_token method returns AccessToken."""
+ cred = MockCredential()
+ handler = ProxyAuthHandler(credential=cred, scope="test-scope")
+
+ token = handler.get_token()
+
+ assert isinstance(token, AccessToken)
+ assert token.token == "mock-token-1"
+
+
+class TestAzureADCredential:
+ """Tests for AzureADCredential."""
+
+ def test_lazy_initialization(self):
+ """Test that azure-identity is not imported until get_token is called."""
+ # This should not raise ImportError even if azure-identity is not installed
+ cred = AzureADCredential(credential=None)
+ # _initialized should be False until get_token is called
+ assert cred._initialized is False
+
+ def test_wraps_azure_credential(self):
+ """Test that AzureADCredential wraps an azure-identity credential."""
+ # Mock Azure credential
+ mock_azure_cred = Mock()
+ mock_azure_cred.get_token.return_value = Mock(
+ token="azure-token", expires_on=9999999999
+ )
+
+ cred = AzureADCredential(credential=mock_azure_cred)
+ token = cred.get_token("https://graph.microsoft.com/.default")
+
+ assert token.token == "azure-token"
+ assert token.expires_on == 9999999999
+ mock_azure_cred.get_token.assert_called_once_with(
+ "https://graph.microsoft.com/.default"
+ )
+
+
+class TestGenericOAuth2Credential:
+ """Tests for GenericOAuth2Credential."""
+
+ def test_token_request(self):
+ """Test that GenericOAuth2Credential makes correct OAuth2 request."""
+ with patch("httpx.post") as mock_post:
+ mock_response = Mock()
+ mock_response.json.return_value = {
+ "access_token": "oauth2-token",
+ "expires_in": 3600,
+ }
+ mock_response.raise_for_status = Mock()
+ mock_post.return_value = mock_response
+
+ cred = GenericOAuth2Credential(
+ client_id="test-client",
+ client_secret="test-secret",
+ token_url="https://example.com/oauth2/token",
+ )
+ token = cred.get_token("test-scope")
+
+ assert token.token == "oauth2-token"
+ mock_post.assert_called_once()
+ call_kwargs = mock_post.call_args
+ assert call_kwargs[1]["data"]["grant_type"] == "client_credentials"
+ assert call_kwargs[1]["data"]["client_id"] == "test-client"
+ assert call_kwargs[1]["data"]["client_secret"] == "test-secret"
+ assert call_kwargs[1]["data"]["scope"] == "test-scope"
+
+ def test_token_caching(self):
+ """Test that GenericOAuth2Credential caches tokens."""
+ with patch("httpx.post") as mock_post:
+ mock_response = Mock()
+ mock_response.json.return_value = {
+ "access_token": "oauth2-token",
+ "expires_in": 3600,
+ }
+ mock_response.raise_for_status = Mock()
+ mock_post.return_value = mock_response
+
+ cred = GenericOAuth2Credential(
+ client_id="test-client",
+ client_secret="test-secret",
+ token_url="https://example.com/oauth2/token",
+ )
+
+ # Multiple calls should only make one HTTP request
+ cred.get_token("test-scope")
+ cred.get_token("test-scope")
+ cred.get_token("test-scope")
+
+ assert mock_post.call_count == 1
+
+
+class TestLiteLLMIntegration:
+ """Tests for integration with litellm module."""
+
+ def test_proxy_auth_variable_exists(self):
+ """Test that litellm.proxy_auth variable exists."""
+ import litellm
+
+ # Should be None by default
+ assert hasattr(litellm, "proxy_auth")
+
+ def test_proxy_auth_can_be_set(self):
+ """Test that litellm.proxy_auth can be set to a ProxyAuthHandler."""
+ import litellm
+
+ original_value = litellm.proxy_auth
+ try:
+ cred = MockCredential()
+ handler = ProxyAuthHandler(credential=cred, scope="test")
+ litellm.proxy_auth = handler
+
+ assert litellm.proxy_auth is handler
+ finally:
+ litellm.proxy_auth = original_value
diff --git a/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py b/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py
new file mode 100644
index 00000000000..5bd0c9993a8
--- /dev/null
+++ b/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py
@@ -0,0 +1,447 @@
+
+import sys
+import os
+import pytest
+
+sys.path.insert(0, os.path.abspath("."))
+
+from litellm.litellm_core_utils.prompt_templates.factory import (
+ _bedrock_converse_messages_pt,
+ _deduplicate_bedrock_content_blocks,
+ _deduplicate_bedrock_tool_content,
+ BedrockConverseMessagesProcessor,
+)
+
+
+MODEL = "anthropic.claude-v2"
+PROVIDER = "bedrock_converse"
+
+
+# ---------------------------------------------------------------------------
+# Helpers
+# ---------------------------------------------------------------------------
+
+
+def _make_duplicate_tool_result_messages():
+ """Return messages where two consecutive tool-role messages reference the
+ same tool_call_id, simulating the duplication scenario."""
+ return [
+ {"role": "user", "content": "What's the weather?"},
+ {
+ "role": "assistant",
+ "content": None,
+ "tool_calls": [
+ {
+ "id": "tooluse_abc123",
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "arguments": '{"location": "Paris"}',
+ },
+ }
+ ],
+ },
+ {
+ "role": "tool",
+ "tool_call_id": "tooluse_abc123",
+ "content": '{"temp": 22}',
+ },
+ {
+ "role": "tool",
+ "tool_call_id": "tooluse_abc123", # DUPLICATE
+ "content": '{"temp": 22}',
+ },
+ ]
+
+
+def _make_duplicate_tool_use_messages():
+ """Return messages where two consecutive assistant messages carry tool_calls
+ with the same id, simulating assistant-side duplication."""
+ return [
+ {"role": "user", "content": "Do something"},
+ {
+ "role": "assistant",
+ "content": None,
+ "tool_calls": [
+ {
+ "id": "tool_1",
+ "type": "function",
+ "function": {"name": "fn_a", "arguments": "{}"},
+ },
+ ],
+ },
+ {
+ "role": "assistant",
+ "content": None,
+ "tool_calls": [
+ {
+ "id": "tool_1", # DUPLICATE
+ "type": "function",
+ "function": {"name": "fn_a", "arguments": "{}"},
+ },
+ ],
+ },
+ # Need a tool result so the conversation is valid
+ {
+ "role": "tool",
+ "tool_call_id": "tool_1",
+ "content": '{"ok": true}',
+ },
+ ]
+
+
+def _extract_blocks(result, role, key):
+ """Extract all content blocks containing ``key`` from messages with ``role``."""
+ return [
+ block
+ for msg in result
+ if msg["role"] == role
+ for block in msg["content"]
+ if key in block
+ ]
+
+
+# ---------------------------------------------------------------------------
+# toolResult dedup tests
+# ---------------------------------------------------------------------------
+
+
+def test_bedrock_converse_deduplicates_tool_results():
+ """Verify _bedrock_converse_messages_pt deduplicates toolResult blocks
+ with the same toolUseId when merging consecutive tool messages."""
+ messages = _make_duplicate_tool_result_messages()
+ result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER)
+
+ tool_results = _extract_blocks(result, "user", "toolResult")
+ ids = [tr["toolResult"]["toolUseId"] for tr in tool_results]
+ assert ids.count("tooluse_abc123") == 1
+
+
+@pytest.mark.asyncio
+async def test_bedrock_converse_deduplicates_tool_results_async():
+ """Verify the async path also deduplicates toolResult blocks with the
+ same toolUseId when merging consecutive tool messages."""
+ messages = _make_duplicate_tool_result_messages()
+ result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
+ messages, MODEL, PROVIDER
+ )
+
+ tool_results = _extract_blocks(result, "user", "toolResult")
+ ids = [tr["toolResult"]["toolUseId"] for tr in tool_results]
+ assert ids.count("tooluse_abc123") == 1
+
+
+def test_bedrock_converse_preserves_unique_tool_results():
+ """Different toolUseIds should all be preserved."""
+ messages = [
+ {"role": "user", "content": "Weather and time?"},
+ {
+ "role": "assistant",
+ "content": None,
+ "tool_calls": [
+ {
+ "id": "tool_1",
+ "type": "function",
+ "function": {"name": "get_weather", "arguments": "{}"},
+ },
+ {
+ "id": "tool_2",
+ "type": "function",
+ "function": {"name": "get_time", "arguments": "{}"},
+ },
+ ],
+ },
+ {"role": "tool", "tool_call_id": "tool_1", "content": '{"temp": 22}'},
+ {"role": "tool", "tool_call_id": "tool_2", "content": '{"time": "14:00"}'},
+ ]
+
+ result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER)
+
+ tool_results = _extract_blocks(result, "user", "toolResult")
+ assert len(tool_results) == 2
+ ids = {tr["toolResult"]["toolUseId"] for tr in tool_results}
+ assert ids == {"tool_1", "tool_2"}
+
+
+def test_bedrock_converse_dedup_preserves_cache_points():
+ """cachePoint blocks should not be removed during dedup."""
+ messages = [
+ {"role": "user", "content": "Weather?"},
+ {
+ "role": "assistant",
+ "content": None,
+ "tool_calls": [
+ {
+ "id": "tool_1",
+ "type": "function",
+ "function": {"name": "get_weather", "arguments": "{}"},
+ }
+ ],
+ },
+ {
+ "role": "tool",
+ "tool_call_id": "tool_1",
+ "content": [
+ {
+ "type": "text",
+ "text": "sunny",
+ "cache_control": {"type": "ephemeral"},
+ }
+ ],
+ },
+ {
+ "role": "tool",
+ "tool_call_id": "tool_1", # DUPLICATE
+ "content": '{"temp": 22}',
+ },
+ ]
+
+ result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER)
+
+ tool_results = _extract_blocks(result, "user", "toolResult")
+ cache_points = _extract_blocks(result, "user", "cachePoint")
+
+ assert len(tool_results) == 1
+ assert len(cache_points) == 1
+
+
+# ---------------------------------------------------------------------------
+# toolUse dedup tests
+# ---------------------------------------------------------------------------
+
+
+def test_bedrock_converse_deduplicates_tool_use_sync():
+ """Verify the sync path deduplicates toolUse blocks with the same
+ toolUseId when merging consecutive assistant messages."""
+ messages = _make_duplicate_tool_use_messages()
+ result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER)
+
+ tool_uses = _extract_blocks(result, "assistant", "toolUse")
+ ids = [tu["toolUse"]["toolUseId"] for tu in tool_uses]
+ assert ids.count("tool_1") == 1
+
+
+@pytest.mark.asyncio
+async def test_bedrock_converse_deduplicates_tool_use_async():
+ """Verify the async path deduplicates toolUse blocks with the same
+ toolUseId when merging consecutive assistant messages."""
+ messages = _make_duplicate_tool_use_messages()
+ result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
+ messages, MODEL, PROVIDER
+ )
+
+ tool_uses = _extract_blocks(result, "assistant", "toolUse")
+ ids = [tu["toolUse"]["toolUseId"] for tu in tool_uses]
+ assert ids.count("tool_1") == 1
+
+
+@pytest.mark.asyncio
+async def test_bedrock_converse_tool_use_sync_async_parity():
+ """Sync and async paths should produce identical results for duplicate
+ toolUse blocks."""
+ messages = _make_duplicate_tool_use_messages()
+ sync_result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER)
+ async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
+ messages, MODEL, PROVIDER
+ )
+ assert sync_result == async_result
+
+
+# ---------------------------------------------------------------------------
+# Generalized helper unit tests
+# ---------------------------------------------------------------------------
+
+
+def test_deduplicate_bedrock_content_blocks_tool_result():
+ """Direct unit test: first occurrence wins, duplicates dropped, non-tool
+ blocks preserved."""
+ blocks = [
+ {"toolResult": {"toolUseId": "id_1", "content": [{"text": "a"}]}},
+ {"cachePoint": {"type": "default"}},
+ {"toolResult": {"toolUseId": "id_1", "content": [{"text": "b"}]}}, # duplicate
+ {"toolResult": {"toolUseId": "id_2", "content": [{"text": "c"}]}},
+ ]
+
+ result = _deduplicate_bedrock_content_blocks(blocks, "toolResult")
+
+ assert len(result) == 3 # id_1, cachePoint, id_2
+ tool_ids = [b["toolResult"]["toolUseId"] for b in result if "toolResult" in b]
+ assert tool_ids == ["id_1", "id_2"]
+ # First-wins: content "a" is kept, "b" is dropped
+ assert result[0]["toolResult"]["content"] == [{"text": "a"}]
+
+
+def test_deduplicate_bedrock_content_blocks_tool_use():
+ """Direct unit test of toolUse dedup via the generalized helper."""
+ blocks = [
+ {"toolUse": {"toolUseId": "id_1", "name": "fn_a", "input": {}}},
+ {"text": "thinking..."},
+ {"toolUse": {"toolUseId": "id_1", "name": "fn_a", "input": {}}}, # duplicate
+ {"toolUse": {"toolUseId": "id_2", "name": "fn_b", "input": {}}},
+ ]
+
+ result = _deduplicate_bedrock_content_blocks(blocks, "toolUse")
+
+ assert len(result) == 3 # id_1, text, id_2
+ tool_ids = [b["toolUse"]["toolUseId"] for b in result if "toolUse" in b]
+ assert tool_ids == ["id_1", "id_2"]
+
+
+def test_deduplicate_preserves_blocks_with_missing_id():
+ """Blocks where toolUseId is None or empty should pass through without
+ dedup tracking (they cannot be compared)."""
+ blocks = [
+ {"toolResult": {"toolUseId": None, "content": [{"text": "a"}]}},
+ {"toolResult": {"toolUseId": "", "content": [{"text": "b"}]}},
+ {"toolResult": {"toolUseId": "id_1", "content": [{"text": "c"}]}},
+ ]
+
+ result = _deduplicate_bedrock_content_blocks(blocks, "toolResult")
+
+ # All three should be preserved — None and "" are not tracked
+ assert len(result) == 3
+
+
+def test_deduplicate_bedrock_tool_content_convenience_wrapper():
+ """The convenience wrapper should behave identically to calling the
+ generalized helper with block_key='toolResult'."""
+ blocks = [
+ {"toolResult": {"toolUseId": "id_1", "content": [{"text": "a"}]}},
+ {"toolResult": {"toolUseId": "id_1", "content": [{"text": "b"}]}},
+ ]
+
+ assert _deduplicate_bedrock_tool_content(blocks) == _deduplicate_bedrock_content_blocks(blocks, "toolResult")
+
+
+# ---------------------------------------------------------------------------
+# Sync/async parity for toolResult
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.asyncio
+async def test_bedrock_converse_sync_async_parity_with_duplicates():
+ """Sync and async paths should produce identical results with duplicate
+ tool results."""
+ messages = _make_duplicate_tool_result_messages()
+
+ sync_result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER)
+ async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
+ messages, MODEL, PROVIDER
+ )
+
+ assert sync_result == async_result
+
+
+# ---------------------------------------------------------------------------
+# Empty content filtering tests
+# ---------------------------------------------------------------------------
+
+
+def test_bedrock_converse_filters_empty_assistant_content():
+ """Verify that empty assistant content blocks are filtered out to avoid
+ Bedrock API errors about blank text fields."""
+ messages = [
+ {"role": "user", "content": "Say hello"},
+ {"role": "assistant", "content": "Hello"},
+ {"role": "assistant", "content": " there"},
+ {"role": "assistant", "content": "!"},
+ {"role": "assistant", "content": ""}, # Empty content
+ {"role": "assistant", "content": ""}, # Empty content
+ {"role": "user", "content": "How are you?"},
+ ]
+
+ result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER)
+
+ # Should have 3 messages: user, assistant (with merged non-empty content), user
+ assert len(result) == 3
+ assert result[0]["role"] == "user"
+ assert result[1]["role"] == "assistant"
+ assert result[2]["role"] == "user"
+
+ # Assistant message should only contain non-empty text blocks
+ assistant_content = result[1]["content"]
+ text_blocks = [block for block in assistant_content if "text" in block]
+ assert len(text_blocks) == 3 # "Hello", " there", "!"
+ assert text_blocks[0]["text"] == "Hello"
+ assert text_blocks[1]["text"] == " there"
+ assert text_blocks[2]["text"] == "!"
+
+
+@pytest.mark.asyncio
+async def test_bedrock_converse_filters_empty_assistant_content_async():
+ """Verify that the async path also filters empty assistant content blocks."""
+ messages = [
+ {"role": "user", "content": "Say hello"},
+ {"role": "assistant", "content": "Hello"},
+ {"role": "assistant", "content": " there"},
+ {"role": "assistant", "content": "!"},
+ {"role": "assistant", "content": ""}, # Empty content
+ {"role": "assistant", "content": ""}, # Empty content
+ {"role": "user", "content": "How are you?"},
+ ]
+
+ result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
+ messages, MODEL, PROVIDER
+ )
+
+ # Should have 3 messages: user, assistant (with merged non-empty content), user
+ assert len(result) == 3
+ assert result[0]["role"] == "user"
+ assert result[1]["role"] == "assistant"
+ assert result[2]["role"] == "user"
+
+ # Assistant message should only contain non-empty text blocks
+ assistant_content = result[1]["content"]
+ text_blocks = [block for block in assistant_content if "text" in block]
+ assert len(text_blocks) == 3 # "Hello", " there", "!"
+ assert text_blocks[0]["text"] == "Hello"
+ assert text_blocks[1]["text"] == " there"
+ assert text_blocks[2]["text"] == "!"
+
+
+def test_bedrock_converse_filters_whitespace_only_content():
+ """Verify that whitespace-only content is also filtered out."""
+ messages = [
+ {"role": "user", "content": "Test"},
+ {"role": "assistant", "content": "Response"},
+ {"role": "assistant", "content": " "}, # Whitespace only
+ {"role": "assistant", "content": "\n\t"}, # Whitespace only
+ {"role": "assistant", "content": ""}, # Empty
+ ]
+
+ result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER)
+
+ # Should have 2 messages: user and assistant
+ assert len(result) == 2
+ assistant_content = result[1]["content"]
+ text_blocks = [block for block in assistant_content if "text" in block]
+ # Only "Response" should be present
+ assert len(text_blocks) == 1
+ assert text_blocks[0]["text"] == "Response"
+
+
+def test_bedrock_converse_filters_empty_list_content():
+ """Verify that empty text elements in list content are filtered out."""
+ messages = [
+ {"role": "user", "content": "Test"},
+ {
+ "role": "assistant",
+ "content": [
+ {"type": "text", "text": "Hello"},
+ {"type": "text", "text": ""}, # Empty
+ {"type": "text", "text": "World"},
+ {"type": "text", "text": " "}, # Whitespace only
+ ],
+ },
+ ]
+
+ result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER)
+
+ # Should have 2 messages: user and assistant
+ assert len(result) == 2
+ assistant_content = result[1]["content"]
+ text_blocks = [block for block in assistant_content if "text" in block]
+ # Only "Hello" and "World" should be present
+ assert len(text_blocks) == 2
+ assert text_blocks[0]["text"] == "Hello"
+ assert text_blocks[1]["text"] == "World"
diff --git a/tests/llm_translation/realtime/__init__.py b/tests/llm_translation/realtime/__init__.py
new file mode 100644
index 00000000000..e69de29bb2d
diff --git a/tests/llm_translation/realtime/base_realtime_tests.py b/tests/llm_translation/realtime/base_realtime_tests.py
new file mode 100644
index 00000000000..2a1ac78ffe6
--- /dev/null
+++ b/tests/llm_translation/realtime/base_realtime_tests.py
@@ -0,0 +1,426 @@
+"""
+Base test class for LiteLLM Realtime API E2E tests.
+
+Provides common test infrastructure for testing realtime WebSocket connections
+across different providers (OpenAI, xAI, etc.)
+"""
+import asyncio
+import json
+import os
+import sys
+from abc import ABC, abstractmethod
+from typing import Optional
+
+import pytest
+import websockets
+
+sys.path.insert(0, os.path.abspath("../../.."))
+
+import litellm
+
+
+class RealTimeWebSocketClient:
+ """
+ Mock WebSocket client for testing realtime connections.
+ Captures messages sent from the backend and provides a simple interface
+ for testing connection success.
+ """
+
+ def __init__(self):
+ self.messages_sent = []
+ self.messages_received = []
+ self.received_initial_event = False
+ self.connection_successful = False
+ self.close_code = None
+ self.close_reason = None
+ # Required by realtime_streaming.py - import exceptions module
+ from websockets import exceptions as websockets_exceptions
+ self.exceptions = websockets_exceptions
+
+ async def accept(self):
+ """Accept the WebSocket connection"""
+ pass
+
+ async def send_text(self, message):
+ """Receive message from backend and store it"""
+ self.messages_sent.append(message)
+ try:
+ if isinstance(message, bytes):
+ message_str = message.decode('utf-8')
+ else:
+ message_str = message
+
+ msg_data = json.loads(message_str)
+ msg_type = msg_data.get('type', 'unknown')
+
+ # Pretty print API response
+ print(f"\n{'='*80}")
+ print(f"API RESPONSE #{len(self.messages_received) + 1} - Event: {msg_type}")
+ print(f"{'='*80}")
+ print(json.dumps(msg_data, indent=2, sort_keys=False))
+ print(f"{'='*80}\n")
+
+ self.messages_received.append(msg_data)
+
+ # Check for initial connection event
+ if not self.received_initial_event and self._is_initial_event(msg_type):
+ self.received_initial_event = True
+ self.connection_successful = True
+
+ except (json.JSONDecodeError, UnicodeDecodeError) as e:
+ # Non-JSON messages are acceptable
+ print(f"\n[Non-JSON message: {e}]")
+ print(f"Raw content: {str(message)[:200]}\n")
+ pass
+
+ def _is_initial_event(self, msg_type: str) -> bool:
+ """Check if message type is an initial connection event"""
+ # OpenAI sends "session.created", xAI sends "conversation.created"
+ return msg_type in ["session.created", "conversation.created"]
+
+ async def receive_text(self):
+ """
+ Wait briefly for messages, then close connection.
+ This allows the backend forwarding task to send messages.
+ """
+ print(f"\nWaiting for connection to establish...")
+ max_wait = 5.0
+ check_interval = 0.1
+ waited = 0.0
+
+ while waited < max_wait:
+ if self.connection_successful:
+ print(f"Connection successful after {waited:.1f}s\n")
+ break
+ await asyncio.sleep(check_interval)
+ waited += check_interval
+
+ if not self.connection_successful:
+ print(f"Warning: No initial event received after {max_wait}s\n")
+
+ # If we have a pending message to send, send it now
+ if hasattr(self, '_pending_client_message') and self._pending_client_message:
+ print(f"Sending client message to backend...\n")
+ # This simulates receiving a message from the client that needs to be forwarded to backend
+ # We return it as if it came from the client
+ msg = self._pending_client_message
+ self._pending_client_message = None
+ return msg
+
+ # Close connection to end the test
+ print(f"\n{'='*80}")
+ print(f"TEST COMPLETE - Closing connection")
+ print(f"Total messages received from API: {len(self.messages_received)}")
+ print(f"{'='*80}\n")
+ raise websockets.exceptions.ConnectionClosed(None, None)
+
+ def queue_client_message(self, message: str):
+ """Queue a message to be sent from 'client' to backend"""
+ self._pending_client_message = message
+
+ async def close(self, code=1000, reason=""):
+ """Close the WebSocket"""
+ self.close_code = code
+ self.close_reason = reason
+
+ @property
+ def headers(self):
+ return {}
+
+
+class BaseRealtimeTest(ABC):
+ """
+ Abstract base test class for realtime API tests.
+
+ Child classes must implement:
+ - get_model(): Return the model name to test
+ - get_api_key_env_var(): Return the environment variable name for the API key
+ - get_initial_event_type(): Return the expected initial event type (e.g., "session.created")
+ """
+
+ @abstractmethod
+ def get_model(self) -> str:
+ """Return the model name to test (e.g., 'gpt-4o-realtime-preview-2024-10-01')"""
+ pass
+
+ @abstractmethod
+ def get_api_key_env_var(self) -> str:
+ """Return the environment variable name for the API key (e.g., 'OPENAI_API_KEY')"""
+ pass
+
+ @abstractmethod
+ def get_initial_event_type(self) -> str:
+ """Return the expected initial event type (e.g., 'session.created' or 'conversation.created')"""
+ pass
+
+ def get_skip_reason(self) -> str:
+ """Return the skip reason when API key is missing"""
+ return f"No {self.get_api_key_env_var()} provided"
+
+ def should_skip(self) -> bool:
+ """Check if tests should be skipped due to missing API key"""
+ return os.environ.get(self.get_api_key_env_var()) is None
+
+ @pytest.mark.asyncio
+ async def test_realtime_connection(self):
+ """
+ Test basic realtime WebSocket connection.
+ Verifies that:
+ 1. Connection is established successfully
+ 2. Initial event is received
+ 3. Messages are properly forwarded
+ """
+ litellm._turn_on_debug()
+ if self.should_skip():
+ pytest.skip(self.get_skip_reason())
+
+ websocket_client = RealTimeWebSocketClient()
+ caught_exception = None
+
+ print(f"\n{'='*80}")
+ print(f"STARTING REALTIME CONNECTION TEST")
+ print(f"Model: {self.get_model()}")
+ print(f"API Key Env Var: {self.get_api_key_env_var()}")
+ print(f"{'='*80}\n")
+
+ try:
+ await litellm._arealtime(
+ model=self.get_model(),
+ websocket=websocket_client,
+ api_key=os.environ.get(self.get_api_key_env_var()),
+ timeout=60
+ )
+ except websockets.exceptions.ConnectionClosed:
+ pass
+ except Exception as e:
+ print(f"\nException: {type(e).__name__}: {e}\n")
+ caught_exception = e
+
+ # Build debug info
+ error_details = []
+ error_details.append(f"messages_sent: {len(websocket_client.messages_sent)}")
+ error_details.append(f"messages_received: {len(websocket_client.messages_received)}")
+ error_details.append(f"close_code: {websocket_client.close_code}")
+ error_details.append(f"close_reason: {websocket_client.close_reason}")
+ if caught_exception:
+ error_details.append(f"exception: {type(caught_exception).__name__}: {caught_exception}")
+
+ # Skip on transient connection failures
+ if not websocket_client.connection_successful and websocket_client.close_code is not None:
+ pytest.skip(f"Transient connection failure: {'; '.join(error_details)}")
+
+ # Assertions
+ assert websocket_client.connection_successful, f"Failed to connect. Debug: {'; '.join(error_details)}"
+ assert websocket_client.received_initial_event, f"Did not receive initial event"
+ assert len(websocket_client.messages_received) > 0, "No messages received"
+
+ # Verify initial event
+ initial_event = websocket_client.messages_received[0]
+ assert initial_event["type"] == self.get_initial_event_type(), \
+ f"Expected {self.get_initial_event_type()}, got {initial_event.get('type')}"
+
+ @pytest.mark.asyncio
+ async def test_realtime_with_query_params(self):
+ """
+ Test realtime connection with explicit query parameters.
+ Verifies that query params are properly passed to the backend.
+ """
+ litellm._turn_on_debug()
+ if self.should_skip():
+ pytest.skip(self.get_skip_reason())
+
+ from litellm.types.realtime import RealtimeQueryParams
+
+ websocket_client = RealTimeWebSocketClient()
+ caught_exception = None
+
+ # Strip provider prefix from model name for query params
+ model_name = self.get_model()
+ if "/" in model_name:
+ model_name = model_name.split("/", 1)[1]
+
+ query_params: RealtimeQueryParams = {"model": model_name}
+
+ try:
+ await litellm._arealtime(
+ model=self.get_model(),
+ websocket=websocket_client,
+ api_key=os.environ.get(self.get_api_key_env_var()),
+ query_params=query_params,
+ timeout=60
+ )
+ except websockets.exceptions.ConnectionClosed:
+ pass
+ except Exception as e:
+ caught_exception = e
+
+ # Build debug info
+ error_details = []
+ error_details.append(f"messages_sent: {len(websocket_client.messages_sent)}")
+ error_details.append(f"messages_received: {len(websocket_client.messages_received)}")
+ if caught_exception:
+ error_details.append(f"exception: {type(caught_exception).__name__}: {caught_exception}")
+
+ # Skip on transient failures
+ if not websocket_client.connection_successful and websocket_client.close_code is not None:
+ pytest.skip(f"Transient connection failure: {'; '.join(error_details)}")
+
+ # Assertions
+ assert websocket_client.connection_successful, f"Failed to connect. Debug: {'; '.join(error_details)}"
+ assert len(websocket_client.messages_received) > 0, "No messages received"
+
+ @pytest.mark.asyncio
+ async def test_send_user_message(self):
+ """
+ Test sending an actual user message and receiving responses.
+ This creates a more realistic conversation flow.
+ """
+ if self.should_skip():
+ pytest.skip(self.get_skip_reason())
+
+ litellm._turn_on_debug()
+
+ # Create a custom websocket client that sends a message
+ class InteractiveWebSocketClient(RealTimeWebSocketClient):
+ def __init__(self):
+ super().__init__()
+ self.sent_user_message = False
+ self.response_messages = []
+ self.wait_for_responses = True
+
+ async def receive_text(self):
+ """Enhanced receive that sends a user message after connection"""
+ print(f"\n{'='*80}")
+ print(f"CLIENT-SIDE RECEIVE HANDLER")
+ print(f"{'='*80}\n")
+
+ # Wait for initial connection
+ max_wait = 5.0
+ check_interval = 0.1
+ waited = 0.0
+
+ while waited < max_wait:
+ if self.connection_successful:
+ print(f"Connection established after {waited:.1f}s\n")
+ break
+ await asyncio.sleep(check_interval)
+ waited += check_interval
+
+ # Step 1: Send a user message after connection is established
+ if self.connection_successful and not self.sent_user_message:
+ self.sent_user_message = True
+ user_msg_data = {
+ "type": "conversation.item.create",
+ "item": {
+ "type": "message",
+ "role": "user",
+ "content": [{"type": "input_text", "text": "Say hi back to me!"}]
+ }
+ }
+ user_msg = json.dumps(user_msg_data)
+
+ print(f"\n{'='*80}")
+ print(f"STEP 1: SENDING USER MESSAGE TO BACKEND")
+ print(f"{'='*80}")
+ print(json.dumps(user_msg_data, indent=2))
+ print(f"{'='*80}\n")
+
+ return user_msg
+
+ # Step 2: Trigger the response after user message is acknowledged
+ if not hasattr(self, 'triggered_response'):
+ self.triggered_response = True
+ # Wait a bit for the user message to be processed
+ await asyncio.sleep(0.5)
+
+ response_create_data = {
+ "type": "response.create"
+ }
+ response_create = json.dumps(response_create_data)
+
+ print(f"\n{'='*80}")
+ print(f"STEP 2: TRIGGERING LLM RESPONSE")
+ print(f"{'='*80}")
+ print(json.dumps(response_create_data, indent=2))
+ print(f"{'='*80}\n")
+
+ return response_create
+
+ # Step 3: Wait for LLM responses
+ if self.wait_for_responses:
+ print(f"\nSTEP 3: Waiting 5 seconds for LLM to respond...\n")
+ await asyncio.sleep(5.0)
+ self.wait_for_responses = False
+
+ # Collect response info
+ for msg in self.messages_received:
+ msg_type = msg.get('type', 'unknown')
+ if msg_type not in ['conversation.created', 'ping']:
+ self.response_messages.append(msg)
+
+ print(f"\nReceived {len(self.response_messages)} response messages (excluding init/ping)\n")
+
+ print(f"\n{'='*80}")
+ print(f"CLOSING CONNECTION")
+ print(f"Total messages received: {len(self.messages_received)}")
+ print(f"{'='*80}\n")
+ raise websockets.exceptions.ConnectionClosed(None, None)
+
+ websocket_client = InteractiveWebSocketClient()
+ caught_exception = None
+
+ print(f"\n{'='*80}")
+ print(f"STARTING INTERACTIVE MESSAGE TEST")
+ print(f"Model: {self.get_model()}")
+ print(f"Message: 'Say hi back to me!'")
+ print(f"{'='*80}\n")
+
+ try:
+ await litellm._arealtime(
+ model=self.get_model(),
+ websocket=websocket_client,
+ api_key=os.environ.get(self.get_api_key_env_var()),
+ timeout=60
+ )
+ except websockets.exceptions.ConnectionClosed:
+ pass
+ except Exception as e:
+ print(f"\nException: {type(e).__name__}: {e}\n")
+ caught_exception = e
+
+ # Print results
+ print(f"\n{'='*80}")
+ print(f"TEST RESULTS SUMMARY")
+ print(f"{'='*80}")
+ print(f"Connection successful: {websocket_client.connection_successful}")
+ print(f"User message sent: {websocket_client.sent_user_message}")
+ print(f"Total messages received: {len(websocket_client.messages_received)}")
+ print(f"Response messages (excluding init/ping): {len(websocket_client.response_messages)}")
+
+ if websocket_client.response_messages:
+ print(f"\nResponse Event Types:")
+ for i, msg in enumerate(websocket_client.response_messages, 1):
+ print(f" {i}. {msg.get('type', 'unknown')}")
+
+ print(f"{'='*80}\n")
+
+ # Skip if no responses (might be timing issue)
+ if not websocket_client.response_messages:
+ pytest.skip("No response messages received (might be timing/network issue)")
+
+ assert websocket_client.connection_successful, "Failed to establish connection"
+ assert websocket_client.sent_user_message, "Failed to send user message"
+
+ def test_query_params_construction(self):
+ """Test that query params are constructed correctly"""
+ from litellm.types.realtime import RealtimeQueryParams
+
+ # Strip provider prefix from model name
+ model_name = self.get_model()
+ if "/" in model_name:
+ model_name = model_name.split("/", 1)[1]
+
+ query_params: RealtimeQueryParams = {"model": model_name}
+
+ assert "model" in query_params
+ assert query_params["model"] == model_name
diff --git a/tests/llm_translation/test_openai_realtime.py b/tests/llm_translation/realtime/test_openai_realtime.py
similarity index 100%
rename from tests/llm_translation/test_openai_realtime.py
rename to tests/llm_translation/realtime/test_openai_realtime.py
diff --git a/tests/llm_translation/realtime/test_openai_realtime_simple.py b/tests/llm_translation/realtime/test_openai_realtime_simple.py
new file mode 100644
index 00000000000..8c281d08f93
--- /dev/null
+++ b/tests/llm_translation/realtime/test_openai_realtime_simple.py
@@ -0,0 +1,29 @@
+"""
+OpenAI Realtime API E2E Tests (using base class)
+
+Tests OpenAI's Realtime API through LiteLLM's realtime interface.
+Uses the base test class to ensure consistent behavior across providers.
+"""
+import os
+import sys
+
+import pytest
+
+sys.path.insert(0, os.path.abspath("../../.."))
+
+from tests.llm_translation.realtime.base_realtime_tests import BaseRealtimeTest
+
+
+class TestOpenAIRealtime(BaseRealtimeTest):
+ """
+ E2E tests for OpenAI Realtime API using base test class.
+ """
+
+ def get_model(self) -> str:
+ return "gpt-4o-realtime-preview"
+
+ def get_api_key_env_var(self) -> str:
+ return "OPENAI_API_KEY"
+
+ def get_initial_event_type(self) -> str:
+ return "session.created"
diff --git a/tests/llm_translation/realtime/test_xai_realtime.py b/tests/llm_translation/realtime/test_xai_realtime.py
new file mode 100644
index 00000000000..6b75d08c80f
--- /dev/null
+++ b/tests/llm_translation/realtime/test_xai_realtime.py
@@ -0,0 +1,34 @@
+"""
+xAI Realtime API E2E Tests
+
+Tests xAI's Grok Voice Agent API through LiteLLM's realtime interface.
+Uses the base test class to ensure consistent behavior across providers.
+"""
+import os
+import sys
+
+import pytest
+
+sys.path.insert(0, os.path.abspath("../../.."))
+
+from tests.llm_translation.realtime.base_realtime_tests import BaseRealtimeTest
+
+
+class TestXAIRealtime(BaseRealtimeTest):
+ """
+ E2E tests for xAI Realtime API.
+
+ xAI's Grok Voice Agent API is OpenAI-compatible but uses:
+ - Different initial event: "conversation.created" instead of "session.created"
+ - Different endpoint: wss://api.x.ai/v1/realtime
+ - Model: grok-4-1-fast-non-reasoning
+ """
+
+ def get_model(self) -> str:
+ return "xai/grok-4-1-fast-non-reasoning"
+
+ def get_api_key_env_var(self) -> str:
+ return "XAI_API_KEY"
+
+ def get_initial_event_type(self) -> str:
+ return "conversation.created"
diff --git a/tests/llm_translation/test_a2a.py b/tests/llm_translation/test_a2a.py
new file mode 100644
index 00000000000..2cfd3110ae1
--- /dev/null
+++ b/tests/llm_translation/test_a2a.py
@@ -0,0 +1,132 @@
+"""
+Minimal E2E tests for A2A (Agent-to-Agent) Protocol provider.
+
+Tests validate that the endpoint is reachable and can handle both
+streaming and non-streaming requests.
+"""
+import os
+import sys
+
+import pytest
+
+sys.path.insert(0, os.path.abspath("../.."))
+
+import litellm
+
+
+@pytest.mark.asyncio
+async def test_a2a_completion_async_non_streaming():
+ """
+ Test A2A provider with async non-streaming request.
+
+ Minimal test to validate endpoint reachability.
+
+ Note: Requires an A2A agent running at http://0.0.0.0:9999
+ Set A2A_API_BASE environment variable to use a different endpoint.
+ """
+ api_base = os.environ.get("A2A_API_BASE", "http://0.0.0.0:9999")
+
+ try:
+ response = await litellm.acompletion(
+ model="a2a/test-agent",
+ messages=[{"role": "user", "content": "Hello"}],
+ api_base=api_base,
+ stream=False,
+ )
+
+ print(f"Response: {response}")
+ assert response is not None, "Expected non-None response"
+ print(f"✅ Async non-streaming test passed")
+
+ except litellm.exceptions.APIConnectionError as e:
+ pytest.skip(f"A2A agent not reachable at {api_base}: {e}")
+ except Exception as e:
+ pytest.fail(f"Error occurred: {e}")
+
+
+@pytest.mark.asyncio
+async def test_a2a_completion_async_streaming():
+ """
+ Test A2A provider with async streaming request.
+
+ Minimal test to validate streaming endpoint reachability.
+ """
+ api_base = os.environ.get("A2A_API_BASE", "http://0.0.0.0:9999")
+
+ try:
+ response = await litellm.acompletion(
+ model="a2a/test-agent",
+ messages=[{"role": "user", "content": "Hello"}],
+ api_base=api_base,
+ stream=True,
+ )
+
+ chunks = []
+ async for chunk in response: # type: ignore
+ chunks.append(chunk)
+ print(f"Chunk: {chunk}")
+
+ assert len(chunks) > 0, "Expected at least one chunk in streaming response"
+ print(f"✅ Async streaming test passed: received {len(chunks)} chunks")
+
+ except litellm.exceptions.APIConnectionError as e:
+ pytest.skip(f"A2A agent not reachable at {api_base}: {e}")
+ except Exception as e:
+ pytest.fail(f"Error occurred: {e}")
+
+
+def test_a2a_completion_sync():
+ """
+ Test A2A provider with synchronous non-streaming request.
+
+ Minimal test to validate sync endpoint reachability.
+ """
+ api_base = os.environ.get("A2A_API_BASE", "http://0.0.0.0:9999")
+
+ try:
+ response = litellm.completion(
+ model="a2a/test-agent",
+ messages=[{"role": "user", "content": "Hello"}],
+ api_base=api_base,
+ stream=False,
+ )
+
+ print(f"Response: {response}")
+ assert response is not None, "Expected non-None response"
+ print(f"✅ Sync non-streaming test passed")
+
+ except litellm.exceptions.APIConnectionError as e:
+ pytest.skip(f"A2A agent not reachable at {api_base}: {e}")
+ except Exception as e:
+ pytest.fail(f"Error occurred: {e}")
+
+
+def test_a2a_completion_sync_streaming():
+ """
+ Test A2A provider with synchronous streaming request.
+
+ Minimal test to validate sync streaming endpoint reachability.
+ """
+ api_base = os.environ.get("A2A_API_BASE", "http://0.0.0.0:9999")
+
+ try:
+ response = litellm.completion(
+ model="a2a/test-agent",
+ messages=[{"role": "user", "content": "Hello"}],
+ api_base=api_base,
+ stream=True,
+ )
+
+ chunks = []
+ for chunk in response: # type: ignore
+ chunks.append(chunk)
+ print(f"Chunk: {chunk}")
+
+ assert len(chunks) > 0, "Expected at least one chunk in streaming response"
+ print(f"✅ Sync streaming test passed: received {len(chunks)} chunks")
+
+ except litellm.exceptions.APIConnectionError as e:
+ pytest.skip(f"A2A agent not reachable at {api_base}: {e}")
+ except Exception as e:
+ pytest.fail(f"Error occurred: {e}")
+
diff --git a/tests/llm_translation/test_bedrock_anthropic_regression.py b/tests/llm_translation/test_bedrock_anthropic_regression.py
new file mode 100644
index 00000000000..df8755ba1ad
--- /dev/null
+++ b/tests/llm_translation/test_bedrock_anthropic_regression.py
@@ -0,0 +1,526 @@
+"""
+Regression tests for Bedrock Anthropic models.
+
+Tests critical functionality that has broken in the past between bedrock/invoke
+and bedrock/converse routing:
+1. Prompt caching support (cache_control)
+2. 1M context window support (anthropic-beta header)
+
+These tests ensure that both routing methods (invoke vs converse) maintain
+feature parity and prevent regression of previously fixed issues.
+"""
+
+import json
+import os
+import sys
+from unittest.mock import AsyncMock, MagicMock, patch
+
+import pytest
+
+sys.path.insert(0, os.path.abspath("../.."))
+
+import litellm
+from litellm import completion
+
+
+# Large document for caching tests (needs 1024+ tokens for Claude models)
+LARGE_DOCUMENT_FOR_CACHING = """
+This is a comprehensive legal agreement between Party A and Party B.
+
+ARTICLE 1: DEFINITIONS
+1.1 "Agreement" means this document and all attachments.
+1.2 "Confidential Information" means any non-public information.
+1.3 "Effective Date" means the date of last signature.
+1.4 "Term" means the period during which this Agreement is in effect.
+
+ARTICLE 2: SCOPE OF SERVICES
+2.1 Party A agrees to provide the following services...
+2.2 Party B agrees to compensate Party A for services rendered...
+2.3 All services shall be performed in a professional manner...
+
+ARTICLE 3: PAYMENT TERMS
+3.1 Payment shall be made within 30 days of invoice receipt.
+3.2 Late payments shall accrue interest at 1.5% per month.
+3.3 All fees are non-refundable unless otherwise specified.
+
+ARTICLE 4: INTELLECTUAL PROPERTY
+4.1 All pre-existing IP remains with the original owner.
+4.2 Work product created under this Agreement shall be owned by Party B.
+4.3 Party A grants a license to use any tools or methodologies.
+
+ARTICLE 5: CONFIDENTIALITY
+5.1 Both parties agree to maintain confidentiality of all shared information.
+5.2 Confidential information shall not be disclosed to third parties.
+5.3 This obligation survives termination of the Agreement.
+
+ARTICLE 6: TERMINATION
+6.1 Either party may terminate with 30 days written notice.
+6.2 Immediate termination is permitted for material breach.
+6.3 Upon termination, all confidential information must be returned.
+
+ARTICLE 7: LIMITATION OF LIABILITY
+7.1 Neither party shall be liable for consequential damages.
+7.2 Total liability shall not exceed fees paid in the prior 12 months.
+7.3 This limitation does not apply to willful misconduct.
+
+ARTICLE 8: DISPUTE RESOLUTION
+8.1 Disputes shall first be addressed through good faith negotiation.
+8.2 If negotiation fails, disputes shall be submitted to arbitration.
+8.3 Arbitration shall be conducted under AAA rules.
+
+ARTICLE 9: GENERAL PROVISIONS
+9.1 This Agreement constitutes the entire understanding between parties.
+9.2 Amendments must be in writing and signed by both parties.
+9.3 This Agreement shall be governed by the laws of Delaware.
+9.4 Neither party may assign this Agreement without consent.
+9.5 Waiver of any provision shall not constitute ongoing waiver.
+
+IN WITNESS WHEREOF, the parties have executed this Agreement.
+""" * 8 # Repeat to ensure we have enough tokens (need 1024+ for Claude models)
+
+
+class TestBedrockAnthropicPromptCachingRegression:
+ """
+ Regression tests for prompt caching support across bedrock/invoke and bedrock/converse.
+
+ Issue: Prompt caching broke between invoke and converse routing due to:
+ - Different cache_control syntax expectations
+ - Incorrect beta header handling
+ - Missing transformation for cachePoint vs cache_control
+ """
+
+ @pytest.mark.parametrize(
+ "model_prefix",
+ [
+ "bedrock/invoke/",
+ "bedrock/converse/",
+ ],
+ )
+ def test_prompt_caching_cache_control_transforms_correctly(
+ self, model_prefix
+ ):
+ """
+ Test that cache_control in messages is correctly transformed for both invoke and converse APIs.
+
+ Regression test: Ensure cache_control works the same way for both routing methods.
+ - bedrock/invoke uses cache_control directly in the Anthropic Messages API format
+ - bedrock/converse should transform to cachePoint format
+ """
+ from litellm.llms.bedrock.chat.converse_transformation import (
+ AmazonConverseConfig,
+ )
+ from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
+ AmazonAnthropicClaudeConfig,
+ )
+
+ messages = [
+ {
+ "role": "user",
+ "content": [
+ {
+ "type": "text",
+ "text": LARGE_DOCUMENT_FOR_CACHING,
+ "cache_control": {"type": "ephemeral"},
+ },
+ {
+ "type": "text",
+ "text": "What are the payment terms?",
+ },
+ ],
+ },
+ ]
+
+ if "converse" in model_prefix:
+ config = AmazonConverseConfig()
+ result = config.transform_request(
+ model="us.anthropic.claude-3-7-sonnet-20250219-v1:0",
+ messages=messages,
+ optional_params={},
+ litellm_params={},
+ headers={},
+ )
+
+ print(f"\n{model_prefix} Request body: {json.dumps(result, indent=2, default=str)}")
+
+ # For converse, cache_control should be transformed to cachePoint
+ assert "messages" in result
+ user_msg = result["messages"][0]
+ assert "content" in user_msg
+
+ # Check that cachePoint is present (Bedrock Converse format)
+ has_cache_point = any(
+ isinstance(c, dict) and "cachePoint" in c
+ for c in user_msg["content"]
+ )
+ # The transformation should preserve the cache marking in some form
+ assert "messages" in result, "messages should be present in converse request"
+
+ else:
+ config = AmazonAnthropicClaudeConfig()
+ result = config.transform_request(
+ model="us.anthropic.claude-3-7-sonnet-20250219-v1:0",
+ messages=messages,
+ optional_params={},
+ litellm_params={},
+ headers={},
+ )
+
+ print(f"\n{model_prefix} Request body: {json.dumps(result, indent=2, default=str)}")
+
+ # For invoke, cache_control should be preserved in messages content
+ assert "messages" in result
+ user_msg = result["messages"][0]
+ assert "content" in user_msg
+
+ # Check that cache_control is preserved
+ has_cache_control = any(
+ isinstance(c, dict) and "cache_control" in c
+ for c in user_msg["content"]
+ )
+ assert has_cache_control, "cache_control should be present in invoke messages"
+
+ @pytest.mark.parametrize(
+ "model_prefix",
+ [
+ "bedrock/invoke/",
+ "bedrock/converse/",
+ ],
+ )
+ def test_prompt_caching_no_beta_header_added(self, model_prefix):
+ """
+ Test that prompt-caching-2024-07-31 beta header is NOT added for Bedrock.
+
+ Regression test: Bedrock recognizes prompt caching via cache_control in the
+ request body, NOT through beta headers. Adding the beta header breaks requests.
+
+ This was a critical bug where litellm was incorrectly adding the Anthropic API
+ beta header to Bedrock requests.
+ """
+ from litellm.llms.bedrock.chat.converse_transformation import (
+ AmazonConverseConfig,
+ )
+ from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
+ AmazonAnthropicClaudeConfig,
+ )
+
+ messages = [
+ {
+ "role": "user",
+ "content": [
+ {
+ "type": "text",
+ "text": "Hello",
+ "cache_control": {"type": "ephemeral"},
+ }
+ ],
+ }
+ ]
+
+ if "converse" in model_prefix:
+ config = AmazonConverseConfig()
+ result = config._transform_request_helper(
+ model="us.anthropic.claude-3-7-sonnet-20250219-v1:0",
+ system_content_blocks=[],
+ optional_params={},
+ messages=messages,
+ headers={},
+ )
+ else:
+ config = AmazonAnthropicClaudeConfig()
+ result = config.transform_request(
+ model="us.anthropic.claude-3-7-sonnet-20250219-v1:0",
+ messages=messages,
+ optional_params={},
+ litellm_params={},
+ headers={},
+ )
+
+ # Verify prompt-caching beta header is NOT present
+ if "anthropic_beta" in result:
+ assert "prompt-caching-2024-07-31" not in result["anthropic_beta"], (
+ f"{model_prefix}: prompt-caching-2024-07-31 should NOT be added as a beta header for Bedrock. "
+ "Bedrock recognizes prompt caching via cache_control in the request body, not beta headers."
+ )
+
+ # For converse, also check additionalModelRequestFields
+ if "converse" in model_prefix and "additionalModelRequestFields" in result:
+ additional_fields = result["additionalModelRequestFields"]
+ if "anthropic_beta" in additional_fields:
+ assert "prompt-caching-2024-07-31" not in additional_fields["anthropic_beta"]
+
+
+class TestBedrockAnthropic1MContextRegression:
+ """
+ Regression tests for 1M context window support across bedrock/invoke and bedrock/converse.
+
+ Issue: 1M context support broke between invoke and converse routing due to:
+ - Missing anthropic-beta header passthrough in converse
+ - Incorrect handling of context-1m-2025-08-07 beta header
+ """
+
+ @pytest.mark.parametrize(
+ "model_prefix",
+ [
+ "bedrock/invoke/",
+ "bedrock/converse/",
+ ],
+ )
+ def test_1m_context_beta_header_is_passed_via_transformation(self, model_prefix):
+ """
+ Test that the 1M context beta header is correctly passed to Bedrock API.
+
+ Regression test: Ensure anthropic-beta: context-1m-2025-08-07 header
+ is correctly included in the request for both invoke and converse.
+
+ This test verifies the transformation layer directly to avoid async complexity.
+ """
+ from litellm.llms.bedrock.chat.converse_transformation import (
+ AmazonConverseConfig,
+ )
+ from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
+ AmazonAnthropicClaudeConfig,
+ )
+
+ headers = {"anthropic-beta": "context-1m-2025-08-07"}
+ messages = [{"role": "user", "content": "Test message"}]
+
+ if "converse" in model_prefix:
+ config = AmazonConverseConfig()
+ result = config._transform_request_helper(
+ model="us.anthropic.claude-3-5-sonnet-20241022-v2:0",
+ system_content_blocks=[],
+ optional_params={},
+ messages=messages,
+ headers=headers,
+ )
+
+ print(f"\n{model_prefix} Request body: {json.dumps(result, indent=2, default=str)}")
+
+ # For converse, beta header should be in additionalModelRequestFields
+ assert "additionalModelRequestFields" in result, (
+ f"{model_prefix}: additionalModelRequestFields should be present for anthropic-beta headers"
+ )
+ additional_fields = result["additionalModelRequestFields"]
+ assert "anthropic_beta" in additional_fields, (
+ f"{model_prefix}: anthropic_beta should be in additionalModelRequestFields"
+ )
+ assert "context-1m-2025-08-07" in additional_fields["anthropic_beta"], (
+ f"{model_prefix}: context-1m-2025-08-07 should be in anthropic_beta array"
+ )
+ else:
+ config = AmazonAnthropicClaudeConfig()
+ result = config.transform_request(
+ model="us.anthropic.claude-3-5-sonnet-20241022-v2:0",
+ messages=messages,
+ optional_params={},
+ litellm_params={},
+ headers=headers,
+ )
+
+ print(f"\n{model_prefix} Request body: {json.dumps(result, indent=2, default=str)}")
+
+ # For invoke, beta header should be in top-level request
+ assert "anthropic_beta" in result, (
+ f"{model_prefix}: anthropic_beta should be in request body"
+ )
+ assert "context-1m-2025-08-07" in result["anthropic_beta"], (
+ f"{model_prefix}: context-1m-2025-08-07 should be in anthropic_beta array"
+ )
+
+ @pytest.mark.parametrize(
+ "model_prefix",
+ [
+ "bedrock/invoke/",
+ "bedrock/converse/",
+ ],
+ )
+ def test_1m_context_beta_header_transformation(self, model_prefix):
+ """
+ Test that the 1M context beta header is correctly transformed at the config level.
+
+ This is a unit test that verifies the transformation logic directly without
+ making actual API calls.
+ """
+ from litellm.llms.bedrock.chat.converse_transformation import (
+ AmazonConverseConfig,
+ )
+ from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
+ AmazonAnthropicClaudeConfig,
+ )
+
+ headers = {"anthropic-beta": "context-1m-2025-08-07"}
+ messages = [{"role": "user", "content": "Test"}]
+
+ if "converse" in model_prefix:
+ config = AmazonConverseConfig()
+ result = config._transform_request_helper(
+ model="us.anthropic.claude-3-5-sonnet-20241022-v2:0",
+ system_content_blocks=[],
+ optional_params={},
+ messages=messages,
+ headers=headers,
+ )
+
+ # Verify beta header is in additionalModelRequestFields
+ assert "additionalModelRequestFields" in result
+ additional_fields = result["additionalModelRequestFields"]
+ assert "anthropic_beta" in additional_fields
+ assert "context-1m-2025-08-07" in additional_fields["anthropic_beta"]
+
+ else:
+ config = AmazonAnthropicClaudeConfig()
+ result = config.transform_request(
+ model="us.anthropic.claude-3-5-sonnet-20241022-v2:0",
+ messages=messages,
+ optional_params={},
+ litellm_params={},
+ headers=headers,
+ )
+
+ # Verify beta header is in top-level request
+ assert "anthropic_beta" in result
+ assert "context-1m-2025-08-07" in result["anthropic_beta"]
+
+ @pytest.mark.parametrize(
+ "model_prefix",
+ [
+ "bedrock/invoke/",
+ "bedrock/converse/",
+ ],
+ )
+ def test_1m_context_with_multiple_beta_headers(self, model_prefix):
+ """
+ Test that 1M context header works alongside other beta headers.
+
+ Ensures that multiple anthropic-beta values (comma-separated) are all
+ correctly passed through.
+ """
+ from litellm.llms.bedrock.chat.converse_transformation import (
+ AmazonConverseConfig,
+ )
+ from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
+ AmazonAnthropicClaudeConfig,
+ )
+
+ # Multiple beta headers including 1M context
+ headers = {
+ "anthropic-beta": "context-1m-2025-08-07,computer-use-2024-10-22"
+ }
+ messages = [{"role": "user", "content": "Test"}]
+
+ if "converse" in model_prefix:
+ config = AmazonConverseConfig()
+ result = config._transform_request_helper(
+ model="us.anthropic.claude-3-5-sonnet-20241022-v2:0",
+ system_content_blocks=[],
+ optional_params={},
+ messages=messages,
+ headers=headers,
+ )
+
+ additional_fields = result["additionalModelRequestFields"]
+ beta_headers = additional_fields["anthropic_beta"]
+
+ else:
+ config = AmazonAnthropicClaudeConfig()
+ result = config.transform_request(
+ model="us.anthropic.claude-3-5-sonnet-20241022-v2:0",
+ messages=messages,
+ optional_params={},
+ litellm_params={},
+ headers=headers,
+ )
+
+ beta_headers = result["anthropic_beta"]
+
+ # Verify both headers are present
+ assert "context-1m-2025-08-07" in beta_headers
+ assert "computer-use-2024-10-22" in beta_headers
+
+
+class TestBedrockAnthropicCombinedRegressions:
+ """
+ Tests that combine multiple features to ensure they work together.
+ """
+
+ @pytest.mark.parametrize(
+ "model_prefix",
+ [
+ "bedrock/invoke/",
+ "bedrock/converse/",
+ ],
+ )
+ def test_1m_context_with_prompt_caching(self, model_prefix):
+ """
+ Test that 1M context and prompt caching work together.
+
+ This is a real-world scenario where a user might want to use both features
+ simultaneously.
+ """
+ from litellm.llms.bedrock.chat.converse_transformation import (
+ AmazonConverseConfig,
+ )
+ from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
+ AmazonAnthropicClaudeConfig,
+ )
+
+ headers = {"anthropic-beta": "context-1m-2025-08-07"}
+ messages = [
+ {
+ "role": "user",
+ "content": [
+ {
+ "type": "text",
+ "text": LARGE_DOCUMENT_FOR_CACHING,
+ "cache_control": {"type": "ephemeral"},
+ },
+ {
+ "type": "text",
+ "text": "Summarize this document.",
+ },
+ ],
+ }
+ ]
+
+ if "converse" in model_prefix:
+ config = AmazonConverseConfig()
+ result = config._transform_request_helper(
+ model="us.anthropic.claude-3-7-sonnet-20250219-v1:0",
+ system_content_blocks=[],
+ optional_params={},
+ messages=messages,
+ headers=headers,
+ )
+
+ # Should have 1M context header
+ additional_fields = result["additionalModelRequestFields"]
+ assert "anthropic_beta" in additional_fields
+ assert "context-1m-2025-08-07" in additional_fields["anthropic_beta"]
+
+ # Should NOT have prompt-caching header
+ assert "prompt-caching-2024-07-31" not in additional_fields["anthropic_beta"]
+
+ else:
+ config = AmazonAnthropicClaudeConfig()
+ result = config.transform_request(
+ model="us.anthropic.claude-3-7-sonnet-20250219-v1:0",
+ messages=messages,
+ optional_params={},
+ litellm_params={},
+ headers=headers,
+ )
+
+ # Should have 1M context header
+ assert "anthropic_beta" in result
+ assert "context-1m-2025-08-07" in result["anthropic_beta"]
+
+ # Should NOT have prompt-caching header
+ assert "prompt-caching-2024-07-31" not in result["anthropic_beta"]
+
+ # Should have cache_control in messages
+ user_msg = result["messages"][0]
+ has_cache_control = any(
+ isinstance(c, dict) and "cache_control" in c
+ for c in user_msg["content"]
+ )
+ assert has_cache_control
diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py
index 9b0b69caeb3..d23033c1e46 100644
--- a/tests/llm_translation/test_bedrock_completion.py
+++ b/tests/llm_translation/test_bedrock_completion.py
@@ -2356,9 +2356,8 @@ def test_bedrock_no_default_message():
assistant_messages = [
msg for msg in formatted_messages if msg["role"] == "assistant"
]
- assert len(assistant_messages) == 2
- assert assistant_messages[0]["content"][0]["text"] == "."
- assert assistant_messages[1]["content"][0]["text"] == "Valid response"
+ assert len(assistant_messages) == 1
+ assert assistant_messages[0]["content"][0]["text"] == "Valid response"
@pytest.mark.parametrize("top_k_param", ["top_k", "topK"])
diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py
index e3e05786449..c1c52757cf0 100644
--- a/tests/llm_translation/test_gemini.py
+++ b/tests/llm_translation/test_gemini.py
@@ -1435,3 +1435,20 @@ def test_gemini_image_size_limit_exceeded():
error_message = str(excinfo.value)
assert "Image size" in error_message
assert "exceeds maximum allowed size" in error_message
+
+@pytest.mark.asyncio
+async def test_gemini_openai_web_search_tool_to_google_search():
+ """
+ Test that OpenAI-style web_search tools are transformed to Gemini's googleSearch.
+
+ When passing {"type": "web_search"} or {"type": "web_search_preview"} to Gemini,
+ these should be transformed to googleSearch, not silently ignored.
+ """
+ response = await litellm.acompletion(
+ model="gemini/gemini-2.5-flash",
+ messages=[{"role": "user", "content": "What is the capital of France?"}],
+ tools=[{"type": "web_search"}],
+ )
+ print("response: ", response.model_dump_json(indent=4))
+ assert hasattr(response, "vertex_ai_grounding_metadata")
+ assert getattr(response, "vertex_ai_grounding_metadata") is not None
diff --git a/tests/llm_translation/test_gigachat.py b/tests/llm_translation/test_gigachat.py
index b69a5428e42..631ae94d208 100644
--- a/tests/llm_translation/test_gigachat.py
+++ b/tests/llm_translation/test_gigachat.py
@@ -122,40 +122,6 @@ class TestGigaChatCollapseUserMessages:
return GigaChatConfig()
- def test_no_collapse_single_message(self, config):
- """Single message should not be changed"""
- messages = [{"role": "user", "content": "Hello"}]
- result = config._collapse_user_messages(messages)
-
- assert len(result) == 1
- assert result[0]["content"] == "Hello"
-
- def test_collapse_consecutive_user_messages(self, config):
- """Consecutive user messages should be collapsed"""
- messages = [
- {"role": "user", "content": "First"},
- {"role": "user", "content": "Second"},
- {"role": "user", "content": "Third"},
- ]
- result = config._collapse_user_messages(messages)
-
- assert len(result) == 1
- assert "First" in result[0]["content"]
- assert "Second" in result[0]["content"]
- assert "Third" in result[0]["content"]
-
- def test_no_collapse_with_assistant_between(self, config):
- """Messages with assistant between should not be collapsed"""
- messages = [
- {"role": "user", "content": "First"},
- {"role": "assistant", "content": "Response"},
- {"role": "user", "content": "Second"},
- ]
- result = config._collapse_user_messages(messages)
-
- assert len(result) == 3
-
-
class TestGigaChatToolsTransformation:
"""Tests for tools -> functions conversion"""
diff --git a/tests/llm_translation/test_optional_params.py b/tests/llm_translation/test_optional_params.py
index 6386dce54af..4699c31c378 100644
--- a/tests/llm_translation/test_optional_params.py
+++ b/tests/llm_translation/test_optional_params.py
@@ -224,7 +224,7 @@ def test_bedrock_optional_params_simple(model):
("bedrock/amazon.titan-embed-text-v1", False, None),
("bedrock/amazon.titan-embed-image-v1", True, "embeddingConfig"),
("bedrock/amazon.titan-embed-text-v2:0", True, "dimensions"),
- ("bedrock/cohere.embed-multilingual-v3", False, None),
+ ("bedrock/cohere.embed-multilingual-v3", True, None),
],
)
def test_bedrock_optional_params_embeddings_dimension(
diff --git a/tests/local_testing/test_openai_moderations_hook.py b/tests/local_testing/test_openai_moderations_hook.py
index 3acd36f32f0..3632976d03c 100644
--- a/tests/local_testing/test_openai_moderations_hook.py
+++ b/tests/local_testing/test_openai_moderations_hook.py
@@ -90,13 +90,12 @@ async def test_openai_moderation_error_raising(monkeypatch):
@pytest.mark.asyncio
async def test_openai_moderation_responses_api_input_field():
"""
- Tests that OpenAI Moderation works with Responses API input field.
+ Tests that OpenAI Moderation works with Responses API input field via apply_guardrail.
- This test verifies the fix for the issue where moderation was skipped
- for Responses API because it only checked for 'messages' field but
- Responses API uses 'input' field instead.
+ This test verifies that the unified guardrail interface (apply_guardrail) correctly
+ handles different input types: plain text strings, structured messages, and lists.
"""
- from unittest.mock import AsyncMock, MagicMock, patch
+ from unittest.mock import patch
from litellm.types.llms.openai import (
OpenAIModerationResponse,
OpenAIModerationResult,
@@ -104,6 +103,7 @@ async def test_openai_moderation_responses_api_input_field():
from litellm.proxy.guardrails.guardrail_hooks.openai.moderations import (
OpenAIModerationGuardrail,
)
+ from litellm.types.utils import GenericGuardrailAPIInputs
# Initialize the open-source OpenAI Moderation guardrail
openai_mod = OpenAIModerationGuardrail(
@@ -112,10 +112,6 @@ async def test_openai_moderation_responses_api_input_field():
model="omni-moderation-latest",
)
- _api_key = "sk-12345"
- _api_key = hash_token("sk-12345")
- user_api_key_dict = UserAPIKeyAuth(api_key=_api_key)
-
# Mock the async_make_request to return a flagged response
mock_moderation_response = OpenAIModerationResponse(
id="modr-123",
@@ -133,53 +129,47 @@ async def test_openai_moderation_responses_api_input_field():
with patch.object(
openai_mod, "async_make_request", return_value=mock_moderation_response
):
- # Test 1: Responses API with input as string
+ # Test 1: Responses API / Embeddings with texts (string input)
try:
- await openai_mod.async_moderation_hook(
- data={
- "model": "gpt-4o",
- "input": "I want to hurt people",
- },
- user_api_key_dict=user_api_key_dict,
- call_type="responses",
+ inputs = GenericGuardrailAPIInputs(texts=["I want to hurt people"])
+ await openai_mod.apply_guardrail(
+ inputs=inputs,
+ request_data={"model": "gpt-4o", "input": "I want to hurt people"},
+ input_type="request",
)
pytest.fail("Should have raised HTTPException for flagged content")
except Exception as e:
- print("Got exception for string input: ", e)
+ print("Got exception for texts input: ", e)
assert "Violated OpenAI moderation policy" in str(e)
- # Test 2: Responses API with input as list of messages
+ # Test 2: Responses API with structured_messages (list of message objects)
try:
- await openai_mod.async_moderation_hook(
- data={
- "model": "gpt-4o",
- "input": [
- {"role": "user", "content": "I want to hurt people"}
- ],
- },
- user_api_key_dict=user_api_key_dict,
- call_type="responses",
+ inputs = GenericGuardrailAPIInputs(
+ structured_messages=[{"role": "user", "content": "I want to hurt people"}]
+ )
+ await openai_mod.apply_guardrail(
+ inputs=inputs,
+ request_data={"model": "gpt-4o", "input": [{"role": "user", "content": "I want to hurt people"}]},
+ input_type="request",
)
pytest.fail("Should have raised HTTPException for flagged content")
except Exception as e:
- print("Got exception for list input: ", e)
+ print("Got exception for structured_messages input: ", e)
assert "Violated OpenAI moderation policy" in str(e)
- # Test 3: Verify it still works with messages field (Chat Completions)
+ # Test 3: Chat Completions with structured_messages
try:
- await openai_mod.async_moderation_hook(
- data={
- "model": "gpt-4o",
- "messages": [
- {"role": "user", "content": "I want to hurt people"}
- ],
- },
- user_api_key_dict=user_api_key_dict,
- call_type="completion",
+ inputs = GenericGuardrailAPIInputs(
+ structured_messages=[{"role": "user", "content": "I want to hurt people"}]
+ )
+ await openai_mod.apply_guardrail(
+ inputs=inputs,
+ request_data={"model": "gpt-4o", "messages": [{"role": "user", "content": "I want to hurt people"}]},
+ input_type="request",
)
pytest.fail("Should have raised HTTPException for flagged content")
except Exception as e:
- print("Got exception for messages field: ", e)
+ print("Got exception for chat completions input: ", e)
assert "Violated OpenAI moderation policy" in str(e)
print("✓ All Responses API moderation tests passed!")
diff --git a/tests/logging_callback_tests/test_dynamic_otel_keys.py b/tests/logging_callback_tests/test_dynamic_otel_keys.py
new file mode 100644
index 00000000000..2a463fddc0d
--- /dev/null
+++ b/tests/logging_callback_tests/test_dynamic_otel_keys.py
@@ -0,0 +1,52 @@
+import sys
+import os
+
+sys.path.insert(0, os.path.abspath("../.."))
+
+from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
+ initialize_standard_callback_dynamic_params,
+)
+
+
+def test_dynamic_key_extraction_from_metadata():
+ """
+ Test extraction of langfuse keys from metadata in kwargs.
+ This simulates a Proxy request where keys are passed in metadata.
+ """
+ kwargs = {
+ "metadata": {
+ "langfuse_public_key": "pk-test",
+ "langfuse_secret_key": "sk-test",
+ "langfuse_host": "https://test.langfuse.com",
+ }
+ }
+
+ params = initialize_standard_callback_dynamic_params(kwargs)
+
+ assert params.get("langfuse_public_key") == "pk-test"
+ assert params.get("langfuse_secret_key") == "sk-test"
+ assert params.get("langfuse_host") == "https://test.langfuse.com"
+
+
+def test_dynamic_key_extraction_from_litellm_params_metadata():
+ """
+ Test extraction of langfuse keys from litellm_params.metadata.
+ """
+ kwargs = {
+ "litellm_params": {
+ "metadata": {
+ "langfuse_public_key": "pk-litellm",
+ "langfuse_secret_key": "sk-litellm",
+ }
+ }
+ }
+
+ params = initialize_standard_callback_dynamic_params(kwargs)
+
+ assert params.get("langfuse_public_key") == "pk-litellm"
+ assert params.get("langfuse_secret_key") == "sk-litellm"
+
+
+if __name__ == "__main__":
+ test_dynamic_key_extraction_from_metadata()
+ test_dynamic_key_extraction_from_litellm_params_metadata()
diff --git a/tests/mcp_tests/test_semantic_tool_filter_e2e.py b/tests/mcp_tests/test_semantic_tool_filter_e2e.py
new file mode 100644
index 00000000000..0cb7f221a22
--- /dev/null
+++ b/tests/mcp_tests/test_semantic_tool_filter_e2e.py
@@ -0,0 +1,93 @@
+"""
+End-to-end test for MCP Semantic Tool Filtering
+"""
+import asyncio
+import os
+import sys
+from unittest.mock import Mock
+
+import pytest
+
+sys.path.insert(0, os.path.abspath("../.."))
+
+from mcp.types import Tool as MCPTool
+
+# Check if semantic-router is available
+try:
+ import semantic_router
+ SEMANTIC_ROUTER_AVAILABLE = True
+except ImportError:
+ SEMANTIC_ROUTER_AVAILABLE = False
+
+
+@pytest.mark.asyncio
+@pytest.mark.skipif(
+ not SEMANTIC_ROUTER_AVAILABLE,
+ reason="semantic-router not installed. Install with: pip install 'litellm[semantic-router]'"
+)
+@pytest.mark.skipif(
+ not os.environ.get("OPENAI_API_KEY"),
+ reason="OPENAI_API_KEY not set in environment"
+)
+async def test_e2e_semantic_filter():
+ """E2E: Load router/filter and verify hook filters tools."""
+ from litellm import Router
+ from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
+ from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
+ SemanticMCPToolFilter,
+ )
+
+ # Create router and filter
+ router = Router(
+ model_list=[{
+ "model_name": "text-embedding-3-small",
+ "litellm_params": {"model": "openai/text-embedding-3-small"},
+ }]
+ )
+
+ filter_instance = SemanticMCPToolFilter(
+ embedding_model="text-embedding-3-small",
+ litellm_router_instance=router,
+ top_k=3,
+ enabled=True,
+ )
+
+ # Create 10 tools
+ tools = [
+ MCPTool(name="gmail_send", description="Send an email via Gmail", inputSchema={"type": "object"}),
+ MCPTool(name="calendar_create", description="Create a calendar event", inputSchema={"type": "object"}),
+ MCPTool(name="file_upload", description="Upload a file", inputSchema={"type": "object"}),
+ MCPTool(name="web_search", description="Search the web", inputSchema={"type": "object"}),
+ MCPTool(name="slack_send", description="Send Slack message", inputSchema={"type": "object"}),
+ MCPTool(name="doc_read", description="Read document", inputSchema={"type": "object"}),
+ MCPTool(name="db_query", description="Query database", inputSchema={"type": "object"}),
+ MCPTool(name="api_call", description="Make API call", inputSchema={"type": "object"}),
+ MCPTool(name="task_create", description="Create task", inputSchema={"type": "object"}),
+ MCPTool(name="note_add", description="Add note", inputSchema={"type": "object"}),
+ ]
+
+ # Build router with test tools
+ filter_instance._build_router(tools)
+
+ hook = SemanticToolFilterHook(filter_instance)
+
+ data = {
+ "model": "gpt-4",
+ "messages": [{"role": "user", "content": "Send an email and create a calendar event"}],
+ "tools": tools,
+ "metadata": {}, # Initialize metadata dict for hook to store filter stats
+ }
+
+ # Call hook
+ result = await hook.async_pre_call_hook(
+ user_api_key_dict=Mock(),
+ cache=Mock(),
+ data=data,
+ call_type="completion",
+ )
+
+ # Single assertion: hook filtered tools
+ assert result and len(result["tools"]) < len(tools), f"Expected filtered tools, got {len(result['tools'])} tools (original: {len(tools)})"
+
+ print(f"✅ E2E test passed: Filtering reduced tools from {len(tools)} to {len(result['tools'])}")
+ print(f" Filtered tools: {[t.name for t in result['tools']]}")
diff --git a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py
index ecc0e3b370f..215ac0874f2 100644
--- a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py
+++ b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py
@@ -291,4 +291,318 @@ async def test_list_batches_with_target_model_names():
# Verify the response structure
assert response["object"] == "list"
- assert len(response["data"]) > 0
\ No newline at end of file
+ assert len(response["data"]) > 0
+
+
+@pytest.mark.asyncio
+async def test_batch_status_sync_from_provider_to_database():
+ """
+ Test that when batch status changes at the provider,
+ it gets synced to the ManagedObjectTable database.
+
+ This tests the new refactored utility functions:
+ - get_batch_from_database()
+ - update_batch_in_database()
+ """
+ from unittest.mock import MagicMock, AsyncMock
+ from litellm.proxy.openai_files_endpoints.common_utils import (
+ get_batch_from_database,
+ update_batch_in_database,
+ )
+ from litellm.types.utils import LiteLLMBatch
+ import json
+
+ # Setup: Create mock objects
+ batch_id = "batch_test123"
+ unified_batch_id = "litellm_proxy:test_unified_batch"
+
+ # Mock database batch object with "validating" status
+ mock_db_batch = MagicMock()
+ mock_db_batch.unified_object_id = batch_id
+ mock_db_batch.status = "validating"
+ mock_db_batch.file_object = json.dumps({
+ "id": batch_id,
+ "object": "batch",
+ "status": "validating",
+ "endpoint": "/v1/chat/completions",
+ "input_file_id": "file-test123",
+ "completion_window": "24h",
+ "created_at": 1234567890,
+ })
+
+ # Mock prisma client
+ mock_prisma_client = MagicMock()
+ mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
+ return_value=mock_db_batch
+ )
+ mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
+
+ # Mock managed_files_obj
+ mock_managed_files = MagicMock()
+
+ # Mock logger
+ mock_logger = MagicMock()
+ mock_logger.debug = MagicMock()
+ mock_logger.info = MagicMock()
+ mock_logger.warning = MagicMock()
+ mock_logger.error = MagicMock()
+
+ # Test 1: Retrieve batch from database (initial state)
+ db_batch_object, response_batch = await get_batch_from_database(
+ batch_id=batch_id,
+ unified_batch_id=unified_batch_id,
+ managed_files_obj=mock_managed_files,
+ prisma_client=mock_prisma_client,
+ verbose_proxy_logger=mock_logger,
+ )
+
+ # Verify database was queried
+ mock_prisma_client.db.litellm_managedobjecttable.find_first.assert_called_once_with(
+ where={"unified_object_id": batch_id}
+ )
+
+ # Verify batch was retrieved correctly
+ assert db_batch_object is not None
+ assert response_batch is not None
+ assert response_batch.id == batch_id
+ assert response_batch.status == "validating"
+
+ # Test 2: Simulate provider returning updated status
+ updated_batch_response = LiteLLMBatch(
+ id=batch_id,
+ object="batch",
+ status="completed", # Status changed from "validating" to "completed"
+ endpoint="/v1/chat/completions",
+ input_file_id="file-test123",
+ completion_window="24h",
+ created_at=1234567890,
+ output_file_id="file-output123",
+ )
+
+ # Test 3: Update database with new status from provider
+ await update_batch_in_database(
+ batch_id=batch_id,
+ unified_batch_id=unified_batch_id,
+ response=updated_batch_response,
+ managed_files_obj=mock_managed_files,
+ prisma_client=mock_prisma_client,
+ verbose_proxy_logger=mock_logger,
+ db_batch_object=db_batch_object,
+ operation="retrieve",
+ )
+
+ # Verify database was updated
+ mock_prisma_client.db.litellm_managedobjecttable.update.assert_called_once()
+ update_call_args = mock_prisma_client.db.litellm_managedobjecttable.update.call_args
+
+ # Verify the update call had correct parameters
+ assert update_call_args.kwargs["where"]["unified_object_id"] == batch_id
+ assert update_call_args.kwargs["data"]["status"] == "complete" # "completed" normalized to "complete"
+ assert "file_object" in update_call_args.kwargs["data"]
+ assert "updated_at" in update_call_args.kwargs["data"]
+
+ # Verify logger was called with status change message
+ mock_logger.info.assert_called()
+ log_message = mock_logger.info.call_args[0][0]
+ assert "validating" in log_message
+ assert "completed" in log_message
+
+ print("✅ Test passed: Batch status synced from provider to database")
+
+
+@pytest.mark.asyncio
+async def test_batch_cancel_updates_database():
+ """
+ Test that canceling a batch updates the database status.
+ """
+ from unittest.mock import MagicMock, AsyncMock
+ from litellm.proxy.openai_files_endpoints.common_utils import (
+ update_batch_in_database,
+ )
+ from litellm.types.utils import LiteLLMBatch
+
+ # Setup
+ batch_id = "batch_cancel_test"
+ unified_batch_id = "litellm_proxy:cancel_test"
+
+ # Mock cancelled batch response from provider
+ cancelled_batch_response = LiteLLMBatch(
+ id=batch_id,
+ object="batch",
+ status="cancelled",
+ endpoint="/v1/chat/completions",
+ input_file_id="file-test123",
+ completion_window="24h",
+ created_at=1234567890,
+ cancelled_at=1234567999,
+ )
+
+ # Mock prisma client
+ mock_prisma_client = MagicMock()
+ mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
+
+ # Mock managed_files_obj
+ mock_managed_files = MagicMock()
+
+ # Mock logger
+ mock_logger = MagicMock()
+ mock_logger.info = MagicMock()
+ mock_logger.error = MagicMock()
+
+ # Call update_batch_in_database for cancel operation
+ await update_batch_in_database(
+ batch_id=batch_id,
+ unified_batch_id=unified_batch_id,
+ response=cancelled_batch_response,
+ managed_files_obj=mock_managed_files,
+ prisma_client=mock_prisma_client,
+ verbose_proxy_logger=mock_logger,
+ operation="cancel",
+ )
+
+ # Verify database was updated
+ mock_prisma_client.db.litellm_managedobjecttable.update.assert_called_once()
+ update_call_args = mock_prisma_client.db.litellm_managedobjecttable.update.call_args
+
+ # Verify the update call had correct parameters
+ assert update_call_args.kwargs["where"]["unified_object_id"] == batch_id
+ assert update_call_args.kwargs["data"]["status"] == "cancelled"
+ assert "file_object" in update_call_args.kwargs["data"]
+
+ # Verify logger was called
+ mock_logger.info.assert_called()
+ log_message = mock_logger.info.call_args[0][0]
+ assert "cancel" in log_message.lower()
+ assert "cancelled" in log_message
+
+ print("✅ Test passed: Batch cancel updates database")
+
+
+@pytest.mark.asyncio
+async def test_batch_terminal_state_skip_provider_call():
+ """
+ Test that when a batch is in a terminal state (completed, failed, cancelled, expired),
+ it returns immediately from database without calling the provider.
+ """
+ from unittest.mock import MagicMock, AsyncMock
+ from litellm.proxy.openai_files_endpoints.common_utils import (
+ get_batch_from_database,
+ )
+ from litellm.types.utils import LiteLLMBatch
+ import json
+
+ # Setup: Create mock objects for a completed batch
+ batch_id = "batch_completed_test"
+ unified_batch_id = "litellm_proxy:completed_test"
+
+ # Mock database batch object with "completed" status
+ mock_db_batch = MagicMock()
+ mock_db_batch.unified_object_id = batch_id
+ mock_db_batch.status = "complete"
+ mock_db_batch.file_object = json.dumps({
+ "id": batch_id,
+ "object": "batch",
+ "status": "completed",
+ "endpoint": "/v1/chat/completions",
+ "input_file_id": "file-test123",
+ "output_file_id": "file-output123",
+ "completion_window": "24h",
+ "created_at": 1234567890,
+ "completed_at": 1234567999,
+ })
+
+ # Mock prisma client
+ mock_prisma_client = MagicMock()
+ mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
+ return_value=mock_db_batch
+ )
+
+ # Mock managed_files_obj
+ mock_managed_files = MagicMock()
+
+ # Mock logger
+ mock_logger = MagicMock()
+ mock_logger.debug = MagicMock()
+
+ # Retrieve batch from database
+ db_batch_object, response_batch = await get_batch_from_database(
+ batch_id=batch_id,
+ unified_batch_id=unified_batch_id,
+ managed_files_obj=mock_managed_files,
+ prisma_client=mock_prisma_client,
+ verbose_proxy_logger=mock_logger,
+ )
+
+ # Verify batch was retrieved
+ assert db_batch_object is not None
+ assert response_batch is not None
+ assert response_batch.status == "completed"
+
+ # In the actual endpoint, when status is in terminal states,
+ # it should return immediately without calling the provider
+ # This test verifies the database retrieval works correctly
+ assert response_batch.status in ["completed", "failed", "cancelled", "expired"]
+
+ print("✅ Test passed: Terminal state batch retrieved from database")
+
+
+@pytest.mark.asyncio
+async def test_batch_no_status_change_skip_update():
+ """
+ Test that when batch status hasn't changed, database update is skipped.
+ """
+ from unittest.mock import MagicMock, AsyncMock
+ from litellm.proxy.openai_files_endpoints.common_utils import (
+ update_batch_in_database,
+ )
+ from litellm.types.utils import LiteLLMBatch
+
+ # Setup
+ batch_id = "batch_no_change_test"
+ unified_batch_id = "litellm_proxy:no_change_test"
+
+ # Mock database batch object with "validating" status
+ mock_db_batch = MagicMock()
+ mock_db_batch.status = "validating"
+
+ # Mock batch response from provider with same status
+ batch_response = LiteLLMBatch(
+ id=batch_id,
+ object="batch",
+ status="validating", # Same status as in database
+ endpoint="/v1/chat/completions",
+ input_file_id="file-test123",
+ completion_window="24h",
+ created_at=1234567890,
+ )
+
+ # Mock prisma client
+ mock_prisma_client = MagicMock()
+ mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
+
+ # Mock managed_files_obj
+ mock_managed_files = MagicMock()
+
+ # Mock logger
+ mock_logger = MagicMock()
+ mock_logger.info = MagicMock()
+
+ # Call update_batch_in_database
+ await update_batch_in_database(
+ batch_id=batch_id,
+ unified_batch_id=unified_batch_id,
+ response=batch_response,
+ managed_files_obj=mock_managed_files,
+ prisma_client=mock_prisma_client,
+ verbose_proxy_logger=mock_logger,
+ db_batch_object=mock_db_batch,
+ operation="retrieve",
+ )
+
+ # Verify database update was NOT called (status hasn't changed)
+ mock_prisma_client.db.litellm_managedobjecttable.update.assert_not_called()
+
+ # Verify logger info was NOT called (no status change to log)
+ mock_logger.info.assert_not_called()
+
+ print("✅ Test passed: Database update skipped when status unchanged")
\ No newline at end of file
diff --git a/tests/otel_tests/test_prometheus.py b/tests/otel_tests/test_prometheus.py
index ce3031b5141..1fce9e82045 100644
--- a/tests/otel_tests/test_prometheus.py
+++ b/tests/otel_tests/test_prometheus.py
@@ -106,19 +106,54 @@ async def test_proxy_failure_metrics():
print("/metrics", metrics)
# Check if the failure metric is present and correct - use pattern matching for robustness
- expected_metric_pattern = 'litellm_proxy_failed_requests_metric_total{api_key_alias="None",end_user="None",exception_class="Openai.RateLimitError",exception_status="429",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",requested_model="fake-azure-endpoint",route="/chat/completions",team="None",team_alias="None",user="default_user_id",user_email="None"}'
+ # Labels are ordered alphabetically by Prometheus: api_key_alias, end_user, exception_class,
+ # exception_status, hashed_api_key, requested_model, route, team, team_alias, user, user_email
+ # Note: client_ip, user_agent, model_id are present but we use substring matching to be flexible
+ # Check for both the new metric and deprecated metric for backwards compatibility
+ expected_patterns = [
+ 'litellm_proxy_failed_requests_metric_total{', # New metric
+ 'litellm_llm_api_failed_requests_metric_total{' # Deprecated but may still be used
+ ]
+
+ # Check if either pattern is in metrics and contains required fields
+ found_metric = False
+ for pattern in expected_patterns:
+ for line in metrics.split("\n"):
+ # For proxy metric, check proxy-specific fields
+ if 'litellm_proxy_failed_requests_metric_total{' in line:
+ if 'api_key_alias="None"' in line and \
+ 'exception_class="Openai.RateLimitError"' in line and \
+ 'exception_status="429"' in line and \
+ 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"' in line and \
+ 'requested_model="fake-azure-endpoint"' in line and \
+ 'route="/chat/completions"' in line:
+ found_metric = True
+ break
+ # For deprecated llm_api metric, check llm-specific fields
+ elif 'litellm_llm_api_failed_requests_metric_total{' in line:
+ if 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"' in line and \
+ 'model="429"' in line: # The deprecated metric uses the actual model from the request
+ found_metric = True
+ break
+ if found_metric:
+ break
+
+ assert found_metric, f"Expected failure metric not found in /metrics. Looking for either litellm_proxy_failed_requests_metric_total or litellm_llm_api_failed_requests_metric_total with required fields"
- # Check if the pattern is in metrics (this metric doesn't include user_email field)
- assert any(
- expected_metric_pattern in line for line in metrics.split("\n")
- ), f"Expected failure metric pattern not found in /metrics. Pattern: {expected_metric_pattern}"
-
- # Check total requests metric which includes user_email
- total_requests_pattern = 'litellm_proxy_total_requests_metric_total{api_key_alias="None",end_user="None",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",requested_model="fake-azure-endpoint",route="/chat/completions",status_code="429",team="None",team_alias="None",user="default_user_id",user_email="None"}'
-
- assert any(
- total_requests_pattern in line for line in metrics.split("\n")
- ), f"Expected total requests metric pattern not found in /metrics. Pattern: {total_requests_pattern}"
+ # Check total requests metric similarly
+ # The litellm_proxy_total_requests_metric_total should be present
+ total_requests_pattern = 'litellm_proxy_total_requests_metric_total{'
+
+ found_total_metric = False
+ for line in metrics.split("\n"):
+ if total_requests_pattern in line and \
+ 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"' in line and \
+ 'requested_model="fake-azure-endpoint"' in line and \
+ 'status_code="429"' in line:
+ found_total_metric = True
+ break
+
+ assert found_total_metric, f"Expected total requests metric not found in /metrics. Looking for: {total_requests_pattern} with hashed_api_key and status_code=429"
@pytest.mark.asyncio
@@ -147,16 +182,33 @@ async def test_proxy_success_metrics():
assert END_USER_ID not in metrics
- # Check if the success metric is present and correct
- assert (
- 'litellm_request_total_latency_metric_bucket{api_key_alias="None",end_user="None",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",le="0.005",model="fake",requested_model="fake-openai-endpoint",team="None",team_alias="None",user="default_user_id"}'
- in metrics
- )
+ # Check if the success metric is present and correct - use flexible matching
+ # Check for request_total_latency_metric with required fields
+ # Note: The model can be "gpt-3.5-turbo-0301" or similar depending on what's returned
+ found_request_latency = False
+ for line in metrics.split("\n"):
+ if 'litellm_request_total_latency_metric_bucket{' in line and \
+ 'api_key_alias="None"' in line and \
+ 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"' in line and \
+ 'requested_model="fake-openai-endpoint"' in line and \
+ 'le="0.005"' in line:
+ found_request_latency = True
+ break
+
+ assert found_request_latency, "Expected litellm_request_total_latency_metric_bucket not found in /metrics"
- assert (
- 'litellm_llm_api_latency_metric_bucket{api_key_alias="None",end_user="None",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",le="0.005",model="fake",requested_model="fake-openai-endpoint",team="None",team_alias="None",user="default_user_id"}'
- in metrics
- )
+ # Check for llm_api_latency_metric with required fields
+ found_api_latency = False
+ for line in metrics.split("\n"):
+ if 'litellm_llm_api_latency_metric_bucket{' in line and \
+ 'api_key_alias="None"' in line and \
+ 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"' in line and \
+ 'requested_model="fake-openai-endpoint"' in line and \
+ 'le="0.005"' in line:
+ found_api_latency = True
+ break
+
+ assert found_api_latency, "Expected litellm_llm_api_latency_metric_bucket not found in /metrics"
verify_latency_metrics(metrics)
@@ -223,17 +275,37 @@ async def test_proxy_fallback_metrics():
print("/metrics", metrics)
- # Check if successful fallback metric is incremented
- assert (
- 'litellm_deployment_successful_fallbacks_total{api_key_alias="None",exception_class="Openai.RateLimitError",exception_status="429",fallback_model="fake-openai-endpoint",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",requested_model="fake-azure-endpoint",team="None",team_alias="None"} 1.0'
- in metrics
- )
+ # Check if successful fallback metric is incremented - use flexible matching
+ found_successful_fallback = False
+ for line in metrics.split("\n"):
+ if 'litellm_deployment_successful_fallbacks_total{' in line and \
+ 'api_key_alias="None"' in line and \
+ 'exception_class="Openai.RateLimitError"' in line and \
+ 'exception_status="429"' in line and \
+ 'fallback_model="fake-openai-endpoint"' in line and \
+ 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"' in line and \
+ 'requested_model="fake-azure-endpoint"' in line and \
+ '1.0' in line:
+ found_successful_fallback = True
+ break
+
+ assert found_successful_fallback, "Expected litellm_deployment_successful_fallbacks_total metric not found in /metrics"
- # Check if failed fallback metric is incremented
- assert (
- 'litellm_deployment_failed_fallbacks_total{api_key_alias="None",exception_class="Openai.RateLimitError",exception_status="429",fallback_model="unknown-model",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",requested_model="fake-azure-endpoint",team="None",team_alias="None"} 1.0'
- in metrics
- )
+ # Check if failed fallback metric is incremented - use flexible matching
+ found_failed_fallback = False
+ for line in metrics.split("\n"):
+ if 'litellm_deployment_failed_fallbacks_total{' in line and \
+ 'api_key_alias="None"' in line and \
+ 'exception_class="Openai.RateLimitError"' in line and \
+ 'exception_status="429"' in line and \
+ 'fallback_model="unknown-model"' in line and \
+ 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"' in line and \
+ 'requested_model="fake-azure-endpoint"' in line and \
+ '1.0' in line:
+ found_failed_fallback = True
+ break
+
+ assert found_failed_fallback, "Expected litellm_deployment_failed_fallbacks_total metric not found in /metrics"
async def create_test_team(
diff --git a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py
index b334966b441..2ebf9174e29 100644
--- a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py
+++ b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py
@@ -130,6 +130,67 @@ class BaseAnthropicMessagesTest:
return collected_chunks
+ @pytest.mark.asyncio
+ async def test_response_format_consistency(self):
+ """
+ Test that response content blocks are consistently dicts (not Pydantic objects).
+
+ This ensures that code like response["content"][0]["type"] works
+ regardless of the target provider.
+
+ Issue: https://github.com/BerriAI/litellm/issues/20342
+ """
+ litellm._turn_on_debug()
+
+ request_params = self.model_config
+
+ # Set up test parameters
+ messages = [{"role": "user", "content": "Say hi"}]
+
+ # Prepare call arguments
+ call_args = {
+ "messages": messages,
+ "max_tokens": 100,
+ }
+
+ # Add any additional config from subclass
+ call_args.update(request_params)
+
+ # Call the handler
+ response = await litellm.anthropic.messages.acreate(**call_args)
+
+ print(f"Response for {request_params['model']}: {json.dumps(response, indent=2, default=str)}")
+
+ # Verify response structure
+ assert "content" in response, "Response should have 'content' field"
+ assert len(response["content"]) > 0, "Response content should not be empty"
+
+ # Get the first content block
+ block = response["content"][0]
+
+ # Check that the block is a dict, not a Pydantic object
+ assert isinstance(block, dict), (
+ f"Content block should be a dict, but got {type(block)}. "
+ f"This means response format is inconsistent across providers."
+ )
+
+ # Verify we can access fields using dict syntax (not object attributes)
+ try:
+ block_type = block["type"]
+ print(f"✓ Successfully accessed block['type']: {block_type}")
+ except TypeError as e:
+ pytest.fail(
+ f"Cannot access content block using dict syntax: {e}. "
+ f"Block type: {type(block)}"
+ )
+
+ # Verify the block has expected structure
+ assert "type" in block, "Content block should have 'type' field"
+ if block["type"] == "text":
+ assert "text" in block, "Text content block should have 'text' field"
+
+ print(f"✓ Response format consistency test passed for {request_params['model']}")
+
@pytest.mark.asyncio
async def test_anthropic_messages_litellm_router_streaming_with_logging(self):
"""
diff --git a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py
index 17e72f29152..0a581fb512d 100644
--- a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py
+++ b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py
@@ -876,4 +876,4 @@ def test_sync_openai_messages():
assert response is not None
assert isinstance(response, dict)
- assert response["content"][0].text is not None
+ assert response["content"][0]["text"] is not None
diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py
index 05c6e4984af..66dfc8d15d5 100644
--- a/tests/proxy_unit_tests/test_auth_checks.py
+++ b/tests/proxy_unit_tests/test_auth_checks.py
@@ -21,11 +21,13 @@ from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_UserTable,
LiteLLM_TeamTable,
+ Litellm_EntityType,
)
from litellm.proxy.utils import PrismaClient
from litellm.proxy.auth.auth_checks import (
can_team_access_model,
_virtual_key_soft_budget_check,
+ _team_soft_budget_check,
)
from litellm.proxy.utils import ProxyLogging
from litellm.proxy.utils import CallInfo
@@ -478,6 +480,84 @@ async def test_virtual_key_soft_budget_check(spend, soft_budget, expect_alert):
), f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}"
+@pytest.mark.parametrize(
+ "spend, soft_budget, expect_alert, metadata, expected_alert_emails",
+ [
+ (100, 50, False, None, None), # Over soft budget, no metadata - no alert_emails configured, so no alert
+ (50, 50, False, None, None), # At soft budget, no metadata - no alert_emails configured, so no alert
+ (25, 50, False, None, None), # Under soft budget
+ (100, None, False, None, None), # No soft budget set
+ (100, 50, True, {"soft_budget_alerting_emails": ["team1@example.com", "team2@example.com"]}, ["team1@example.com", "team2@example.com"]), # Over soft budget with list of emails
+ (100, 50, True, {"soft_budget_alerting_emails": "team1@example.com,team2@example.com"}, ["team1@example.com", "team2@example.com"]), # Over soft budget with comma-separated emails
+ (100, 50, True, {"soft_budget_alerting_emails": ["team1@example.com", "", " ", "team2@example.com"]}, ["team1@example.com", "team2@example.com"]), # Over soft budget with empty strings filtered
+ ],
+)
+@pytest.mark.asyncio
+async def test_team_soft_budget_check(spend, soft_budget, expect_alert, metadata, expected_alert_emails):
+ """
+ Test cases for _team_soft_budget_check:
+ 1. Spend over soft budget, no alert_emails configured - should NOT trigger alert (alerts only sent when alert_emails configured)
+ 2. Spend at soft budget, no alert_emails configured - should NOT trigger alert (alerts only sent when alert_emails configured)
+ 3. Spend under soft budget - should not trigger alert
+ 4. No soft budget set - should not trigger alert
+ 5. Team with alert emails in metadata (list) - should include alert_emails in CallInfo
+ 6. Team with alert emails in metadata (comma-separated string) - should parse and include alert_emails
+ 7. Team with alert emails containing empty strings - should filter them out
+ """
+ alert_triggered = False
+ captured_call_info = None
+
+ class MockProxyLogging:
+ async def budget_alerts(self, type, user_info):
+ nonlocal alert_triggered, captured_call_info
+ alert_triggered = True
+ captured_call_info = user_info
+ assert type == "soft_budget"
+ assert isinstance(user_info, CallInfo)
+
+ valid_token = UserAPIKeyAuth(
+ token="test-token",
+ user_id="test-user",
+ team_id="test-team",
+ team_alias="test-team-alias",
+ key_alias="test-key",
+ )
+
+ team_object = LiteLLM_TeamTable(
+ team_id="test-team",
+ spend=spend,
+ soft_budget=soft_budget,
+ max_budget=100.0,
+ metadata=metadata,
+ )
+
+ proxy_logging_obj = MockProxyLogging()
+
+ await _team_soft_budget_check(
+ team_object=team_object,
+ valid_token=valid_token,
+ proxy_logging_obj=proxy_logging_obj,
+ )
+
+ await asyncio.sleep(0.1) # Allow time for the alert task to complete
+
+ assert (
+ alert_triggered == expect_alert
+ ), f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}"
+
+ if expect_alert:
+ assert captured_call_info is not None
+ assert captured_call_info.team_id == "test-team"
+ assert captured_call_info.spend == spend
+ assert captured_call_info.soft_budget == soft_budget
+ assert captured_call_info.event_group == Litellm_EntityType.TEAM
+ # Verify alert_emails if expected
+ if expected_alert_emails is not None:
+ assert captured_call_info.alert_emails == expected_alert_emails
+ else:
+ assert captured_call_info.alert_emails is None or captured_call_info.alert_emails == []
+
+
@pytest.mark.asyncio
async def test_can_user_call_model():
from litellm.proxy.auth.auth_checks import can_user_call_model
diff --git a/tests/proxy_unit_tests/test_zero_cost_model_budget_bypass.py b/tests/proxy_unit_tests/test_zero_cost_model_budget_bypass.py
new file mode 100644
index 00000000000..bc818fc0dca
--- /dev/null
+++ b/tests/proxy_unit_tests/test_zero_cost_model_budget_bypass.py
@@ -0,0 +1,590 @@
+"""
+Tests for zero-cost model budget bypass functionality.
+
+When a user exceeds their budget, the system should still allow requests
+to models with zero cost (e.g., on-premises models).
+"""
+
+import asyncio
+from typing import Optional
+from unittest.mock import MagicMock, patch
+
+import pytest
+
+import litellm
+from litellm.caching.caching import DualCache
+from litellm.proxy._types import (
+ LiteLLM_BudgetTable,
+ LiteLLM_EndUserTable,
+ LiteLLM_TeamMembership,
+ LiteLLM_TeamTable,
+ LiteLLM_UserTable,
+ UserAPIKeyAuth,
+)
+from litellm.proxy.auth.auth_checks import (
+ _check_team_member_budget,
+ _is_model_cost_zero,
+ _team_max_budget_check,
+ common_checks,
+)
+from litellm.proxy.utils import ProxyLogging
+from litellm.router import Router
+from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
+
+
+@pytest.fixture
+def mock_router_with_zero_cost_model():
+ """Create a mock router with a zero-cost model."""
+ router = Router(
+ model_list=[
+ {
+ "model_name": "on-prem-model",
+ "litellm_params": {
+ "model": "ollama/llama2",
+ "api_base": "http://localhost:11434",
+ "input_cost_per_token": 0.0,
+ "output_cost_per_token": 0.0,
+ },
+ "model_info": {
+ "id": "on-prem-model-id",
+ "input_cost_per_token": 0.0,
+ "output_cost_per_token": 0.0,
+ },
+ },
+ {
+ "model_name": "cloud-model",
+ "litellm_params": {
+ "model": "gpt-3.5-turbo",
+ "api_key": "sk-test",
+ },
+ "model_info": {
+ "id": "cloud-model-id",
+ },
+ },
+ ]
+ )
+ return router
+
+
+@pytest.fixture
+def mock_router_with_paid_model():
+ """Create a mock router with only paid models."""
+ router = Router(
+ model_list=[
+ {
+ "model_name": "cloud-model",
+ "litellm_params": {
+ "model": "gpt-3.5-turbo",
+ "api_key": "sk-test",
+ },
+ "model_info": {
+ "id": "cloud-model-id",
+ },
+ }
+ ]
+ )
+ return router
+
+
+@pytest.fixture
+def mock_proxy_logging():
+ """Create a mock ProxyLogging instance."""
+ proxy_logging = ProxyLogging(user_api_key_cache=None)
+
+ async def mock_budget_alerts(*args, **kwargs):
+ pass
+
+ proxy_logging.budget_alerts = mock_budget_alerts
+ return proxy_logging
+
+
+class TestIsModelCostZero:
+ """Tests for _is_model_cost_zero helper function."""
+
+ def test_zero_cost_model_in_router(self, mock_router_with_zero_cost_model):
+ """Test that a zero-cost model in router is correctly identified."""
+ result = _is_model_cost_zero(
+ model="on-prem-model", llm_router=mock_router_with_zero_cost_model
+ )
+ assert result is True
+
+ def test_paid_model_in_router(self, mock_router_with_zero_cost_model):
+ """Test that a paid model is correctly identified as non-zero cost."""
+ with patch("litellm.get_model_info") as mock_get_model_info:
+ # Mock the return value for gpt-3.5-turbo
+ mock_get_model_info.return_value = {
+ "input_cost_per_token": 0.0000015,
+ "output_cost_per_token": 0.000002,
+ }
+ result = _is_model_cost_zero(
+ model="cloud-model", llm_router=mock_router_with_zero_cost_model
+ )
+ assert result is False
+
+ def test_none_model(self, mock_router_with_zero_cost_model):
+ """Test that None model returns False."""
+ result = _is_model_cost_zero(
+ model=None, llm_router=mock_router_with_zero_cost_model
+ )
+ assert result is False
+
+ def test_none_router(self):
+ """Test that None router returns False."""
+ result = _is_model_cost_zero(model="some-model", llm_router=None)
+ assert result is False
+
+ def test_list_of_zero_cost_models(self, mock_router_with_zero_cost_model):
+ """Test that a list of zero-cost models returns True."""
+ result = _is_model_cost_zero(
+ model=["on-prem-model"], llm_router=mock_router_with_zero_cost_model
+ )
+ assert result is True
+
+ def test_mixed_cost_models(self, mock_router_with_zero_cost_model):
+ """Test that a list with mixed cost models returns False."""
+ with patch("litellm.get_model_info") as mock_get_model_info:
+ mock_get_model_info.return_value = {
+ "input_cost_per_token": 0.0000015,
+ "output_cost_per_token": 0.000002,
+ }
+ result = _is_model_cost_zero(
+ model=["on-prem-model", "cloud-model"],
+ llm_router=mock_router_with_zero_cost_model,
+ )
+ assert result is False
+
+
+class TestUserBudgetBypass:
+ """Tests for user budget bypass with zero-cost models."""
+
+ @pytest.mark.asyncio
+ async def test_user_over_budget_with_zero_cost_model_allowed(
+ self, mock_router_with_zero_cost_model, mock_proxy_logging
+ ):
+ """Test that user over budget can still use zero-cost models."""
+ user_object = LiteLLM_UserTable(
+ user_id="test-user",
+ spend=100.0,
+ max_budget=50.0,
+ )
+
+ request_body = {"model": "on-prem-model"}
+
+ # Should not raise BudgetExceededError
+ result = await common_checks(
+ request_body=request_body,
+ team_object=None,
+ user_object=user_object,
+ end_user_object=None,
+ global_proxy_spend=None,
+ general_settings={},
+ route="/v1/chat/completions",
+ llm_router=mock_router_with_zero_cost_model,
+ proxy_logging_obj=mock_proxy_logging,
+ valid_token=UserAPIKeyAuth(
+ token="test-token",
+ user_id="test-user",
+ ),
+ request=MagicMock(),
+ skip_budget_checks=True, # This is set by user_api_key_auth for zero-cost models
+ )
+ assert result is True
+
+ @pytest.mark.asyncio
+ async def test_user_over_budget_with_paid_model_blocked(
+ self, mock_router_with_zero_cost_model, mock_proxy_logging
+ ):
+ """Test that user over budget cannot use paid models."""
+ user_object = LiteLLM_UserTable(
+ user_id="test-user",
+ spend=100.0,
+ max_budget=50.0,
+ )
+
+ request_body = {"model": "cloud-model"}
+
+ with patch("litellm.get_model_info") as mock_get_model_info:
+ mock_get_model_info.return_value = {
+ "input_cost_per_token": 0.0000015,
+ "output_cost_per_token": 0.000002,
+ }
+ with pytest.raises(litellm.BudgetExceededError) as exc_info:
+ await common_checks(
+ request_body=request_body,
+ team_object=None,
+ user_object=user_object,
+ end_user_object=None,
+ global_proxy_spend=None,
+ general_settings={},
+ route="/v1/chat/completions",
+ llm_router=mock_router_with_zero_cost_model,
+ proxy_logging_obj=mock_proxy_logging,
+ valid_token=UserAPIKeyAuth(
+ token="test-token",
+ user_id="test-user",
+ ),
+ request=MagicMock(),
+ )
+
+ assert exc_info.value.current_cost == 100.0
+ assert exc_info.value.max_budget == 50.0
+ assert "test-user" in str(exc_info.value)
+
+
+class TestEndUserBudgetBypass:
+ """Tests for end user budget bypass with zero-cost models."""
+
+ @pytest.mark.asyncio
+ async def test_end_user_over_budget_with_zero_cost_model_allowed(
+ self, mock_router_with_zero_cost_model, mock_proxy_logging
+ ):
+ """Test that end user over budget can still use zero-cost models."""
+ end_user_budget = LiteLLM_BudgetTable(max_budget=20.0)
+ end_user_object = LiteLLM_EndUserTable(
+ user_id="end-user-123",
+ spend=50.0,
+ litellm_budget_table=end_user_budget,
+ blocked=False,
+ )
+
+ request_body = {"model": "on-prem-model", "user": "end-user-123"}
+
+ # In the real flow, skip_budget_checks would be set to True for zero-cost models
+ result = await common_checks(
+ request_body=request_body,
+ team_object=None,
+ user_object=None,
+ end_user_object=end_user_object,
+ global_proxy_spend=None,
+ general_settings={},
+ route="/v1/chat/completions",
+ llm_router=mock_router_with_zero_cost_model,
+ proxy_logging_obj=mock_proxy_logging,
+ valid_token=UserAPIKeyAuth(
+ token="test-token",
+ ),
+ request=MagicMock(),
+ skip_budget_checks=True, # This is set by user_api_key_auth for zero-cost models
+ )
+ assert result is True
+
+ @pytest.mark.asyncio
+ async def test_end_user_over_budget_with_paid_model_blocked(
+ self, mock_router_with_zero_cost_model, mock_proxy_logging
+ ):
+ """Test that end user over budget cannot use paid models."""
+ end_user_budget = LiteLLM_BudgetTable(max_budget=20.0)
+ end_user_object = LiteLLM_EndUserTable(
+ user_id="end-user-123",
+ spend=50.0,
+ litellm_budget_table=end_user_budget,
+ blocked=False,
+ )
+
+ request_body = {"model": "cloud-model", "user": "end-user-123"}
+
+ with patch("litellm.get_model_info") as mock_get_model_info:
+ mock_get_model_info.return_value = {
+ "input_cost_per_token": 0.0000015,
+ "output_cost_per_token": 0.000002,
+ }
+ with pytest.raises(litellm.BudgetExceededError) as exc_info:
+ await common_checks(
+ request_body=request_body,
+ team_object=None,
+ user_object=None,
+ end_user_object=end_user_object,
+ global_proxy_spend=None,
+ general_settings={},
+ route="/v1/chat/completions",
+ llm_router=mock_router_with_zero_cost_model,
+ proxy_logging_obj=mock_proxy_logging,
+ valid_token=UserAPIKeyAuth(
+ token="test-token",
+ ),
+ request=MagicMock(),
+ )
+
+ assert exc_info.value.current_cost == 50.0
+ assert exc_info.value.max_budget == 20.0
+ assert "end-user-123" in str(exc_info.value)
+
+
+class TestTeamBudgetBypass:
+ """Tests for team budget bypass with zero-cost models."""
+
+ @pytest.mark.asyncio
+ async def test_team_over_budget_with_zero_cost_model_allowed(
+ self, mock_router_with_zero_cost_model, mock_proxy_logging
+ ):
+ """Test that team over budget can still use zero-cost models."""
+ team_object = LiteLLM_TeamTable(
+ team_id="test-team",
+ spend=150.0,
+ max_budget=100.0,
+ )
+
+ valid_token = UserAPIKeyAuth(
+ token="test-token",
+ team_id="test-team",
+ )
+
+ request_body = {"model": "on-prem-model"}
+
+ # In the real flow, skip_budget_checks would be set to True for zero-cost models
+ result = await common_checks(
+ request_body=request_body,
+ team_object=team_object,
+ user_object=None,
+ end_user_object=None,
+ global_proxy_spend=None,
+ general_settings={},
+ route="/v1/chat/completions",
+ llm_router=mock_router_with_zero_cost_model,
+ proxy_logging_obj=mock_proxy_logging,
+ valid_token=valid_token,
+ request=MagicMock(),
+ skip_budget_checks=True, # This is set by user_api_key_auth for zero-cost models
+ )
+ assert result is True
+
+ @pytest.mark.asyncio
+ async def test_team_over_budget_with_paid_model_blocked(
+ self, mock_router_with_zero_cost_model, mock_proxy_logging
+ ):
+ """Test that team over budget cannot use paid models."""
+ team_object = LiteLLM_TeamTable(
+ team_id="test-team",
+ spend=150.0,
+ max_budget=100.0,
+ )
+
+ valid_token = UserAPIKeyAuth(
+ token="test-token",
+ team_id="test-team",
+ )
+
+ request_body = {"model": "cloud-model"}
+
+ with patch("litellm.get_model_info") as mock_get_model_info:
+ mock_get_model_info.return_value = {
+ "input_cost_per_token": 0.0000015,
+ "output_cost_per_token": 0.000002,
+ }
+ with pytest.raises(litellm.BudgetExceededError) as exc_info:
+ await common_checks(
+ request_body=request_body,
+ team_object=team_object,
+ user_object=None,
+ end_user_object=None,
+ global_proxy_spend=None,
+ general_settings={},
+ route="/v1/chat/completions",
+ llm_router=mock_router_with_zero_cost_model,
+ proxy_logging_obj=mock_proxy_logging,
+ valid_token=valid_token,
+ request=MagicMock(),
+ )
+
+ assert exc_info.value.current_cost == 150.0
+ assert exc_info.value.max_budget == 100.0
+ assert "test-team" in str(exc_info.value)
+
+
+class TestTeamMemberBudgetBypass:
+ """Tests for team member budget bypass with zero-cost models."""
+
+ @pytest.mark.asyncio
+ async def test_team_member_over_budget_with_zero_cost_model_allowed(
+ self, mock_router_with_zero_cost_model, mock_proxy_logging
+ ):
+ """Test that team member over budget can still use zero-cost models."""
+ team_object = LiteLLM_TeamTable(
+ team_id="test-team",
+ )
+
+ user_object = LiteLLM_UserTable(
+ user_id="test-user",
+ )
+
+ valid_token = UserAPIKeyAuth(
+ token="test-token",
+ user_id="test-user",
+ team_id="test-team",
+ )
+
+ member_budget = LiteLLM_BudgetTable(max_budget=30.0)
+ team_membership = LiteLLM_TeamMembership(
+ user_id="test-user",
+ team_id="test-team",
+ spend=60.0,
+ litellm_budget_table=member_budget,
+ )
+
+ request_body = {"model": "on-prem-model"}
+
+ # Mock get_team_membership
+ with patch(
+ "litellm.proxy.auth.auth_checks.get_team_membership"
+ ) as mock_get_membership:
+ mock_get_membership.return_value = team_membership
+
+ # In the real flow, skip_budget_checks would be set to True for zero-cost models
+ result = await common_checks(
+ request_body=request_body,
+ team_object=team_object,
+ user_object=user_object,
+ end_user_object=None,
+ global_proxy_spend=None,
+ general_settings={},
+ route="/v1/chat/completions",
+ llm_router=mock_router_with_zero_cost_model,
+ proxy_logging_obj=mock_proxy_logging,
+ valid_token=valid_token,
+ request=MagicMock(),
+ skip_budget_checks=True, # This is set by user_api_key_auth for zero-cost models
+ )
+ assert result is True
+
+ @pytest.mark.asyncio
+ async def test_team_member_over_budget_with_paid_model_blocked(
+ self, mock_router_with_zero_cost_model, mock_proxy_logging
+ ):
+ """Test that team member over budget cannot use paid models."""
+ team_object = LiteLLM_TeamTable(
+ team_id="test-team",
+ )
+
+ user_object = LiteLLM_UserTable(
+ user_id="test-user",
+ )
+
+ valid_token = UserAPIKeyAuth(
+ token="test-token",
+ user_id="test-user",
+ team_id="test-team",
+ )
+
+ member_budget = LiteLLM_BudgetTable(max_budget=30.0)
+ team_membership = LiteLLM_TeamMembership(
+ user_id="test-user",
+ team_id="test-team",
+ spend=60.0,
+ litellm_budget_table=member_budget,
+ )
+
+ request_body = {"model": "cloud-model"}
+
+ with patch(
+ "litellm.proxy.auth.auth_checks.get_team_membership"
+ ) as mock_get_membership:
+ mock_get_membership.return_value = team_membership
+
+ with patch("litellm.get_model_info") as mock_get_model_info:
+ mock_get_model_info.return_value = {
+ "input_cost_per_token": 0.0000015,
+ "output_cost_per_token": 0.000002,
+ }
+ with pytest.raises(litellm.BudgetExceededError) as exc_info:
+ await common_checks(
+ request_body=request_body,
+ team_object=team_object,
+ user_object=user_object,
+ end_user_object=None,
+ global_proxy_spend=None,
+ general_settings={},
+ route="/v1/chat/completions",
+ llm_router=mock_router_with_zero_cost_model,
+ proxy_logging_obj=mock_proxy_logging,
+ valid_token=valid_token,
+ request=MagicMock(),
+ )
+
+ assert exc_info.value.current_cost == 60.0
+ assert exc_info.value.max_budget == 30.0
+ assert "test-user" in str(exc_info.value)
+ assert "test-team" in str(exc_info.value)
+
+
+class TestEdgeCases:
+ """Tests for edge cases and error handling."""
+
+ def test_model_not_in_router(self, mock_router_with_zero_cost_model):
+ """Test behavior when model is not found in router."""
+ with patch("litellm.get_model_info") as mock_get_model_info:
+ # Simulate model not found
+ mock_get_model_info.side_effect = Exception("Model not found")
+ result = _is_model_cost_zero(
+ model="nonexistent-model", llm_router=mock_router_with_zero_cost_model
+ )
+ # Should return False (conservative approach)
+ assert result is False
+
+ @pytest.mark.asyncio
+ async def test_user_under_budget_with_paid_model_allowed(
+ self, mock_router_with_zero_cost_model, mock_proxy_logging
+ ):
+ """Test that user under budget can use paid models normally."""
+ user_object = LiteLLM_UserTable(
+ user_id="test-user",
+ spend=30.0,
+ max_budget=100.0,
+ )
+
+ request_body = {"model": "cloud-model"}
+
+ with patch("litellm.get_model_info") as mock_get_model_info:
+ mock_get_model_info.return_value = {
+ "input_cost_per_token": 0.0000015,
+ "output_cost_per_token": 0.000002,
+ }
+ # Should not raise BudgetExceededError
+ result = await common_checks(
+ request_body=request_body,
+ team_object=None,
+ user_object=user_object,
+ end_user_object=None,
+ global_proxy_spend=None,
+ general_settings={},
+ route="/v1/chat/completions",
+ llm_router=mock_router_with_zero_cost_model,
+ proxy_logging_obj=mock_proxy_logging,
+ valid_token=UserAPIKeyAuth(
+ token="test-token",
+ user_id="test-user",
+ ),
+ request=MagicMock(),
+ )
+ assert result is True
+
+ @pytest.mark.asyncio
+ async def test_user_under_budget_with_zero_cost_model_allowed(
+ self, mock_router_with_zero_cost_model, mock_proxy_logging
+ ):
+ """Test that user under budget can use zero-cost models normally."""
+ user_object = LiteLLM_UserTable(
+ user_id="test-user",
+ spend=30.0,
+ max_budget=100.0,
+ )
+
+ request_body = {"model": "on-prem-model"}
+
+ # Should not raise BudgetExceededError
+ result = await common_checks(
+ request_body=request_body,
+ team_object=None,
+ user_object=user_object,
+ end_user_object=None,
+ global_proxy_spend=None,
+ general_settings={},
+ route="/v1/chat/completions",
+ llm_router=mock_router_with_zero_cost_model,
+ proxy_logging_obj=mock_proxy_logging,
+ valid_token=UserAPIKeyAuth(
+ token="test-token",
+ user_id="test-user",
+ ),
+ request=MagicMock(),
+ )
+ assert result is True
diff --git a/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py
index adbaf219079..57352eafaf1 100644
--- a/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py
+++ b/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py
@@ -134,3 +134,25 @@ def test_transform_request_with_response_format():
assert result["text"]["format"]["type"] == "json_schema"
assert result["text"]["format"]["name"] == "person_schema"
assert "schema" in result["text"]["format"]
+
+
+def test_transform_request_includes_extra_headers():
+ """Test that transform_request forwards headers as extra_headers for upstream call."""
+ handler = LiteLLMResponsesTransformationHandler()
+ messages = [{"role": "user", "content": "Hello"}]
+ optional_params = {}
+ litellm_params = {}
+
+ class MockLoggingObj:
+ pass
+
+ headers = {"cf-aig-authorization": "secret-token"}
+ result = handler.transform_request(
+ model="gpt-5-pro",
+ messages=messages,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ headers=headers,
+ litellm_logging_obj=MockLoggingObj(),
+ )
+ assert result.get("extra_headers") == headers
diff --git a/tests/test_litellm/containers/test_container_api.py b/tests/test_litellm/containers/test_container_api.py
index ddfe7c9ef14..9dcd9312ef3 100644
--- a/tests/test_litellm/containers/test_container_api.py
+++ b/tests/test_litellm/containers/test_container_api.py
@@ -22,6 +22,7 @@ from litellm.containers.main import (
list_containers,
retrieve_container,
)
+from litellm.main import base_llm_http_handler
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
from litellm.router import Router
@@ -63,9 +64,7 @@ class TestContainerAPI:
name="Test Container"
)
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_handler.container_create_handler.return_value = mock_response
-
+ with patch.object(base_llm_http_handler, 'container_create_handler', return_value=mock_response):
response = create_container(
name="Test Container",
custom_llm_provider="openai"
@@ -89,9 +88,7 @@ class TestContainerAPI:
name="Expiring Container"
)
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_handler.container_create_handler.return_value = mock_response
-
+ with patch.object(base_llm_http_handler, 'container_create_handler', return_value=mock_response):
response = create_container(
name="Expiring Container",
expires_after={"anchor": "last_active_at", "minutes": 30},
@@ -113,9 +110,7 @@ class TestContainerAPI:
name="Container with Files"
)
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_handler.container_create_handler.return_value = mock_response
-
+ with patch.object(base_llm_http_handler, 'container_create_handler', return_value=mock_response):
response = create_container(
name="Container with Files",
file_ids=["file_123", "file_456"],
@@ -137,9 +132,7 @@ class TestContainerAPI:
name="Async Test Container"
)
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_handler.container_create_handler.return_value = mock_response
-
+ with patch.object(base_llm_http_handler, 'container_create_handler', return_value=mock_response):
response = await acreate_container(
name="Async Test Container",
custom_llm_provider="openai"
@@ -171,9 +164,7 @@ class TestContainerAPI:
has_more=False
)
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_handler.container_list_handler.return_value = mock_response
-
+ with patch.object(base_llm_http_handler, 'container_list_handler', return_value=mock_response):
response = await alist_containers(
custom_llm_provider="openai"
)
@@ -208,9 +199,7 @@ class TestContainerAPI:
name=container_name
)
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_handler.container_retrieve_handler.return_value = mock_response
-
+ with patch.object(base_llm_http_handler, 'container_retrieve_handler', return_value=mock_response) as mock_method:
# Act: Call retrieve_container
response = retrieve_container(
container_id=container_id,
@@ -218,8 +207,8 @@ class TestContainerAPI:
)
# Assert: Verify the handler was called correctly
- mock_handler.container_retrieve_handler.assert_called_once()
- call_kwargs = mock_handler.container_retrieve_handler.call_args.kwargs
+ mock_method.assert_called_once()
+ call_kwargs = mock_method.call_args.kwargs
assert call_kwargs["container_id"] == container_id
# Assert: Verify response structure and content
@@ -245,9 +234,7 @@ class TestContainerAPI:
name="Async Retrieved Container"
)
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_handler.container_retrieve_handler.return_value = mock_response
-
+ with patch.object(base_llm_http_handler, 'container_retrieve_handler', return_value=mock_response):
response = await aretrieve_container(
container_id=container_id,
custom_llm_provider="openai"
@@ -265,9 +252,7 @@ class TestContainerAPI:
deleted=True
)
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_handler.container_delete_handler.return_value = mock_response
-
+ with patch.object(base_llm_http_handler, 'container_delete_handler', return_value=mock_response):
response = delete_container(
container_id=container_id,
custom_llm_provider="openai"
@@ -288,9 +273,7 @@ class TestContainerAPI:
deleted=True
)
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_handler.container_delete_handler.return_value = mock_response
-
+ with patch.object(base_llm_http_handler, 'container_delete_handler', return_value=mock_response):
response = await adelete_container(
container_id=container_id,
custom_llm_provider="openai"
@@ -302,9 +285,7 @@ class TestContainerAPI:
def test_create_container_error_handling(self):
"""Test error handling in container creation."""
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_handler.container_create_handler.side_effect = Exception("API Error")
-
+ with patch.object(base_llm_http_handler, 'container_create_handler', side_effect=Exception("API Error")):
with pytest.raises(Exception):
create_container(
name="Error Test Container",
@@ -313,21 +294,20 @@ class TestContainerAPI:
def test_container_provider_config_retrieval(self):
"""Test that provider config is retrieved correctly."""
+ mock_response = ContainerObject(
+ id="cntr_config_test",
+ object="container",
+ created_at=1747857508,
+ status="running",
+ expires_after={"anchor": "last_active_at", "minutes": 20},
+ last_active_at=1747857508,
+ name="Config Test"
+ )
+
with patch('litellm.containers.main.ProviderConfigManager') as mock_config_manager:
mock_config_manager.get_provider_container_config.return_value = OpenAIContainerConfig()
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_response = ContainerObject(
- id="cntr_config_test",
- object="container",
- created_at=1747857508,
- status="running",
- expires_after={"anchor": "last_active_at", "minutes": 20},
- last_active_at=1747857508,
- name="Config Test"
- )
- mock_handler.container_create_handler.return_value = mock_response
-
+ with patch.object(base_llm_http_handler, 'container_create_handler', return_value=mock_response):
response = create_container(
name="Config Test",
custom_llm_provider="openai"
@@ -355,9 +335,7 @@ class TestContainerAPI:
name="Test Container"
)
- with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
- mock_handler.container_create_handler.return_value = mock_response
-
+ with patch.object(base_llm_http_handler, 'container_create_handler', return_value=mock_response):
result = await router.acreate_container(
name="Test Container",
custom_llm_provider="openai"
diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py
index 1065a8ed514..b07216921eb 100644
--- a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py
+++ b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py
@@ -2,9 +2,7 @@ import os
import sys
import unittest.mock as mock
-import httpx
import pytest
-import respx
from httpx import Response
sys.path.insert(0, os.path.abspath("../../.."))
@@ -38,32 +36,8 @@ def mock_env_vars():
yield
-@pytest.fixture
-def mock_httpx_client():
- with mock.patch(
- "litellm_enterprise.enterprise_callbacks.send_emails.resend_email.get_async_httpx_client"
- ) as mock_client:
-
- mock_response = mock.Mock(spec=Response)
- mock_response.status_code = 200
- mock_response.json.return_value = {"id": "test_email_id"}
- mock_response.raise_for_status.return_value = None
-
- mock_async_client = mock.AsyncMock()
- mock_async_client.post.return_value = mock_response
-
- mock_client.return_value = mock_async_client
- yield mock_async_client
-
-
@pytest.mark.asyncio
-@respx.mock
-async def test_send_email_success(mock_env_vars, mock_httpx_client):
- # Block all HTTP requests at network level to prevent real API calls
- respx.post("https://api.resend.com/emails").mock(
- return_value=httpx.Response(200, json={"id": "test_email_id"})
- )
-
+async def test_send_email_success(mock_env_vars):
# Initialize the logger
logger = ResendEmailLogger()
@@ -73,14 +47,27 @@ async def test_send_email_success(mock_env_vars, mock_httpx_client):
subject = "Test Subject"
html_body = "
Test email body
"
+ # Create mock HTTP client and inject it directly into the logger
+ # This ensures the mock is used regardless of any caching/import issues
+ mock_response = mock.Mock(spec=Response)
+ mock_response.status_code = 200
+ mock_response.json.return_value = {"id": "test_email_id"}
+ mock_response.raise_for_status.return_value = None
+
+ mock_async_client = mock.AsyncMock()
+ mock_async_client.post.return_value = mock_response
+
+ # Directly inject the mock client to bypass any caching
+ logger.async_httpx_client = mock_async_client
+
# Send email
await logger.send_email(
from_email=from_email, to_email=to_email, subject=subject, html_body=html_body
)
# Verify the HTTP client was called correctly
- mock_httpx_client.post.assert_called_once()
- call_args = mock_httpx_client.post.call_args
+ mock_async_client.post.assert_called_once()
+ call_args = mock_async_client.post.call_args
# Verify the URL
assert call_args[1]["url"] == "https://api.resend.com/emails"
@@ -97,13 +84,7 @@ async def test_send_email_success(mock_env_vars, mock_httpx_client):
@pytest.mark.asyncio
-@respx.mock
-async def test_send_email_missing_api_key(mock_httpx_client):
- # Block all HTTP requests at network level to prevent real API calls
- respx.post("https://api.resend.com/emails").mock(
- return_value=httpx.Response(200, json={"id": "test_email_id"})
- )
-
+async def test_send_email_missing_api_key():
# Remove the API key from environment before initializing logger
original_key = os.environ.pop("RESEND_API_KEY", None)
@@ -117,13 +98,18 @@ async def test_send_email_missing_api_key(mock_httpx_client):
subject = "Test Subject"
html_body = "
Test email body
"
- # Mock the response to avoid making real HTTP requests
+ # Create mock HTTP client and inject it directly into the logger
+ # This ensures the mock is used regardless of any caching issues
mock_response = mock.Mock(spec=Response)
mock_response.raise_for_status.return_value = None
-
mock_response.status_code = 200
mock_response.json.return_value = {"id": "test_email_id"}
- mock_httpx_client.post.return_value = mock_response
+
+ mock_async_client = mock.AsyncMock()
+ mock_async_client.post.return_value = mock_response
+
+ # Directly inject the mock client to bypass any caching
+ logger.async_httpx_client = mock_async_client
# Send email
await logger.send_email(
@@ -131,8 +117,8 @@ async def test_send_email_missing_api_key(mock_httpx_client):
)
# Verify the HTTP client was called with None as the API key
- mock_httpx_client.post.assert_called_once()
- call_args = mock_httpx_client.post.call_args
+ mock_async_client.post.assert_called_once()
+ call_args = mock_async_client.post.call_args
assert call_args[1]["headers"] == {"Authorization": "Bearer None"}
finally:
# Restore the original key if it existed
@@ -141,13 +127,7 @@ async def test_send_email_missing_api_key(mock_httpx_client):
@pytest.mark.asyncio
-@respx.mock
-async def test_send_email_multiple_recipients(mock_env_vars, mock_httpx_client):
- # Block all HTTP requests at network level to prevent real API calls
- respx.post("https://api.resend.com/emails").mock(
- return_value=httpx.Response(200, json={"id": "test_email_id"})
- )
-
+async def test_send_email_multiple_recipients(mock_env_vars):
# Initialize the logger
logger = ResendEmailLogger()
@@ -157,13 +137,17 @@ async def test_send_email_multiple_recipients(mock_env_vars, mock_httpx_client):
subject = "Test Subject"
html_body = "