Merge pull request #19753 from BerriAI/litellm_oss_staging_01_26_2026

fix(proxy): support slashes in google generateContent model names (#1…
This commit is contained in:
Sameer Kankute 2026-01-27 17:09:12 +05:30 • committed by GitHub
commit ea0a264a3c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
29 changed files with 2569 additions and 269 deletions

View file

@ -128,6 +128,7 @@ guardrails:
mode: ["pre_call", "post_call", "during_call"] # Run at multiple stages
api_key: os.environ/ONYX_API_KEY
api_base: os.environ/ONYX_API_BASE
timeout: 10.0 # Optional, defaults to 10 seconds
```
### Required Parameters
@ -137,6 +138,7 @@ guardrails:
### Optional Parameters
- **`api_base`**: Onyx API base URL (defaults to `https://ai-guard.onyx.security`)
- **`timeout`**: Request timeout in seconds (defaults to `10.0`)
## Environment Variables
@ -145,4 +147,5 @@ You can set these environment variables instead of hardcoding values in your con
```shell
export ONYX_API_KEY="your-api-key-here"
export ONYX_API_BASE="https://ai-guard.onyx.security" # Optional
export ONYX_TIMEOUT=10 # Optional, timeout in seconds
```

View file

@ -4,13 +4,17 @@ LiteLLM Proxy uses this MCP Client to connnect to other MCP servers.
import asyncio
import base64
from typing import Any, Awaitable, Callable, Dict, List, Optional, TypeVar, Union
from typing import Any, Awaitable, Callable, Dict, List, Optional, Tuple, TypeVar, Union
import httpx
from mcp import ClientSession, ReadResourceResult, Resource, StdioServerParameters
from mcp.client.sse import sse_client
from mcp.client.stdio import stdio_client
from mcp.client.streamable_http import streamable_http_client
try:
from mcp.client.streamable_http import streamable_http_client # type: ignore
except ImportError:
streamable_http_client = None
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import (
@ -76,104 +80,100 @@ class MCPClient:
def _create_transport_context(
self,
) -> tuple[Any, Optional[httpx.AsyncClient]]:
"""Create the appropriate transport context based on transport type."""
) -> Tuple[Any, Optional[httpx.AsyncClient]]:
"""
Create the appropriate transport context based on transport type.
Returns:
Tuple of (transport_context, http_client).
http_client is only set for HTTP transport and needs cleanup.
"""
http_client: Optional[httpx.AsyncClient] = None
if self.transport_type == MCPTransport.stdio:
if not self.stdio_config:
raise ValueError("stdio_config is required for stdio transport")
server_params = StdioServerParameters(
command=self.stdio_config.get("command", ""),
args=self.stdio_config.get("args", []),
env=self.stdio_config.get("env", {}),
)
transport_ctx = stdio_client(server_params)
elif self.transport_type == MCPTransport.sse:
return stdio_client(server_params), None
if self.transport_type == MCPTransport.sse:
headers = self._get_auth_headers()
httpx_client_factory = self._create_httpx_client_factory()
transport_ctx = sse_client(
return sse_client(
url=self.server_url,
timeout=self.timeout,
headers=headers,
httpx_client_factory=httpx_client_factory,
)
else:
headers = self._get_auth_headers()
httpx_client_factory = self._create_httpx_client_factory()
verbose_logger.debug(
"litellm headers for streamable_http_client: %s", headers
)
http_client = httpx_client_factory(
headers=headers,
timeout=httpx.Timeout(self.timeout),
)
transport_ctx = streamable_http_client(
url=self.server_url,
http_client=http_client,
)
if transport_ctx is None:
raise RuntimeError("Failed to create transport context")
), None
# HTTP transport (default)
headers = self._get_auth_headers()
httpx_client_factory = self._create_httpx_client_factory()
verbose_logger.debug(
"litellm headers for streamable_http_client: %s", headers
)
http_client = httpx_client_factory(
headers=headers,
timeout=httpx.Timeout(self.timeout),
)
transport_ctx = streamable_http_client(
url=self.server_url,
http_client=http_client,
)
return transport_ctx, http_client
async def _execute_session_operation(
self,
transport_ctx: Any,
operation: Callable[[ClientSession], Awaitable[TSessionResult]],
) -> TSessionResult:
"""
Execute an operation within a transport and session context.
Handles entering/exiting contexts and running the operation.
"""
transport = await transport_ctx.__aenter__()
try:
read_stream, write_stream = transport[0], transport[1]
session_ctx = ClientSession(read_stream, write_stream)
session = await session_ctx.__aenter__()
try:
await session.initialize()
return await operation(session)
finally:
try:
await session_ctx.__aexit__(None, None, None)
except BaseException as e:
verbose_logger.debug(f"Error during session context exit: {e}")
finally:
try:
await transport_ctx.__aexit__(None, None, None)
except BaseException as e:
verbose_logger.debug(f"Error during transport context exit: {e}")
async def run_with_session(
self, operation: Callable[[ClientSession], Awaitable[TSessionResult]]
) -> TSessionResult:
"""Open a session, run the provided coroutine, and clean up."""
transport_ctx = None
http_client: Optional[httpx.AsyncClient] = None
session_ctx = None
try:
transport_ctx, http_client = self._create_transport_context()
# Enter transport context
transport = await transport_ctx.__aenter__()
try:
read_stream, write_stream = transport[0], transport[1]
session_ctx = ClientSession(read_stream, write_stream)
# Enter session context
session = await session_ctx.__aenter__()
try:
await session.initialize()
result = await operation(session)
return result
finally:
# Ensure session context is properly exited
if session_ctx is not None:
try:
await session_ctx.__aexit__(None, None, None)
except Exception as e:
verbose_logger.debug(
f"Error during session context exit: {e}"
)
finally:
# Ensure transport context is properly exited
if transport_ctx is not None:
try:
await transport_ctx.__aexit__(None, None, None)
except Exception as e:
verbose_logger.debug(
f"Error during transport context exit: {e}"
)
return await self._execute_session_operation(transport_ctx, operation)
except Exception:
verbose_logger.warning(
"MCP client run_with_session failed for %s", self.server_url or "stdio"
)
raise
finally:
# Always clean up http_client if it was created
if http_client is not None:
try:
await http_client.aclose()
except Exception as e:
verbose_logger.debug(
f"Error during http_client cleanup: {e}"
)
except BaseException as e:
verbose_logger.debug(f"Error during http_client cleanup: {e}")
def update_auth_value(self, mcp_auth_value: Union[str, Dict[str, str]]):
"""

View file

@ -229,14 +229,18 @@ class PrometheusLogger(CustomLogger):
self.litellm_remaining_api_key_requests_for_model = self._gauge_factory(
"litellm_remaining_api_key_requests_for_model",
"Remaining Requests API Key can make for model (model based rpm limit on key)",
labelnames=["hashed_api_key", "api_key_alias", "model"],
labelnames=self.get_labels_for_metric(
"litellm_remaining_api_key_requests_for_model"
),
)
# Remaining MODEL TPM limit for API Key
self.litellm_remaining_api_key_tokens_for_model = self._gauge_factory(
"litellm_remaining_api_key_tokens_for_model",
"Remaining Tokens API Key can make for model (model based tpm limit on key)",
labelnames=["hashed_api_key", "api_key_alias", "model"],
labelnames=self.get_labels_for_metric(
"litellm_remaining_api_key_tokens_for_model"
),
)
########################################
@ -312,6 +316,18 @@ class PrometheusLogger(CustomLogger):
labelnames=self.get_labels_for_metric("litellm_deployment_state"),
)
self.litellm_deployment_tpm_limit = self._gauge_factory(
"litellm_deployment_tpm_limit",
"Deployment TPM limit found in config",
labelnames=self.get_labels_for_metric("litellm_deployment_tpm_limit"),
)
self.litellm_deployment_rpm_limit = self._gauge_factory(
"litellm_deployment_rpm_limit",
"Deployment RPM limit found in config",
labelnames=self.get_labels_for_metric("litellm_deployment_rpm_limit"),
)
self.litellm_deployment_cooled_down = self._counter_factory(
"litellm_deployment_cooled_down",
"LLM Deployment Analytics - Number of times a deployment has been cooled down by LiteLLM load balancing logic. exception_status is the status of the exception that caused the deployment to be cooled down",
@ -373,15 +389,9 @@ class PrometheusLogger(CustomLogger):
self.litellm_llm_api_failed_requests_metric = self._counter_factory(
name="litellm_llm_api_failed_requests_metric",
documentation="deprecated - use litellm_proxy_failed_requests_metric",
labelnames=[
"end_user",
"hashed_api_key",
"api_key_alias",
"model",
"team",
"team_alias",
"user",
],
labelnames=self.get_labels_for_metric(
"litellm_llm_api_failed_requests_metric"
),
)
self.litellm_requests_metric = self._counter_factory(
@ -954,6 +964,8 @@ class PrometheusLogger(CustomLogger):
route=standard_logging_payload["metadata"].get(
"user_api_key_request_route"
),
client_ip=standard_logging_payload["metadata"].get("requester_ip_address"),
user_agent=standard_logging_payload["metadata"].get("user_agent"),
)
if (
@ -1011,6 +1023,7 @@ class PrometheusLogger(CustomLogger):
user_api_key_alias=user_api_key_alias,
kwargs=kwargs,
metadata=_metadata,
model_id=enum_values.model_id,
)
# set latency metrics
@ -1245,6 +1258,7 @@ class PrometheusLogger(CustomLogger):
user_api_key_alias: Optional[str],
kwargs: dict,
metadata: dict,
model_id: Optional[str] = None,
):
from litellm.proxy.common_utils.callback_utils import (
get_model_group_from_litellm_kwargs,
@ -1266,11 +1280,11 @@ class PrometheusLogger(CustomLogger):
)
self.litellm_remaining_api_key_requests_for_model.labels(
user_api_key, user_api_key_alias, model_group
user_api_key, user_api_key_alias, model_group, model_id
).set(remaining_requests)
self.litellm_remaining_api_key_tokens_for_model.labels(
user_api_key, user_api_key_alias, model_group
user_api_key, user_api_key_alias, model_group, model_id
).set(remaining_tokens)
def _set_latency_metrics(
@ -1365,14 +1379,14 @@ class PrometheusLogger(CustomLogger):
standard_logging_payload: StandardLoggingPayload = kwargs.get(
"standard_logging_object", {}
)
if self._should_skip_metrics_for_invalid_key(
kwargs=kwargs, standard_logging_payload=standard_logging_payload
):
return
model = kwargs.get("model", "")
litellm_params = kwargs.get("litellm_params", {}) or {}
get_end_user_id_for_cost_tracking = _get_cached_end_user_id_for_cost_tracking()
@ -1396,6 +1410,7 @@ class PrometheusLogger(CustomLogger):
user_api_team,
user_api_team_alias,
user_id,
standard_logging_payload.get("model_id", ""),
).inc()
self.set_llm_deployment_failure_metrics(kwargs)
except Exception as e:
@ -1413,49 +1428,57 @@ class PrometheusLogger(CustomLogger):
) -> Optional[int]:
"""
Extract HTTP status code from various input formats for validation.
This is a centralized helper to extract status code from different
callback function signatures. Handles both ProxyException (uses 'code')
and standard exceptions (uses 'status_code').
Args:
kwargs: Dictionary potentially containing 'exception' key
enum_values: Object with 'status_code' attribute
exception: Exception object to extract status code from directly
Returns:
Status code as integer if found, None otherwise
"""
status_code = None
# Try from enum_values first (most common in our callbacks)
if enum_values and hasattr(enum_values, "status_code") and enum_values.status_code:
if (
enum_values
and hasattr(enum_values, "status_code")
and enum_values.status_code
):
try:
status_code = int(enum_values.status_code)
except (ValueError, TypeError):
pass
if not status_code and exception:
# ProxyException uses 'code' attribute, other exceptions may use 'status_code'
status_code = getattr(exception, "status_code", None) or getattr(exception, "code", None)
status_code = getattr(exception, "status_code", None) or getattr(
exception, "code", None
)
if status_code is not None:
try:
status_code = int(status_code)
except (ValueError, TypeError):
status_code = None
if not status_code and kwargs:
exception_in_kwargs = kwargs.get("exception")
if exception_in_kwargs:
status_code = getattr(exception_in_kwargs, "status_code", None) or getattr(exception_in_kwargs, "code", None)
status_code = getattr(
exception_in_kwargs, "status_code", None
) or getattr(exception_in_kwargs, "code", None)
if status_code is not None:
try:
status_code = int(status_code)
except (ValueError, TypeError):
status_code = None
return status_code
def _is_invalid_api_key_request(
self,
status_code: Optional[int],
@ -1463,23 +1486,23 @@ class PrometheusLogger(CustomLogger):
) -> bool:
"""
Determine if a request has an invalid API key based on status code and exception.
This method prevents invalid authentication attempts from being recorded in
Prometheus metrics. A 401 status code is the definitive indicator of authentication
failure. Additionally, we check exception messages for authentication error patterns
to catch cases where the exception hasn't been converted to a ProxyException yet.
Args:
status_code: HTTP status code (401 indicates authentication error)
exception: Exception object to check for auth-related error messages
Returns:
True if the request has an invalid API key and metrics should be skipped,
False otherwise
"""
if status_code == 401:
return True
# Handle cases where AssertionError is raised before conversion to ProxyException
if exception is not None:
exception_str = str(exception).lower()
@ -1492,9 +1515,9 @@ class PrometheusLogger(CustomLogger):
]
if any(pattern in exception_str for pattern in auth_error_patterns):
return True
return False
def _should_skip_metrics_for_invalid_key(
self,
kwargs: Optional[dict] = None,
@ -1505,18 +1528,18 @@ class PrometheusLogger(CustomLogger):
) -> bool:
"""
Determine if Prometheus metrics should be skipped for invalid API key requests.
This is a centralized validation method that extracts status code and exception
information from various callback function signatures and determines if the request
represents an invalid API key attempt that should be filtered from metrics.
Args:
kwargs: Dictionary potentially containing exception and other data
user_api_key_dict: User API key authentication object (currently unused)
enum_values: Object with status_code attribute
standard_logging_payload: Standard logging payload dictionary
exception: Exception object to check directly
Returns:
True if metrics should be skipped (invalid key detected), False otherwise
"""
@ -1525,17 +1548,17 @@ class PrometheusLogger(CustomLogger):
enum_values=enum_values,
exception=exception,
)
if exception is None and kwargs:
exception = kwargs.get("exception")
if self._is_invalid_api_key_request(status_code, exception=exception):
verbose_logger.debug(
"Skipping Prometheus metrics for invalid API key request: "
f"status_code={status_code}, exception={type(exception).__name__ if exception else None}"
)
return True
return False
async def async_post_call_failure_hook(
@ -1576,6 +1599,10 @@ class PrometheusLogger(CustomLogger):
litellm_params=request_data,
proxy_server_request=request_data.get("proxy_server_request", {}),
)
_metadata = request_data.get("metadata", {}) or {}
model_id = _metadata.get("model_info", {}).get("id") or request_data.get(
"model_info", {}
).get("id")
enum_values = UserAPIKeyLabelValues(
end_user=user_api_key_dict.end_user_id,
user=user_api_key_dict.user_id,
@ -1590,6 +1617,9 @@ class PrometheusLogger(CustomLogger):
exception_class=self._get_exception_class_name(original_exception),
tags=_tags,
route=user_api_key_dict.request_route,
client_ip=_metadata.get("requester_ip_address"),
user_agent=_metadata.get("user_agent"),
model_id=model_id,
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
@ -1629,6 +1659,7 @@ class PrometheusLogger(CustomLogger):
):
return
_metadata = data.get("metadata", {}) or {}
enum_values = UserAPIKeyLabelValues(
end_user=user_api_key_dict.end_user_id,
hashed_api_key=user_api_key_dict.api_key,
@ -1644,6 +1675,8 @@ class PrometheusLogger(CustomLogger):
litellm_params=data,
proxy_server_request=data.get("proxy_server_request", {}),
),
client_ip=_metadata.get("requester_ip_address"),
user_agent=_metadata.get("user_agent"),
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
@ -1684,7 +1717,7 @@ class PrometheusLogger(CustomLogger):
exception = request_kwargs.get("exception", None)
llm_provider = _litellm_params.get("custom_llm_provider", None)
if self._should_skip_metrics_for_invalid_key(
kwargs=request_kwargs,
standard_logging_payload=standard_logging_payload,
@ -1716,6 +1749,10 @@ class PrometheusLogger(CustomLogger):
"user_api_key_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"),
)
"""
@ -1753,6 +1790,49 @@ class PrometheusLogger(CustomLogger):
)
)
def _set_deployment_tpm_rpm_limit_metrics(
self,
model_info: dict,
litellm_params: dict,
litellm_model_name: Optional[str],
model_id: Optional[str],
api_base: Optional[str],
llm_provider: Optional[str],
):
"""
Set the deployment TPM and RPM limits metrics
"""
tpm = model_info.get("tpm") or litellm_params.get("tpm")
rpm = model_info.get("rpm") or litellm_params.get("rpm")
if tpm is not None:
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_deployment_tpm_limit"
),
enum_values=UserAPIKeyLabelValues(
litellm_model_name=litellm_model_name,
model_id=model_id,
api_base=api_base,
api_provider=llm_provider,
),
)
self.litellm_deployment_tpm_limit.labels(**_labels).set(tpm)
if rpm is not None:
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_deployment_rpm_limit"
),
enum_values=UserAPIKeyLabelValues(
litellm_model_name=litellm_model_name,
model_id=model_id,
api_base=api_base,
api_provider=llm_provider,
),
)
self.litellm_deployment_rpm_limit.labels(**_labels).set(rpm)
def set_llm_deployment_success_metrics(
self,
request_kwargs: dict,
@ -1786,6 +1866,16 @@ class PrometheusLogger(CustomLogger):
_model_info = _metadata.get("model_info") or {}
model_id = _model_info.get("id", None)
if _model_info or _litellm_params:
self._set_deployment_tpm_rpm_limit_metrics(
model_info=_model_info,
litellm_params=_litellm_params,
litellm_model_name=litellm_model_name,
model_id=model_id,
api_base=api_base,
llm_provider=llm_provider,
)
remaining_requests: Optional[int] = None
remaining_tokens: Optional[int] = None
if additional_headers := standard_logging_payload["hidden_params"][
@ -2263,7 +2353,10 @@ class PrometheusLogger(CustomLogger):
async def fetch_keys(
page_size: int, page: int
) -> Tuple[List[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]], Optional[int]]:
) -> Tuple[
List[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]],
Optional[int],
]:
key_list_response = await _list_key_helper(
prisma_client=prisma_client,
page=page,
@ -2379,12 +2472,16 @@ class PrometheusLogger(CustomLogger):
# Get total user count
total_users = await prisma_client.db.litellm_usertable.count()
self.litellm_total_users_metric.set(total_users)
verbose_logger.debug(f"Prometheus: set litellm_total_users to {total_users}")
verbose_logger.debug(
f"Prometheus: set litellm_total_users to {total_users}"
)
# Get total team count
total_teams = await prisma_client.db.litellm_teamtable.count()
self.litellm_teams_count_metric.set(total_teams)
verbose_logger.debug(f"Prometheus: set litellm_teams_count to {total_teams}")
verbose_logger.debug(
f"Prometheus: set litellm_teams_count to {total_teams}"
)
except Exception as e:
verbose_logger.exception(
f"Error initializing user/team count metrics: {str(e)}"

View file

@ -93,8 +93,11 @@ def get_litellm_params(
"text_completion": text_completion,
"azure_ad_token_provider": azure_ad_token_provider,
"user_continue_message": user_continue_message,
"base_model": base_model or (
_get_base_model_from_litellm_call_metadata(metadata=metadata) if metadata else None
"base_model": base_model
or (
_get_base_model_from_litellm_call_metadata(metadata=metadata)
if metadata
else None
),
"litellm_trace_id": litellm_trace_id,
"litellm_session_id": litellm_session_id,
@ -139,5 +142,7 @@ def get_litellm_params(
"aws_sts_endpoint": kwargs.get("aws_sts_endpoint"),
"aws_external_id": kwargs.get("aws_external_id"),
"aws_bedrock_runtime_endpoint": kwargs.get("aws_bedrock_runtime_endpoint"),
"tpm": kwargs.get("tpm"),
"rpm": kwargs.get("rpm"),
}
return litellm_params

View file

@ -335,7 +335,9 @@ class Logging(LiteLLMLoggingBaseClass):
self.start_time = start_time # log the call start time
self.call_type = call_type
self.litellm_call_id = litellm_call_id
self.litellm_trace_id: str = litellm_trace_id if litellm_trace_id else str(uuid.uuid4())
self.litellm_trace_id: str = (
litellm_trace_id if litellm_trace_id else str(uuid.uuid4())
)
self.function_id = function_id
self.streaming_chunks: List[Any] = [] # for generating complete stream response
self.sync_streaming_chunks: List[
@ -544,7 +546,10 @@ class Logging(LiteLLMLoggingBaseClass):
if "stream_options" in additional_params:
self.stream_options = additional_params["stream_options"]
## check if custom pricing set ##
if any(litellm_params.get(key) is not None for key in _CUSTOM_PRICING_KEYS & litellm_params.keys()):
if any(
litellm_params.get(key) is not None
for key in _CUSTOM_PRICING_KEYS & litellm_params.keys()
):
self.custom_pricing = True
if "custom_llm_provider" in self.model_call_details:
@ -4454,6 +4459,7 @@ class StandardLoggingPayloadSetup:
user_api_key_request_route=None,
spend_logs_metadata=None,
requester_ip_address=None,
user_agent=None,
requester_metadata=None,
prompt_management_metadata=prompt_management_metadata,
applied_guardrails=applied_guardrails,
@ -5139,6 +5145,7 @@ def get_standard_logging_object_payload(
model_group=_model_group,
model_id=_model_id,
requester_ip_address=clean_metadata.get("requester_ip_address", None),
user_agent=clean_metadata.get("user_agent", None),
messages=StandardLoggingPayloadSetup.append_system_prompt_messages(
kwargs=kwargs, messages=kwargs.get("messages")
),
@ -5204,6 +5211,7 @@ def get_standard_logging_metadata(
user_api_key_team_alias=None,
spend_logs_metadata=None,
requester_ip_address=None,
user_agent=None,
requester_metadata=None,
user_api_key_end_user_id=None,
prompt_management_metadata=None,

View file

@ -148,7 +148,7 @@ from litellm.utils import (
validate_and_fix_openai_messages,
validate_and_fix_openai_tools,
validate_chat_completion_tool_choice,
validate_openai_optional_params
validate_openai_optional_params,
)
from ._logging import verbose_logger
@ -368,7 +368,7 @@ class AsyncCompletions:
@tracer.wrap()
@client
async def acompletion( # noqa: PLR0915
async def acompletion( # noqa: PLR0915
model: str,
# Optional OpenAI params: see https://platform.openai.com/docs/api-reference/chat/create
messages: List = [],
@ -603,12 +603,11 @@ async def acompletion( # noqa: PLR0915
if timeout is not None and isinstance(timeout, (int, float)):
timeout_value = float(timeout)
init_response = await asyncio.wait_for(
loop.run_in_executor(None, func_with_context),
timeout=timeout_value
loop.run_in_executor(None, func_with_context), timeout=timeout_value
)
else:
init_response = await loop.run_in_executor(None, func_with_context)
if isinstance(init_response, dict) or isinstance(
init_response, ModelResponse
): ## CACHING SCENARIO
@ -640,6 +639,7 @@ async def acompletion( # noqa: PLR0915
except asyncio.TimeoutError:
custom_llm_provider = custom_llm_provider or "openai"
from litellm.exceptions import Timeout
raise Timeout(
message=f"Request timed out after {timeout} seconds",
model=model,
@ -1118,7 +1118,6 @@ def completion( # type: ignore # noqa: PLR0915
# validate optional params
stop = validate_openai_optional_params(stop=stop)
######### unpacking kwargs #####################
args = locals()
@ -1135,7 +1134,9 @@ def completion( # type: ignore # noqa: PLR0915
# Check if MCP tools are present (following responses pattern)
# Cast tools to Optional[Iterable[ToolParam]] for type checking
tools_for_mcp = cast(Optional[Iterable[ToolParam]], tools)
if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools_for_mcp):
if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(
tools=tools_for_mcp
):
# Return coroutine - acompletion will await it
# completion() can return a coroutine when MCP tools are present, which acompletion() awaits
return acompletion_with_mcp( # type: ignore[return-value]
@ -1536,6 +1537,8 @@ def completion( # type: ignore # noqa: PLR0915
max_retries=max_retries,
timeout=timeout,
litellm_request_debug=kwargs.get("litellm_request_debug", False),
tpm=kwargs.get("tpm"),
rpm=kwargs.get("rpm"),
)
cast(LiteLLMLoggingObj, logging).update_environment_variables(
model=model,
@ -2361,11 +2364,7 @@ def completion( # type: ignore # noqa: PLR0915
input=messages, api_key=api_key, original_response=response
)
elif custom_llm_provider == "minimax":
api_key = (
api_key
or get_secret_str("MINIMAX_API_KEY")
or litellm.api_key
)
api_key = api_key or get_secret_str("MINIMAX_API_KEY") or litellm.api_key
api_base = (
api_base
@ -2413,7 +2412,9 @@ def completion( # type: ignore # noqa: PLR0915
or custom_llm_provider == "wandb"
or custom_llm_provider == "clarifai"
or custom_llm_provider in litellm.openai_compatible_providers
or JSONProviderRegistry.exists(custom_llm_provider) # JSON-configured providers
or JSONProviderRegistry.exists(
custom_llm_provider
) # JSON-configured providers
or "ft:gpt-3.5-turbo" in model # finetune gpt-3.5-turbo
): # allow user to make an openai call with a custom base
# note: if a user sets a custom base - we should ensure this works
@ -4724,7 +4725,7 @@ def embedding( # noqa: PLR0915
if headers is not None and headers != {}:
optional_params["extra_headers"] = headers
if encoding_format is not None:
optional_params["encoding_format"] = encoding_format
else:
@ -6789,9 +6790,7 @@ def speech( # noqa: PLR0915
if text_to_speech_provider_config is None:
text_to_speech_provider_config = MinimaxTextToSpeechConfig()
minimax_config = cast(
MinimaxTextToSpeechConfig, text_to_speech_provider_config
)
minimax_config = cast(MinimaxTextToSpeechConfig, text_to_speech_provider_config)
if api_base is not None:
litellm_params_dict["api_base"] = api_base
@ -6931,7 +6930,7 @@ async def ahealth_check(
custom_llm_provider_from_params = model_params.get("custom_llm_provider", None)
api_base_from_params = model_params.get("api_base", None)
api_key_from_params = model_params.get("api_key", None)
model, custom_llm_provider, _, _ = get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider_from_params,
@ -7305,8 +7304,9 @@ def __getattr__(name: str) -> Any:
_encoding = tiktoken.get_encoding("cl100k_base")
# Cache it in the module's __dict__ for subsequent accesses
import sys
sys.modules[__name__].__dict__["encoding"] = _encoding
global _encoding_cache
_encoding_cache = _encoding
return _encoding
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")

View file

@ -387,25 +387,57 @@ async def callback(code: str, state: str):
1. Try resource_metadata from WWW-Authenticate header (if present)
2. Fall back to path-based well-known URI: /.well-known/oauth-protected-resource/{path}
(
If the resource identifier value contains a path or query component, any terminating slash (/)
following the host component MUST be removed before inserting /.well-known/ and the well-known
URI path suffix between the host component and the path(include root path) and/or query components.
If the resource identifier value contains a path or query component, any terminating slash (/)
following the host component MUST be removed before inserting /.well-known/ and the well-known
URI path suffix between the host component and the path(include root path) and/or query components.
https://datatracker.ietf.org/doc/html/rfc9728#section-3.1)
3. Fall back to root-based well-known URI: /.well-known/oauth-protected-resource
Dual Pattern Support:
- Standard MCP pattern: /mcp/{server_name} (recommended, used by mcp-inspector, VSCode Copilot)
- LiteLLM legacy pattern: /{server_name}/mcp (backward compatibility)
The resource URL returned matches the pattern used in the discovery request.
"""
@router.get(f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}/mcp")
@router.get("/.well-known/oauth-protected-resource")
async def oauth_protected_resource_mcp(
request: Request, mcp_server_name: Optional[str] = None
):
def _build_oauth_protected_resource_response(
request: Request,
mcp_server_name: Optional[str],
use_standard_pattern: bool,
) -> dict:
"""
Build OAuth protected resource response with the appropriate URL pattern.
Args:
request: FastAPI Request object
mcp_server_name: Name of the MCP server
use_standard_pattern: If True, use /mcp/{server_name} pattern;
if False, use /{server_name}/mcp pattern
Returns:
OAuth protected resource metadata dict
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
# Get the correct base URL considering X-Forwarded-* headers
request_base_url = get_request_base_url(request)
mcp_server: Optional[MCPServer] = None
if mcp_server_name:
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name)
# Build resource URL based on the pattern
if mcp_server_name:
if use_standard_pattern:
# Standard MCP pattern: /mcp/{server_name}
resource_url = f"{request_base_url}/mcp/{mcp_server_name}"
else:
# LiteLLM legacy pattern: /{server_name}/mcp
resource_url = f"{request_base_url}/{mcp_server_name}/mcp"
else:
resource_url = f"{request_base_url}/mcp"
return {
"authorization_servers": [
(
@ -414,14 +446,55 @@ async def oauth_protected_resource_mcp(
else f"{request_base_url}"
)
],
"resource": (
f"{request_base_url}/{mcp_server_name}/mcp"
if mcp_server_name
else f"{request_base_url}/mcp"
), # this is what Claude will call
"resource": resource_url,
"scopes_supported": mcp_server.scopes if mcp_server else [],
}
# Standard MCP pattern: /.well-known/oauth-protected-resource/mcp/{server_name}
# This is the pattern expected by standard MCP clients (mcp-inspector, VSCode Copilot)
@router.get(f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/mcp/{{mcp_server_name}}")
async def oauth_protected_resource_mcp_standard(
request: Request, mcp_server_name: str
):
"""
OAuth protected resource discovery endpoint using standard MCP URL pattern.
Standard pattern: /mcp/{server_name}
Discovery path: /.well-known/oauth-protected-resource/mcp/{server_name}
This endpoint is compliant with MCP specification and works with standard
MCP clients like mcp-inspector and VSCode Copilot.
"""
return _build_oauth_protected_resource_response(
request=request,
mcp_server_name=mcp_server_name,
use_standard_pattern=True,
)
# LiteLLM legacy pattern: /.well-known/oauth-protected-resource/{server_name}/mcp
# Kept for backward compatibility with existing deployments
@router.get(f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}/mcp")
@router.get("/.well-known/oauth-protected-resource")
async def oauth_protected_resource_mcp(
request: Request, mcp_server_name: Optional[str] = None
):
"""
OAuth protected resource discovery endpoint using LiteLLM legacy URL pattern.
Legacy pattern: /{server_name}/mcp
Discovery path: /.well-known/oauth-protected-resource/{server_name}/mcp
This endpoint is kept for backward compatibility. New integrations should
use the standard MCP pattern (/mcp/{server_name}) instead.
"""
return _build_oauth_protected_resource_response(
request=request,
mcp_server_name=mcp_server_name,
use_standard_pattern=False,
)
"""
https://datatracker.ietf.org/doc/html/rfc8414#section-3.1
RFC 8414: Path-aware OAuth discovery
@ -430,15 +503,26 @@ async def oauth_protected_resource_mcp(
the well-known URI suffix between the host component and the path(include root path)
component.
"""
@router.get(f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}")
@router.get("/.well-known/oauth-authorization-server")
async def oauth_authorization_server_mcp(
request: Request, mcp_server_name: Optional[str] = None
):
def _build_oauth_authorization_server_response(
request: Request,
mcp_server_name: Optional[str],
) -> dict:
"""
Build OAuth authorization server metadata response.
Args:
request: FastAPI Request object
mcp_server_name: Name of the MCP server
Returns:
OAuth authorization server metadata dict
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
# Get the correct base URL considering X-Forwarded-* headers
request_base_url = get_request_base_url(request)
authorization_endpoint = (
@ -470,18 +554,58 @@ async def oauth_authorization_server_mcp(
}
# Standard MCP pattern: /.well-known/oauth-authorization-server/mcp/{server_name}
@router.get(f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/mcp/{{mcp_server_name}}")
async def oauth_authorization_server_mcp_standard(
request: Request, mcp_server_name: str
):
"""
OAuth authorization server discovery endpoint using standard MCP URL pattern.
Standard pattern: /mcp/{server_name}
Discovery path: /.well-known/oauth-authorization-server/mcp/{server_name}
"""
return _build_oauth_authorization_server_response(
request=request,
mcp_server_name=mcp_server_name,
)
# LiteLLM legacy pattern and root endpoint
@router.get(f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}")
@router.get("/.well-known/oauth-authorization-server")
async def oauth_authorization_server_mcp(
request: Request, mcp_server_name: Optional[str] = None
):
"""
OAuth authorization server discovery endpoint.
Supports both legacy pattern (/{server_name}) and root endpoint.
"""
return _build_oauth_authorization_server_response(
request=request,
mcp_server_name=mcp_server_name,
)
# Alias for standard OpenID discovery
@router.get("/.well-known/openid-configuration")
async def openid_configuration(request: Request):
return await oauth_authorization_server_mcp(request)
# Additional legacy pattern support
@router.get("/.well-known/oauth-authorization-server/{mcp_server_name}/mcp")
@router.get("/.well-known/oauth-authorization-server")
async def oauth_authorization_server_root(
request: Request, mcp_server_name: Optional[str] = None
async def oauth_authorization_server_legacy(
request: Request, mcp_server_name: str
):
return await oauth_authorization_server_mcp(request, mcp_server_name)
"""
OAuth authorization server discovery for legacy /{server_name}/mcp pattern.
"""
return _build_oauth_authorization_server_response(
request=request,
mcp_server_name=mcp_server_name,
)
@router.post("/{mcp_server_name}/register")

View file

@ -63,7 +63,20 @@ from litellm.types.mcp_server.mcp_server_manager import (
MCPOAuthMetadata,
MCPServer,
)
from mcp.shared.tool_name_validation import SEP_986_URL, validate_tool_name
try:
from mcp.shared.tool_name_validation import SEP_986_URL, validate_tool_name # type: ignore
except ImportError:
SEP_986_URL = "https://github.com/modelcontextprotocol/protocol/blob/main/proposals/0001-tool-name-validation.md"
def validate_tool_name(name: str):
from pydantic import BaseModel
class MockResult(BaseModel):
is_valid: bool = True
warnings: list = []
return MockResult()
# Probe includes characters on both sides of the separator to mimic real prefixed tool names.
@ -90,7 +103,9 @@ def _warn_on_server_name_fields(
if result.is_valid:
return
warning_text = "; ".join(result.warnings) if result.warnings else "Validation failed"
warning_text = (
"; ".join(result.warnings) if result.warnings else "Validation failed"
)
verbose_logger.warning(
"MCP server '%s' has invalid %s '%s': %s",
server_id,
@ -103,7 +118,6 @@ def _warn_on_server_name_fields(
_warn("server_name", server_name)
def _deserialize_json_dict(data: Any) -> Optional[Dict[str, str]]:
"""
Deserialize optional JSON mappings stored in the database.
@ -391,10 +405,13 @@ class MCPServerManager:
# Note: `extra_headers` on MCPServer is a List[str] of header names to forward
# from the client request (not available in this OpenAPI tool generation step).
# `static_headers` is a dict of concrete headers to always send.
headers = merge_mcp_headers(
extra_headers=headers,
static_headers=server.static_headers,
) or {}
headers = (
merge_mcp_headers(
extra_headers=headers,
static_headers=server.static_headers,
)
or {}
)
verbose_logger.debug(
f"Using headers for OpenAPI tools (excluding sensitive values): "

View file

@ -73,7 +73,11 @@ if MCP_AVAILABLE:
AuthContextMiddleware,
auth_context_var,
)
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
try:
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
except ImportError:
StreamableHTTPSessionManager = None # type: ignore
from mcp.types import (
CallToolResult,
EmbeddedResource,

View file

@ -425,12 +425,12 @@ class LiteLLMRoutes(enum.Enum):
]
google_routes = [
"/v1beta/models/{model_name}:countTokens",
"/v1beta/models/{model_name}:generateContent",
"/v1beta/models/{model_name}:streamGenerateContent",
"/models/{model_name}:countTokens",
"/models/{model_name}:generateContent",
"/models/{model_name}:streamGenerateContent",
"/v1beta/models/{model_name:path}:countTokens",
"/v1beta/models/{model_name:path}:generateContent",
"/v1beta/models/{model_name:path}:streamGenerateContent",
"/models/{model_name:path}:countTokens",
"/models/{model_name:path}:generateContent",
"/models/{model_name:path}:streamGenerateContent",
# Google Interactions API
"/interactions",
"/v1beta/interactions",

View file

@ -758,11 +758,27 @@ def get_model_from_request(
if match:
model = match.group(1)
# If still not found, extract model from Google generateContent-style routes.
# These routes put the model in the path and allow "/" inside the model id.
# Examples:
# - /v1beta/models/gemini-2.0-flash:generateContent
# - /v1beta/models/bedrock/claude-sonnet-3.7:generateContent
# - /models/custom/ns/model:streamGenerateContent
if model is None and not route.lower().startswith("/vertex"):
google_match = re.search(r"/(?:v1beta|beta)/models/([^:]+):", route)
if google_match:
model = google_match.group(1)
if model is None and not route.lower().startswith("/vertex"):
google_match = re.search(r"^/models/([^:]+):", route)
if google_match:
model = google_match.group(1)
# If still not found, extract from Vertex AI passthrough route
# Pattern: /vertex_ai/.../models/{model_id}:*
# Example: /vertex_ai/v1/.../models/gemini-1.5-pro:generateContent
if model is None and "/vertex" in route.lower():
vertex_match = re.search(r"/models/([^/:]+)", route)
if model is None and route.lower().startswith("/vertex"):
vertex_match = re.search(r"/models/([^:]+)", route)
if vertex_match:
model = vertex_match.group(1)

View file

@ -392,7 +392,15 @@ class RouteChecks:
# Ensure route is a string before attempting regex matching
if not isinstance(route, str):
return False
pattern = re.sub(r"\{[^}]+\}", r"[^/]+", pattern)
def _placeholder_to_regex(match: re.Match) -> str:
placeholder = match.group(0).strip("{}")
if placeholder.endswith(":path"):
# allow "/" in the placeholder value, but don't eat the route suffix after ":"
return r"[^:]+"
return r"[^/]+"
pattern = re.sub(r"\{[^}]+\}", _placeholder_to_regex, pattern)
# Anchor the pattern to match the entire string
pattern = f"^{pattern}$"
if re.match(pattern, route):

View file

@ -8,6 +8,7 @@ import os
import uuid
from typing import TYPE_CHECKING, Any, Literal, Optional, Type
import httpx
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
@ -25,10 +26,12 @@ if TYPE_CHECKING:
class OnyxGuardrail(CustomGuardrail):
def __init__(
self, api_base: Optional[str] = None, api_key: Optional[str] = None, **kwargs
self, api_base: Optional[str] = None, api_key: Optional[str] = None, timeout: Optional[float] = 10.0, **kwargs
):
timeout = timeout or int(os.getenv("ONYX_TIMEOUT", 10.0))
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
llm_provider=httpxSpecialProvider.GuardrailCallback,
params={"timeout": httpx.Timeout(timeout=timeout, connect=5.0)},
)
self.api_base = api_base or os.getenv(
"ONYX_API_BASE",

View file

@ -856,7 +856,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915
# Add headers to metadata for guardrails to access (fixes #17477)
# Guardrails use metadata["headers"] to access request headers (e.g., User-Agent)
if _metadata_variable_name in data and isinstance(data[_metadata_variable_name], dict):
if _metadata_variable_name in data and isinstance(
data[_metadata_variable_name], dict
):
data[_metadata_variable_name]["headers"] = _headers
# check for forwardable headers
@ -1012,7 +1014,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915
# User spend, budget - used by prometheus.py
# Follow same pattern as team and API key budgets
data[_metadata_variable_name]["user_api_key_user_spend"] = user_api_key_dict.user_spend
data[_metadata_variable_name][
"user_api_key_user_spend"
] = user_api_key_dict.user_spend
data[_metadata_variable_name][
"user_api_key_user_max_budget"
] = user_api_key_dict.user_max_budget
@ -1039,8 +1043,8 @@ async def add_litellm_data_to_request( # noqa: PLR0915
## [Enterprise Only]
# Add User-IP Address
requester_ip_address = ""
if premium_user is True:
# Only set the IP Address for Enterprise Users
if True: # Always set the IP Address if available
# logic for tracking IP Address
# logic for tracking IP Address
if (
@ -1060,6 +1064,16 @@ async def add_litellm_data_to_request( # noqa: PLR0915
requester_ip_address = request.client.host
data[_metadata_variable_name]["requester_ip_address"] = requester_ip_address
# Add User-Agent
user_agent = ""
if (
request is not None
and hasattr(request, "headers")
and "user-agent" in request.headers
):
user_agent = request.headers["user-agent"]
data[_metadata_variable_name]["user_agent"] = user_agent
# Check if using tag based routing
tags = LiteLLMProxyRequestSetup.add_request_tag_to_metadata(
llm_router=llm_router,
@ -1542,7 +1556,9 @@ def add_guardrails_from_policy_engine(
f"policy_count={len(registry.get_all_policies())}"
)
if not registry.is_initialized():
verbose_proxy_logger.debug("Policy engine not initialized, skipping policy matching")
verbose_proxy_logger.debug(
"Policy engine not initialized, skipping policy matching"
)
return
# Build context from request
@ -1560,13 +1576,17 @@ def add_guardrails_from_policy_engine(
# Get matching policies via attachments
matching_policy_names = PolicyMatcher.get_matching_policies(context=context)
verbose_proxy_logger.debug(f"Policy engine: matched policies via attachments: {matching_policy_names}")
verbose_proxy_logger.debug(
f"Policy engine: matched policies via attachments: {matching_policy_names}"
)
# Combine attachment-based policies with dynamic request body policies
all_policy_names = set(matching_policy_names)
if request_body_policies and isinstance(request_body_policies, list):
all_policy_names.update(request_body_policies)
verbose_proxy_logger.debug(f"Policy engine: added dynamic policies from request body: {request_body_policies}")
verbose_proxy_logger.debug(
f"Policy engine: added dynamic policies from request body: {request_body_policies}"
)
if not all_policy_names:
return
@ -1577,7 +1597,9 @@ def add_guardrails_from_policy_engine(
context=context,
)
verbose_proxy_logger.debug(f"Policy engine: applied policies (conditions matched): {applied_policy_names}")
verbose_proxy_logger.debug(
f"Policy engine: applied policies (conditions matched): {applied_policy_names}"
)
# Track applied policies in metadata for response headers
for policy_name in applied_policy_names:
@ -1588,7 +1610,9 @@ def add_guardrails_from_policy_engine(
# Resolve guardrails from matching policies
resolved_guardrails = PolicyResolver.resolve_guardrails_for_context(context=context)
verbose_proxy_logger.debug(f"Policy engine: resolved guardrails: {resolved_guardrails}")
verbose_proxy_logger.debug(
f"Policy engine: resolved guardrails: {resolved_guardrails}"
)
if not resolved_guardrails:
return

View file

@ -56,7 +56,19 @@ except ImportError as e:
MCP_AVAILABLE = False
if MCP_AVAILABLE:
from mcp.shared.tool_name_validation import validate_tool_name
try:
from mcp.shared.tool_name_validation import validate_tool_name # type: ignore
except ImportError:
def validate_tool_name(name: str):
from pydantic import BaseModel
class MockResult(BaseModel):
is_valid: bool = True
warnings: list = []
return MockResult()
from litellm.proxy._experimental.mcp_server.db import (
create_mcp_server,
delete_mcp_server,
@ -122,9 +134,7 @@ if MCP_AVAILABLE:
)
if validation_result.warnings:
error_messages_text = (
error_messages_text
+ "\n"
+ "\n".join(validation_result.warnings)
error_messages_text + "\n" + "\n".join(validation_result.warnings)
)
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,

View file

@ -9700,7 +9700,7 @@ def get_logo_url():
@app.get("/get_image", include_in_schema=False)
def get_image():
async def get_image():
"""Get logo to show on admin UI"""
# get current_dir
@ -9719,25 +9719,37 @@ def get_image():
if is_non_root and not os.path.exists(default_logo):
default_logo = default_site_logo
cache_dir = assets_dir if is_non_root else current_dir
cache_path = os.path.join(cache_dir, "cached_logo.jpg")
# [OPTIMIZATION] Check if the cached image exists first
if os.path.exists(cache_path):
return FileResponse(cache_path, media_type="image/jpeg")
logo_path = os.getenv("UI_LOGO_PATH", default_logo)
verbose_proxy_logger.debug("Reading logo from path: %s", logo_path)
# Check if the logo path is an HTTP/HTTPS URL
if logo_path.startswith(("http://", "https://")):
# Download the image and cache it
client = HTTPHandler()
response = client.get(logo_path)
if response.status_code == 200:
# Save the image to a local file
cache_dir = assets_dir if is_non_root else current_dir
cache_path = os.path.join(cache_dir, "cached_logo.jpg")
with open(cache_path, "wb") as f:
f.write(response.content)
try:
# Download the image and cache it
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
# Return the cached image as a FileResponse
return FileResponse(cache_path, media_type="image/jpeg")
else:
# Handle the case when the image cannot be downloaded
async_client = AsyncHTTPHandler(timeout=5.0)
response = await async_client.get(logo_path)
if response.status_code == 200:
# Save the image to a local file
with open(cache_path, "wb") as f:
f.write(response.content)
# Return the cached image as a FileResponse
return FileResponse(cache_path, media_type="image/jpeg")
else:
# Handle the case when the image cannot be downloaded
return FileResponse(default_logo, media_type="image/jpeg")
except Exception as e:
# Handle any exceptions during the download (e.g., timeout, connection error)
verbose_proxy_logger.debug(f"Error downloading logo from {logo_path}: {e}")
return FileResponse(default_logo, media_type="image/jpeg")
else:
# Return the local image file if the logo path is not an HTTP/HTTPS URL

View file

@ -150,6 +150,9 @@ class UserAPIKeyLabelNames(Enum):
FALLBACK_MODEL = "fallback_model"
ROUTE = "route"
MODEL_GROUP = "model_group"
CLIENT_IP = "client_ip"
USER_AGENT = "user_agent"
CALLBACK_NAME = "callback_name"
DEFINED_PROMETHEUS_METRICS = Literal[
@ -196,6 +199,12 @@ DEFINED_PROMETHEUS_METRICS = Literal[
"litellm_cache_hits_metric",
"litellm_cache_misses_metric",
"litellm_cached_tokens_metric",
"litellm_deployment_tpm_limit",
"litellm_deployment_rpm_limit",
"litellm_remaining_api_key_requests_for_model",
"litellm_remaining_api_key_tokens_for_model",
"litellm_llm_api_failed_requests_metric",
"litellm_callback_logging_failures_metric",
]
@ -209,6 +218,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
UserAPIKeyLabelNames.END_USER.value,
UserAPIKeyLabelNames.USER.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_llm_api_time_to_first_token_metric = [
@ -217,6 +227,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.TEAM.value,
UserAPIKeyLabelNames.TEAM_ALIAS.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_request_total_latency_metric = [
@ -228,6 +239,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.TEAM_ALIAS.value,
UserAPIKeyLabelNames.USER.value,
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_request_queue_time_seconds = [
@ -239,6 +251,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.TEAM_ALIAS.value,
UserAPIKeyLabelNames.USER.value,
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
# Guardrail metrics - these use custom labels (guardrail_name, status, error_type, hook_type)
@ -258,6 +271,9 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.STATUS_CODE.value,
UserAPIKeyLabelNames.USER_EMAIL.value,
UserAPIKeyLabelNames.ROUTE.value,
UserAPIKeyLabelNames.CLIENT_IP.value,
UserAPIKeyLabelNames.USER_AGENT.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_proxy_failed_requests_metric = [
@ -272,6 +288,9 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.EXCEPTION_STATUS.value,
UserAPIKeyLabelNames.EXCEPTION_CLASS.value,
UserAPIKeyLabelNames.ROUTE.value,
UserAPIKeyLabelNames.CLIENT_IP.value,
UserAPIKeyLabelNames.USER_AGENT.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_deployment_latency_per_output_token = [
@ -292,6 +311,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value,
UserAPIKeyLabelNames.API_KEY_HASH.value,
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_remaining_requests_metric = [
@ -301,6 +321,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value,
UserAPIKeyLabelNames.API_KEY_HASH.value,
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_remaining_tokens_metric = [
@ -310,6 +331,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value,
UserAPIKeyLabelNames.API_KEY_HASH.value,
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_requests_metric = [
@ -321,6 +343,9 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.TEAM_ALIAS.value,
UserAPIKeyLabelNames.USER.value,
UserAPIKeyLabelNames.USER_EMAIL.value,
UserAPIKeyLabelNames.CLIENT_IP.value,
UserAPIKeyLabelNames.USER_AGENT.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_spend_metric = [
@ -332,6 +357,9 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.TEAM_ALIAS.value,
UserAPIKeyLabelNames.USER.value,
UserAPIKeyLabelNames.USER_EMAIL.value,
UserAPIKeyLabelNames.CLIENT_IP.value,
UserAPIKeyLabelNames.USER_AGENT.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_input_tokens_metric = [
@ -344,6 +372,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.USER.value,
UserAPIKeyLabelNames.USER_EMAIL.value,
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_total_tokens_metric = [
@ -356,6 +385,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.USER.value,
UserAPIKeyLabelNames.USER_EMAIL.value,
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_output_tokens_metric = [
@ -368,6 +398,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.USER.value,
UserAPIKeyLabelNames.USER_EMAIL.value,
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_deployment_state = [
@ -377,6 +408,15 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.API_PROVIDER.value,
]
litellm_deployment_tpm_limit = [
UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value,
UserAPIKeyLabelNames.MODEL_ID.value,
UserAPIKeyLabelNames.API_BASE.value,
UserAPIKeyLabelNames.API_PROVIDER.value,
]
litellm_deployment_rpm_limit = litellm_deployment_tpm_limit
litellm_deployment_cooled_down = [
UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value,
UserAPIKeyLabelNames.MODEL_ID.value,
@ -394,6 +434,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.TEAM_ALIAS.value,
UserAPIKeyLabelNames.EXCEPTION_STATUS.value,
UserAPIKeyLabelNames.EXCEPTION_CLASS.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_deployment_failed_fallbacks = litellm_deployment_successful_fallbacks
@ -436,6 +477,26 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.USER.value,
]
litellm_user_budget_remaining_hours_metric = [
UserAPIKeyLabelNames.USER.value,
]
litellm_remaining_api_key_requests_for_model = [
UserAPIKeyLabelNames.API_KEY_HASH.value,
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
]
litellm_remaining_api_key_tokens_for_model = [
UserAPIKeyLabelNames.API_KEY_HASH.value,
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
]
litellm_callback_logging_failures_metric = [
UserAPIKeyLabelNames.CALLBACK_NAME.value,
]
# Add deployment metrics
litellm_deployment_failure_responses = [
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
@ -449,6 +510,8 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.TEAM.value,
UserAPIKeyLabelNames.TEAM_ALIAS.value,
UserAPIKeyLabelNames.CLIENT_IP.value,
UserAPIKeyLabelNames.USER_AGENT.value,
]
litellm_deployment_total_requests = [
@ -461,10 +524,37 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.TEAM.value,
UserAPIKeyLabelNames.TEAM_ALIAS.value,
UserAPIKeyLabelNames.CLIENT_IP.value,
UserAPIKeyLabelNames.USER_AGENT.value,
]
litellm_deployment_success_responses = litellm_deployment_total_requests
litellm_remaining_api_key_requests_for_model = [
UserAPIKeyLabelNames.API_KEY_HASH.value,
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_remaining_api_key_tokens_for_model = [
UserAPIKeyLabelNames.API_KEY_HASH.value,
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_llm_api_failed_requests_metric = [
UserAPIKeyLabelNames.END_USER.value,
UserAPIKeyLabelNames.API_KEY_HASH.value,
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
UserAPIKeyLabelNames.TEAM.value,
UserAPIKeyLabelNames.TEAM_ALIAS.value,
UserAPIKeyLabelNames.USER.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
# Buffer monitoring metrics - these typically don't need additional labels
litellm_pod_lock_manager_size: List[str] = []
@ -485,6 +575,7 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.TEAM_ALIAS.value,
UserAPIKeyLabelNames.END_USER.value,
UserAPIKeyLabelNames.USER.value,
UserAPIKeyLabelNames.MODEL_ID.value,
]
litellm_cache_hits_metric = _cache_metric_labels
@ -577,6 +668,12 @@ class UserAPIKeyLabelValues(BaseModel):
route: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.ROUTE.value)
] = None
client_ip: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.CLIENT_IP.value)
] = None
user_agent: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.USER_AGENT.value)
] = None
class PrometheusMetricsConfig(BaseModel):

View file

@ -16,6 +16,11 @@ class OnyxGuardrailConfigModel(GuardrailConfigModel):
description="The API key for the Onyx Guard server. If not provided, the `ONYX_API_KEY` environment variable is checked.",
)
timeout: Optional[float] = Field(
default=None,
description="The timeout for the Onyx Guard server in seconds. If not provided, the `ONYX_TIMEOUT` environment variable is checked.",
)
@staticmethod
def ui_friendly_name() -> str:
return "Onyx Guardrail"

View file

@ -3,25 +3,26 @@ import time
from enum import Enum
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Mapping, Optional, Union
from aiohttp import FormData
from openai._models import BaseModel as OpenAIObject
from openai.types.audio.transcription_create_params import FileTypes # type: ignore
from openai.types.chat.chat_completion import ChatCompletion
from openai.types.audio.transcription_create_params import FileTypes as FileTypes # type: ignore
from openai.types.chat.chat_completion import ChatCompletion as ChatCompletion
from openai.types.completion_usage import (
CompletionTokensDetails,
CompletionUsage,
PromptTokensDetails,
)
from openai.types.moderation import (
Categories,
CategoryAppliedInputTypes,
CategoryScores,
Categories as Categories,
CategoryAppliedInputTypes as CategoryAppliedInputTypes,
CategoryScores as CategoryScores,
)
from openai.types.moderation_create_response import (
Moderation as Moderation,
ModerationCreateResponse as ModerationCreateResponse,
)
from openai.types.moderation_create_response import Moderation, ModerationCreateResponse
from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, model_validator
from typing_extensions import Callable, Dict, Required, TypedDict, override
from typing_extensions import Required, TypedDict
import litellm
from litellm._uuid import uuid
from litellm.types.llms.base import (
BaseLiteLLMOpenAIResponseObject,
@ -52,7 +53,7 @@ from .llms.openai import (
ResponsesAPIResponse,
WebSearchOptions,
)
from .rerank import RerankResponse
from .rerank import RerankResponse as RerankResponse
if TYPE_CHECKING:
from .vector_stores import VectorStoreSearchResponse
@ -1411,7 +1412,7 @@ class Usage(SafeAttributeModel, CompletionUsage):
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None
"""Breakdown of tokens used in the prompt."""
def __init__(
def __init__( # noqa: PLR0915
self,
prompt_tokens: Optional[int] = None,
completion_tokens: Optional[int] = None,
@ -2501,6 +2502,7 @@ class StandardLoggingMetadata(StandardLoggingUserAPIKeyMetadata):
dict
] # special param to log k,v pairs to spendlogs for a call
requester_ip_address: Optional[str]
user_agent: Optional[str]
requester_metadata: Optional[dict]
requester_custom_headers: Optional[
Dict[str, str]
@ -2686,6 +2688,7 @@ class StandardLoggingPayload(TypedDict):
request_tags: list
end_user: Optional[str]
requester_ip_address: Optional[str]
user_agent: Optional[str]
messages: Optional[Union[str, list, dict]]
response: Optional[Union[str, list, dict]]
error_str: Optional[str]

File diff suppressed because it is too large Load diff

View file

@ -683,9 +683,11 @@ async def test_streaming_responses_api_with_mcp_tools(
Return the user the result of request 2
"""
# Skip test if ANTHROPIC_API_KEY is not set for anthropic models
if "anthropic" in model.lower() and not os.getenv("ANTHROPIC_API_KEY"):
# Skip test if required API keys are not set
if ("anthropic" in model.lower() or "claude" in model.lower()) and not os.getenv("ANTHROPIC_API_KEY"):
pytest.skip("ANTHROPIC_API_KEY not set, skipping anthropic model test")
if ("gpt" in model.lower() or "openai" in model.lower()) and not os.getenv("OPENAI_API_KEY"):
pytest.skip("OPENAI_API_KEY not set, skipping openai model test")
from unittest.mock import AsyncMock, patch

View file

@ -1105,14 +1105,24 @@ async def test_mcp_server_manager_config_integration_with_database():
test_manager.get_allowed_mcp_servers = mock_get_allowed_servers
# Mock _create_mcp_client to return a client that completes immediately
# This avoids network calls while preserving the actual conversion logic
def mock_create_mcp_client(*args, **kwargs):
mock_client = MagicMock()
mock_client.run_with_session = AsyncMock(return_value="ok")
return mock_client
test_manager._create_mcp_client = mock_create_mcp_client
# Mock health_check_server to avoid real network calls that timeout
async def mock_health_check(server_id: str, mcp_auth_header=None):
server = test_manager.get_mcp_server_by_id(server_id)
if not server:
return None
return LiteLLM_MCPServerTable(
server_id=server_id,
server_name=server.name,
url=server.url,
transport=server.transport,
description=server.mcp_info.get("description") if server.mcp_info else None,
mcp_access_groups=server.access_groups,
status="healthy",
last_health_check=datetime.datetime.now(),
mcp_info=server.mcp_info,
)
test_manager.health_check_server = mock_health_check
# Test the method (this tests our second fix)
servers_list = await test_manager.get_all_mcp_servers_with_health_and_teams(

View file

@ -0,0 +1,89 @@
import os
import sys
from unittest import mock
# Standard path insertion
sys.path.insert(0, os.path.abspath("../.."))
import pytest
import httpx
from litellm.proxy.proxy_server import app
@pytest.mark.asyncio
async def test_get_image_error_handling():
"""
Test that get_image handles network errors gracefully and doesn't hang.
"""
# Set an unreachable URL
os.environ["UI_LOGO_PATH"] = "http://invalid-url-12345.com/logo.jpg"
# Clear cache
parent_dir = os.path.dirname(
os.path.dirname(
app.__file__
if hasattr(app, "__file__")
else "litellm/proxy/proxy_server.py"
)
)
cache_path = os.path.join(parent_dir, "proxy", "cached_logo.jpg")
if os.path.exists(cache_path):
os.remove(cache_path)
# Mock AsyncHTTPHandler to simulate a timeout or connection error
with mock.patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get"
) as mock_get:
mock_get.side_effect = httpx.ConnectError("Network is unreachable")
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app), base_url="http://testserver"
) as ac:
response = await ac.get("/get_image")
assert response.status_code == 200
assert response.headers["content-type"] == "image/jpeg"
@pytest.mark.asyncio
async def test_get_image_cache_logic():
"""
Test that once cached, get_image doesn't hit the network.
"""
os.environ["UI_LOGO_PATH"] = "http://example.com/logo.jpg"
# Clear cache
parent_dir = os.path.dirname(
os.path.dirname(
app.__file__
if hasattr(app, "__file__")
else "litellm/proxy/proxy_server.py"
)
)
cache_path = os.path.join(parent_dir, "proxy", "cached_logo.jpg")
if os.path.exists(cache_path):
os.remove(cache_path)
# Mock response
mock_response = mock.Mock()
mock_response.status_code = 200
mock_response.content = b"fake image data"
with mock.patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get"
) as mock_get:
mock_get.return_value = mock_response
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app), base_url="http://testserver"
) as ac:
# First call - should hit download logic
response1 = await ac.get("/get_image")
assert response1.status_code == 200
assert mock_get.call_count == 1
# Second call - should hit cache
response2 = await ac.get("/get_image")
assert response2.status_code == 200
# If cache works, mock_get shouldn't be called again
assert mock_get.call_count == 1

View file

@ -0,0 +1,203 @@
import pytest
from unittest.mock import MagicMock, patch
from litellm.integrations.prometheus import PrometheusLogger
from litellm.types.integrations.prometheus import (
UserAPIKeyLabelValues,
)
from litellm.proxy._types import UserAPIKeyAuth
@pytest.mark.asyncio
async def test_async_post_call_failure_hook_includes_client_ip_user_agent():
"""
Test that async_post_call_failure_hook includes client_ip and user_agent in UserAPIKeyLabelValues
"""
# Mocking
# Mocking
with patch(
"litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None
):
logger = PrometheusLogger()
# Initialize attributes manually as __init__ is mocked
logger.litellm_proxy_failed_requests_metric = MagicMock()
logger.litellm_proxy_total_requests_metric = MagicMock()
logger.get_labels_for_metric = MagicMock(
return_value=["client_ip", "user_agent"]
)
request_data = {
"model": "gpt-4",
"metadata": {
"requester_ip_address": "127.0.0.1",
"user_agent": "test-agent",
},
}
user_api_key_dict = UserAPIKeyAuth(token="test_token")
original_exception = Exception("Test exception")
# Mock prometheus_label_factory to inspect arguments
with patch(
"litellm.integrations.prometheus.prometheus_label_factory"
) as mock_label_factory:
mock_label_factory.return_value = {}
await logger.async_post_call_failure_hook(
request_data=request_data,
original_exception=original_exception,
user_api_key_dict=user_api_key_dict,
)
# Verification
assert mock_label_factory.call_count >= 1
# Check calls
calls = mock_label_factory.call_args_list
found = False
for call in calls:
kwargs = call.kwargs
enum_values = kwargs.get("enum_values")
if isinstance(enum_values, UserAPIKeyLabelValues):
if (
enum_values.client_ip == "127.0.0.1"
and enum_values.user_agent == "test-agent"
):
found = True
break
assert (
found
), "UserAPIKeyLabelValues should contain client_ip='127.0.0.1' and user_agent='test-agent'"
@pytest.mark.asyncio
async def test_async_post_call_success_hook_includes_client_ip_user_agent():
"""
Test that async_post_call_success_hook includes client_ip and user_agent in UserAPIKeyLabelValues
"""
# Mocking
# Mocking
with patch(
"litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None
):
logger = PrometheusLogger()
logger.litellm_proxy_total_requests_metric = MagicMock()
logger.get_labels_for_metric = MagicMock(
return_value=["client_ip", "user_agent"]
)
data = {
"model": "gpt-4",
"metadata": {
"requester_ip_address": "192.168.1.1",
"user_agent": "success-agent",
},
}
user_api_key_dict = UserAPIKeyAuth(token="test_token")
response = MagicMock()
# Mock prometheus_label_factory to inspect arguments
with patch(
"litellm.integrations.prometheus.prometheus_label_factory"
) as mock_label_factory:
mock_label_factory.return_value = {}
await logger.async_post_call_success_hook(
data=data,
user_api_key_dict=user_api_key_dict,
response=response,
)
# Verification
assert mock_label_factory.call_count >= 1
# Check calls
calls = mock_label_factory.call_args_list
found = False
for call in calls:
kwargs = call.kwargs
enum_values = kwargs.get("enum_values")
if isinstance(enum_values, UserAPIKeyLabelValues):
if (
enum_values.client_ip == "192.168.1.1"
and enum_values.user_agent == "success-agent"
):
found = True
break
assert (
found
), "UserAPIKeyLabelValues should contain client_ip='192.168.1.1' and user_agent='success-agent'"
def test_set_llm_deployment_failure_metrics_includes_client_ip_user_agent():
"""
Test that set_llm_deployment_failure_metrics includes client_ip and user_agent in UserAPIKeyLabelValues
"""
# Mocking
# Mocking
with patch(
"litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None
):
logger = PrometheusLogger()
logger.litellm_deployment_failure_responses = MagicMock()
logger.litellm_deployment_total_requests = MagicMock()
logger.get_labels_for_metric = MagicMock(
return_value=["client_ip", "user_agent"]
)
logger.set_deployment_partial_outage = MagicMock()
request_kwargs = {
"model": "gpt-4",
"standard_logging_object": {
"metadata": {
"requester_ip_address": "10.0.0.1",
"user_agent": "failure-deployment",
"user_api_key_team_id": "team_1",
"user_api_key_team_alias": "team_alias_1",
"user_api_key_alias": "key_alias_1",
},
"model_group": "group_1",
"api_base": "http://api.base",
"model_id": "model_1",
},
"litellm_params": {},
"exception": Exception("Deployment failure"),
}
# Mock prometheus_label_factory to inspect arguments
with patch(
"litellm.integrations.prometheus.prometheus_label_factory"
) as mock_label_factory:
mock_label_factory.return_value = {}
logger.set_llm_deployment_failure_metrics(request_kwargs=request_kwargs)
# Verification
assert mock_label_factory.call_count >= 1
# Check calls
calls = mock_label_factory.call_args_list
found = False
for call in calls:
kwargs = call.kwargs
enum_values = kwargs.get("enum_values")
if isinstance(enum_values, UserAPIKeyLabelValues):
if (
enum_values.client_ip == "10.0.0.1"
and enum_values.user_agent == "failure-deployment"
):
found = True
break
assert (
found
), "UserAPIKeyLabelValues should contain client_ip='10.0.0.1' and user_agent='failure-deployment'"
if __name__ == "__main__":
import asyncio
asyncio.run(test_async_post_call_failure_hook_includes_client_ip_user_agent())
asyncio.run(test_async_post_call_success_hook_includes_client_ip_user_agent())
test_set_llm_deployment_failure_metrics_includes_client_ip_user_agent()
print("✅ All client_ip and user_agent tests passed!")

View file

@ -26,15 +26,49 @@ def test_user_email_in_required_metrics():
"litellm_input_tokens_metric",
"litellm_output_tokens_metric",
"litellm_requests_metric",
"litellm_spend_metric"
"litellm_spend_metric",
]
for metric_name in metrics_with_user_email:
labels = PrometheusMetricLabels.get_labels(metric_name)
assert user_email_label in labels, f"Metric {metric_name} should contain user_email label"
assert (
user_email_label in labels
), f"Metric {metric_name} should contain user_email label"
print(f"✅ {metric_name} contains user_email label")
def test_model_id_in_required_metrics():
"""
Test that model_id label is present in all the metrics that should have it
"""
model_id_label = UserAPIKeyLabelNames.MODEL_ID.value
# Metrics that should have model_id
metrics_with_model_id = [
"litellm_proxy_total_requests_metric",
"litellm_proxy_failed_requests_metric",
"litellm_input_tokens_metric",
"litellm_output_tokens_metric",
"litellm_requests_metric",
"litellm_spend_metric",
"litellm_llm_api_latency_metric",
"litellm_remaining_requests_metric",
"litellm_deployment_successful_fallbacks",
"litellm_cache_hits_metric",
"litellm_cache_misses_metric",
"litellm_remaining_api_key_requests_for_model",
"litellm_remaining_api_key_tokens_for_model",
"litellm_llm_api_failed_requests_metric",
]
for metric_name in metrics_with_model_id:
labels = PrometheusMetricLabels.get_labels(metric_name)
assert (
model_id_label in labels
), f"Metric {metric_name} should contain model_id label"
print(f"✅ {metric_name} contains model_id label")
def test_user_email_label_exists():
"""Test that the USER_EMAIL label is properly defined"""
assert UserAPIKeyLabelNames.USER_EMAIL.value == "user_email"
@ -52,12 +86,14 @@ def test_prometheus_metric_labels_structure():
"litellm_proxy_failed_requests_metric",
"litellm_input_tokens_metric",
"litellm_output_tokens_metric",
"litellm_spend_metric"
"litellm_spend_metric",
]
for metric_name in test_metrics:
# Check metric is in DEFINED_PROMETHEUS_METRICS
assert metric_name in get_args(DEFINED_PROMETHEUS_METRICS), f"{metric_name} should be in DEFINED_PROMETHEUS_METRICS"
assert metric_name in get_args(
DEFINED_PROMETHEUS_METRICS
), f"{metric_name} should be in DEFINED_PROMETHEUS_METRICS"
# Check labels can be retrieved
labels = PrometheusMetricLabels.get_labels(metric_name)
@ -74,11 +110,11 @@ def test_route_normalization_for_responses_api():
"""
Test that route normalization prevents high cardinality in Prometheus metrics
for the /v1/responses/{response_id} endpoint.
Issue: https://github.com/BerriAI/litellm/issues/XXXX
Each unique response ID was creating a separate metric line, causing the
/metrics endpoint to grow to ~30MB and take ~40 seconds to respond.
Fix: Routes are normalized to collapse dynamic IDs into placeholders.
"""
from litellm.proxy.auth.auth_utils import normalize_request_route
@ -91,43 +127,53 @@ def test_route_normalization_for_responses_api():
("/v1/responses/resp_abc123", "/v1/responses/{response_id}"),
("/v1/responses/litellm_poll_xyz", "/v1/responses/{response_id}"),
]
for original, expected in responses_routes:
normalized = normalize_request_route(original)
assert normalized == expected, \
f"Failed: {original} -> {normalized} (expected {expected})"
assert (
normalized == expected
), f"Failed: {original} -> {normalized} (expected {expected})"
# Verify cardinality reduction
unique_normalized = set(normalize_request_route(route) for route, _ in responses_routes)
assert len(unique_normalized) == 1, \
f"Expected 1 unique normalized route, got {len(unique_normalized)}: {unique_normalized}"
print(f"✅ Responses API routes: {len(responses_routes)} different IDs normalized to 1 metric label")
unique_normalized = set(
normalize_request_route(route) for route, _ in responses_routes
)
assert (
len(unique_normalized) == 1
), f"Expected 1 unique normalized route, got {len(unique_normalized)}: {unique_normalized}"
print(
f"✅ Responses API routes: {len(responses_routes)} different IDs normalized to 1 metric label"
)
def test_route_normalization_for_sub_routes():
"""Test that sub-routes like /cancel and /input_items are normalized correctly"""
from litellm.proxy.auth.auth_utils import normalize_request_route
sub_routes = [
("/v1/responses/id1/cancel", "/v1/responses/{response_id}/cancel"),
("/v1/responses/id2/cancel", "/v1/responses/{response_id}/cancel"),
("/v1/responses/id3/input_items", "/v1/responses/{response_id}/input_items"),
("/openai/v1/responses/id4/input_items", "/openai/v1/responses/{response_id}/input_items"),
(
"/openai/v1/responses/id4/input_items",
"/openai/v1/responses/{response_id}/input_items",
),
]
for original, expected in sub_routes:
normalized = normalize_request_route(original)
assert normalized == expected, \
f"Failed: {original} -> {normalized} (expected {expected})"
assert (
normalized == expected
), f"Failed: {original} -> {normalized} (expected {expected})"
print("✅ Sub-routes normalized correctly")
def test_route_normalization_preserves_static_routes():
"""Test that static routes are not affected by normalization"""
from litellm.proxy.auth.auth_utils import normalize_request_route
static_routes = [
"/chat/completions",
"/v1/chat/completions",
@ -137,46 +183,47 @@ def test_route_normalization_preserves_static_routes():
"/v1/models",
"/v1/responses", # List endpoint without ID
]
for route in static_routes:
normalized = normalize_request_route(route)
assert normalized == route, \
f"Static route should not be modified: {route} -> {normalized}"
assert (
normalized == route
), f"Static route should not be modified: {route} -> {normalized}"
print(f"✅ {len(static_routes)} static routes preserved")
def test_route_normalization_other_dynamic_apis():
"""Test normalization for other OpenAI-compatible APIs with dynamic IDs"""
from litellm.proxy.auth.auth_utils import normalize_request_route
test_cases = [
# Threads API
("/v1/threads/thread_123", "/v1/threads/{thread_id}"),
("/v1/threads/thread_abc/messages", "/v1/threads/{thread_id}/messages"),
("/v1/threads/thread_abc/runs/run_123", "/v1/threads/{thread_id}/runs/{run_id}"),
(
"/v1/threads/thread_abc/runs/run_123",
"/v1/threads/{thread_id}/runs/{run_id}",
),
# Vector Stores API
("/v1/vector_stores/vs_123", "/v1/vector_stores/{vector_store_id}"),
("/v1/vector_stores/vs_123/files", "/v1/vector_stores/{vector_store_id}/files"),
# Assistants API
("/v1/assistants/asst_123", "/v1/assistants/{assistant_id}"),
# Files API
("/v1/files/file_123", "/v1/files/{file_id}"),
("/v1/files/file_123/content", "/v1/files/{file_id}/content"),
# Batches API
("/v1/batches/batch_123", "/v1/batches/{batch_id}"),
("/v1/batches/batch_123/cancel", "/v1/batches/{batch_id}/cancel"),
]
for original, expected in test_cases:
normalized = normalize_request_route(original)
assert normalized == expected, \
f"Failed: {original} -> {normalized} (expected {expected})"
assert (
normalized == expected
), f"Failed: {original} -> {normalized} (expected {expected})"
print(f"✅ {len(test_cases)} other API routes normalized correctly")
@ -195,26 +242,29 @@ def test_prometheus_metrics_use_normalized_routes():
# Create a mock PrometheusLogger
prometheus_logger = MagicMock()
prometheus_logger.get_labels_for_metric = PrometheusLogger.get_labels_for_metric.__get__(prometheus_logger)
prometheus_logger.get_labels_for_metric = (
PrometheusLogger.get_labels_for_metric.__get__(prometheus_logger)
)
# Test with a normalized route
enum_values = UserAPIKeyLabelValues(
route="/v1/responses/{response_id}", # Normalized route
status_code="200",
requested_model="gpt-4",
)
labels = prometheus_label_factory(
supported_enum_labels=prometheus_logger.get_labels_for_metric(
metric_name="litellm_proxy_total_requests_metric"
),
enum_values=enum_values,
)
# Verify the route is normalized in labels
assert labels["route"] == "/v1/responses/{response_id}", \
f"Expected normalized route in labels, got: {labels.get('route')}"
assert (
labels["route"] == "/v1/responses/{response_id}"
), f"Expected normalized route in labels, got: {labels.get('route')}"
print("✅ Prometheus metrics use normalized routes in labels")
@ -227,4 +277,4 @@ if __name__ == "__main__":
test_route_normalization_preserves_static_routes()
test_route_normalization_other_dynamic_apis()
test_prometheus_metrics_use_normalized_routes()
print("\n✅ All prometheus label tests passed!")
print("\n✅ All prometheus label tests passed!")

View file

@ -0,0 +1,77 @@
"""
Unit tests for the new Prometheus metrics that were previously missing from validation.
Tests for:
- litellm_remaining_api_key_requests_for_model
- litellm_remaining_api_key_tokens_for_model
- litellm_callback_logging_failures_metric
"""
from typing import get_args
from litellm.types.integrations.prometheus import (
DEFINED_PROMETHEUS_METRICS,
PrometheusMetricLabels,
UserAPIKeyLabelNames,
)
def test_new_metrics_in_defined_metrics():
"""
Test that the new metrics are present in DEFINED_PROMETHEUS_METRICS.
"""
defined_metrics = get_args(DEFINED_PROMETHEUS_METRICS)
new_metrics = [
"litellm_remaining_api_key_requests_for_model",
"litellm_remaining_api_key_tokens_for_model",
"litellm_callback_logging_failures_metric",
]
for metric in new_metrics:
assert (
metric in defined_metrics
), f"{metric} should be in DEFINED_PROMETHEUS_METRICS"
def test_new_metrics_have_correct_labels():
"""
Test that the new metrics have the correct labels defined.
"""
# Test API Key limits metrics labels
api_key_metrics = [
"litellm_remaining_api_key_requests_for_model",
"litellm_remaining_api_key_tokens_for_model",
]
expected_api_key_labels = [
UserAPIKeyLabelNames.API_KEY_HASH.value,
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
]
for metric in api_key_metrics:
labels = PrometheusMetricLabels.get_labels(metric)
for expected_label in expected_api_key_labels:
assert (
expected_label in labels
), f"{metric} should have label {expected_label}"
# Test Callback failure metric labels
callback_metric = "litellm_callback_logging_failures_metric"
callback_labels = PrometheusMetricLabels.get_labels(callback_metric)
assert (
UserAPIKeyLabelNames.CALLBACK_NAME.value in callback_labels
), f"{callback_metric} should have label {UserAPIKeyLabelNames.CALLBACK_NAME.value}"
def test_callback_name_label_definition():
"""
Test that CALLBACK_NAME is defined correctly in UserAPIKeyLabelNames.
"""
assert UserAPIKeyLabelNames.CALLBACK_NAME.value == "callback_name"
if __name__ == "__main__":
test_new_metrics_in_defined_metrics()
test_new_metrics_have_correct_labels()
test_callback_name_label_definition()

View file

@ -8,6 +8,7 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
_get_customer_id_from_standard_headers,
get_end_user_id_from_request_body,
get_model_from_request,
get_key_model_rpm_limit,
get_key_model_tpm_limit,
)
@ -186,3 +187,25 @@ class TestGetEndUserIdFromRequestBodyWithStandardHeaders:
request_body=request_body, request_headers=headers
)
assert result == "body-user"
def test_get_model_from_request_supports_google_model_names_with_slashes():
assert (
get_model_from_request(
request_data={},
route="/v1beta/models/bedrock/claude-sonnet-3.7:generateContent",
)
== "bedrock/claude-sonnet-3.7"
)
assert (
get_model_from_request(
request_data={},
route="/models/hosted_vllm/gpt-oss-20b:generateContent",
)
== "hosted_vllm/gpt-oss-20b"
)
def test_get_model_from_request_vertex_passthrough_still_works():
route = "/vertex_ai/v1/projects/p/locations/l/publishers/google/models/gemini-1.5-pro:generateContent"
assert get_model_from_request(request_data={}, route=route) == "gemini-1.5-pro"

View file

@ -161,9 +161,11 @@ def test_virtual_key_llm_api_route_includes_passthrough_prefix(route):
[
"/v1beta/models/gemini-2.5-flash:countTokens",
"/v1beta/models/gemini-2.0-flash:generateContent",
"/v1beta/models/bedrock/claude-sonnet-3.7:generateContent",
"/v1beta/models/gemini-1.5-pro:streamGenerateContent",
"/models/gemini-2.5-flash:countTokens",
"/models/gemini-2.0-flash:generateContent",
"/models/bedrock/claude-sonnet-3.7:generateContent",
"/models/gemini-1.5-pro:streamGenerateContent",
],
)
@ -187,9 +189,11 @@ def test_virtual_key_llm_api_routes_allows_google_routes(route):
"/v1beta/models/google-gemini-2-5-pro-code-reviewer-k8s:generateContent",
"/v1beta/models/gemini-2.5-flash-exp:countTokens",
"/v1beta/models/custom-model-name-123:streamGenerateContent",
"/v1beta/models/bedrock/claude-sonnet-3.7:generateContent",
"/models/google-gemini-2-5-pro-code-reviewer-k8s:generateContent",
"/models/gemini-2.5-flash-exp:countTokens",
"/models/custom-model-name-123:streamGenerateContent",
"/models/bedrock/claude-sonnet-3.7:generateContent",
],
)
def test_google_routes_with_dynamic_model_names_recognized_as_llm_api_route(route):

View file

@ -3,6 +3,7 @@ import sys
import uuid
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import HTTPException
from httpx import Request, Response
@ -47,20 +48,129 @@ def test_onyx_guard_config():
del os.environ["ONYX_API_KEY"]
def test_onyx_guard_with_custom_timeout_from_kwargs():
"""Test Onyx guard instantiation with custom timeout passed via kwargs."""
# Set environment variables for testing
os.environ["ONYX_API_BASE"] = "https://test.onyx.security"
os.environ["ONYX_API_KEY"] = "test-api-key"
with patch(
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
) as mock_get_client:
mock_get_client.return_value = MagicMock()
# Simulate how guardrail is instantiated from config with timeout
guardrail = OnyxGuardrail(
guardrail_name="onyx-guard-custom-timeout",
event_hook="pre_call",
default_on=True,
timeout=45.0,
)
# Verify the client was initialized with custom timeout
mock_get_client.assert_called()
call_kwargs = mock_get_client.call_args.kwargs
timeout_param = call_kwargs["params"]["timeout"]
assert timeout_param.read == 45.0
assert timeout_param.connect == 5.0
# Clean up
if "ONYX_API_BASE" in os.environ:
del os.environ["ONYX_API_BASE"]
if "ONYX_API_KEY" in os.environ:
del os.environ["ONYX_API_KEY"]
def test_onyx_guard_with_timeout_none_uses_env_var():
"""Test Onyx guard with timeout=None uses ONYX_TIMEOUT env var.
When timeout=None is passed (as it would be from config model with default None),
the ONYX_TIMEOUT environment variable should be used.
"""
# Set environment variables for testing
os.environ["ONYX_API_BASE"] = "https://test.onyx.security"
os.environ["ONYX_API_KEY"] = "test-api-key"
os.environ["ONYX_TIMEOUT"] = "60"
with patch(
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
) as mock_get_client:
mock_get_client.return_value = MagicMock()
# Pass timeout=None to simulate config model behavior
guardrail = OnyxGuardrail(
guardrail_name="onyx-guard-env-timeout",
event_hook="pre_call",
default_on=True,
timeout=None, # This triggers env var lookup
)
# Verify the client was initialized with timeout from env var
mock_get_client.assert_called()
call_kwargs = mock_get_client.call_args.kwargs
timeout_param = call_kwargs["params"]["timeout"]
assert timeout_param.read == 60.0
assert timeout_param.connect == 5.0
# Clean up
if "ONYX_API_BASE" in os.environ:
del os.environ["ONYX_API_BASE"]
if "ONYX_API_KEY" in os.environ:
del os.environ["ONYX_API_KEY"]
if "ONYX_TIMEOUT" in os.environ:
del os.environ["ONYX_TIMEOUT"]
def test_onyx_guard_with_timeout_none_defaults_to_10():
"""Test Onyx guard with timeout=None and no env var defaults to 10 seconds."""
# Set environment variables for testing
os.environ["ONYX_API_BASE"] = "https://test.onyx.security"
os.environ["ONYX_API_KEY"] = "test-api-key"
# Ensure ONYX_TIMEOUT is not set
if "ONYX_TIMEOUT" in os.environ:
del os.environ["ONYX_TIMEOUT"]
with patch(
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
) as mock_get_client:
mock_get_client.return_value = MagicMock()
# Pass timeout=None with no env var - should default to 10.0
guardrail = OnyxGuardrail(
guardrail_name="onyx-guard-default-timeout",
event_hook="pre_call",
default_on=True,
timeout=None,
)
# Verify the client was initialized with default timeout of 10.0
mock_get_client.assert_called()
call_kwargs = mock_get_client.call_args.kwargs
timeout_param = call_kwargs["params"]["timeout"]
assert timeout_param.read == 10.0
assert timeout_param.connect == 5.0
# Clean up
if "ONYX_API_BASE" in os.environ:
del os.environ["ONYX_API_BASE"]
if "ONYX_API_KEY" in os.environ:
del os.environ["ONYX_API_KEY"]
class TestOnyxGuardrail:
"""Test suite for Onyx Security Guardrail integration."""
def setup_method(self):
"""Setup test environment."""
# Clean up any existing environment variables
for key in ["ONYX_API_BASE", "ONYX_API_KEY"]:
for key in ["ONYX_API_BASE", "ONYX_API_KEY", "ONYX_TIMEOUT"]:
if key in os.environ:
del os.environ[key]
def teardown_method(self):
"""Clean up test environment."""
# Clean up any environment variables set during tests
for key in ["ONYX_API_BASE", "ONYX_API_KEY"]:
for key in ["ONYX_API_BASE", "ONYX_API_KEY", "ONYX_TIMEOUT"]:
if key in os.environ:
del os.environ[key]
@ -103,6 +213,95 @@ class TestOnyxGuardrail:
):
OnyxGuardrail(guardrail_name="test-guard", event_hook="pre_call")
def test_initialization_with_default_timeout(self):
"""Test that default timeout is 10.0 seconds."""
os.environ["ONYX_API_KEY"] = "test-api-key"
with patch(
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
) as mock_get_client:
mock_get_client.return_value = MagicMock()
guardrail = OnyxGuardrail(
guardrail_name="test-guard", event_hook="pre_call", default_on=True
)
# Verify the client was initialized with correct timeout
mock_get_client.assert_called_once()
call_kwargs = mock_get_client.call_args.kwargs
timeout_param = call_kwargs["params"]["timeout"]
assert timeout_param.read == 10.0
assert timeout_param.connect == 5.0
def test_initialization_with_custom_timeout_parameter(self):
"""Test initialization with custom timeout parameter."""
os.environ["ONYX_API_KEY"] = "test-api-key"
with patch(
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
) as mock_get_client:
mock_get_client.return_value = MagicMock()
guardrail = OnyxGuardrail(
guardrail_name="test-guard",
event_hook="pre_call",
default_on=True,
timeout=30.0,
)
# Verify the client was initialized with custom timeout
mock_get_client.assert_called_once()
call_kwargs = mock_get_client.call_args.kwargs
timeout_param = call_kwargs["params"]["timeout"]
assert timeout_param.read == 30.0
assert timeout_param.connect == 5.0
def test_initialization_with_timeout_from_env_var(self):
"""Test initialization with timeout from ONYX_TIMEOUT environment variable.
Note: The env var is only used when timeout=None is explicitly passed,
since the default parameter value is 10.0 (not None).
"""
os.environ["ONYX_API_KEY"] = "test-api-key"
os.environ["ONYX_TIMEOUT"] = "25"
with patch(
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
) as mock_get_client:
mock_get_client.return_value = MagicMock()
# Must pass timeout=None explicitly to trigger env var lookup
guardrail = OnyxGuardrail(
guardrail_name="test-guard", event_hook="pre_call", default_on=True, timeout=None
)
# Verify the client was initialized with timeout from env var
mock_get_client.assert_called_once()
call_kwargs = mock_get_client.call_args.kwargs
timeout_param = call_kwargs["params"]["timeout"]
assert timeout_param.read == 25.0
assert timeout_param.connect == 5.0
def test_initialization_timeout_parameter_overrides_env_var(self):
"""Test that timeout parameter overrides ONYX_TIMEOUT environment variable."""
os.environ["ONYX_API_KEY"] = "test-api-key"
os.environ["ONYX_TIMEOUT"] = "25"
with patch(
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
) as mock_get_client:
mock_get_client.return_value = MagicMock()
guardrail = OnyxGuardrail(
guardrail_name="test-guard",
event_hook="pre_call",
default_on=True,
timeout=15.0,
)
# Verify the client was initialized with parameter timeout (not env var)
mock_get_client.assert_called_once()
call_kwargs = mock_get_client.call_args.kwargs
timeout_param = call_kwargs["params"]["timeout"]
assert timeout_param.read == 15.0
assert timeout_param.connect == 5.0
@pytest.mark.asyncio
async def test_apply_guardrail_request_no_violations(self):
"""Test apply_guardrail for request with no violations detected."""
@ -388,6 +587,105 @@ class TestOnyxGuardrail:
assert result == inputs
@pytest.mark.asyncio
async def test_apply_guardrail_timeout_error_handling(self):
"""Test handling of timeout errors in apply_guardrail (graceful degradation)."""
# Set required API key
os.environ["ONYX_API_KEY"] = "test-api-key"
guardrail = OnyxGuardrail(
guardrail_name="test-guard", event_hook="pre_call", default_on=True, timeout=1.0
)
inputs = GenericGuardrailAPIInputs()
request_data = {
"proxy_server_request": {
"messages": [{"role": "user", "content": "Test message"}],
"model": "gpt-3.5-turbo",
}
}
# Test httpx timeout error
with patch.object(
guardrail.async_handler, "post", side_effect=httpx.TimeoutException("Request timed out")
):
# Should return original inputs on timeout (graceful degradation)
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
logging_obj=None,
)
assert result == inputs
@pytest.mark.asyncio
async def test_apply_guardrail_read_timeout_error_handling(self):
"""Test handling of read timeout errors in apply_guardrail."""
# Set required API key
os.environ["ONYX_API_KEY"] = "test-api-key"
guardrail = OnyxGuardrail(
guardrail_name="test-guard", event_hook="pre_call", default_on=True, timeout=5.0
)
inputs = GenericGuardrailAPIInputs()
request_data = {
"proxy_server_request": {
"messages": [{"role": "user", "content": "Test message"}],
"model": "gpt-3.5-turbo",
}
}
# Test httpx ReadTimeout error
with patch.object(
guardrail.async_handler, "post", side_effect=httpx.ReadTimeout("Read timed out")
):
# Should return original inputs on timeout (graceful degradation)
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
logging_obj=None,
)
assert result == inputs
@pytest.mark.asyncio
async def test_apply_guardrail_connect_timeout_error_handling(self):
"""Test handling of connect timeout errors in apply_guardrail."""
# Set required API key
os.environ["ONYX_API_KEY"] = "test-api-key"
guardrail = OnyxGuardrail(
guardrail_name="test-guard", event_hook="pre_call", default_on=True, timeout=5.0
)
inputs = GenericGuardrailAPIInputs()
request_data = {
"proxy_server_request": {
"messages": [{"role": "user", "content": "Test message"}],
"model": "gpt-3.5-turbo",
}
}
# Test httpx ConnectTimeout error
with patch.object(
guardrail.async_handler, "post", side_effect=httpx.ConnectTimeout("Connect timed out")
):
# Should return original inputs on timeout (graceful degradation)
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
logging_obj=None,
)
assert result == inputs
@pytest.mark.asyncio
async def test_apply_guardrail_no_logging_obj(self):
"""Test apply_guardrail without logging object (uses UUID)."""