mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
commit
ea0a264a3c
29 changed files with 2569 additions and 269 deletions
|
|
@ -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
|
||||
```
|
||||
|
|
|
|||
|
|
@ -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]]):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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): "
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
89
tests/proxy_unit_tests/test_get_image.py
Normal file
89
tests/proxy_unit_tests/test_get_image.py
Normal 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
|
||||
|
|
@ -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!")
|
||||
|
|
@ -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!")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue