mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat: add clientip and user agent in metrics (#19717)
* feat: add clientip and user agent in metrics * fix: lint errors * Add model id and other req labels --------- Co-authored-by: Krish Dholakia <krrishdholakia@gmail.com>
This commit is contained in:
parent
a65ac2af13
commit
344ea3d9f2
7 changed files with 492 additions and 110 deletions
|
|
@ -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"
|
||||
),
|
||||
)
|
||||
|
||||
########################################
|
||||
|
|
@ -373,15 +377,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 +952,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 +1011,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 +1246,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 +1268,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 +1367,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 +1398,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 +1416,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 +1474,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 +1503,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 +1516,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 +1536,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 +1587,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 +1605,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 +1647,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 +1663,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 +1705,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 +1737,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"),
|
||||
)
|
||||
|
||||
"""
|
||||
|
|
@ -2263,7 +2288,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 +2407,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)}"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
@ -4453,6 +4458,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,
|
||||
|
|
@ -5138,6 +5144,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")
|
||||
),
|
||||
|
|
@ -5203,6 +5210,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,
|
||||
|
|
|
|||
|
|
@ -846,7 +846,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
|
||||
|
|
@ -1002,7 +1004,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
|
||||
|
|
@ -1029,8 +1033,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 (
|
||||
|
|
@ -1050,6 +1054,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,
|
||||
|
|
@ -1532,7 +1546,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
|
||||
|
|
@ -1550,13 +1566,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
|
||||
|
|
@ -1567,7 +1587,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:
|
||||
|
|
@ -1578,7 +1600,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
|
||||
|
|
|
|||
|
|
@ -150,6 +150,8 @@ 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"
|
||||
|
||||
|
||||
|
|
@ -199,6 +201,7 @@ DEFINED_PROMETHEUS_METRICS = Literal[
|
|||
"litellm_cached_tokens_metric",
|
||||
"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",
|
||||
]
|
||||
|
||||
|
|
@ -213,6 +216,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 = [
|
||||
|
|
@ -221,6 +225,7 @@ class PrometheusMetricLabels:
|
|||
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
|
||||
UserAPIKeyLabelNames.TEAM.value,
|
||||
UserAPIKeyLabelNames.TEAM_ALIAS.value,
|
||||
UserAPIKeyLabelNames.MODEL_ID.value,
|
||||
]
|
||||
|
||||
litellm_request_total_latency_metric = [
|
||||
|
|
@ -232,6 +237,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 = [
|
||||
|
|
@ -243,6 +249,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)
|
||||
|
|
@ -262,6 +269,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 = [
|
||||
|
|
@ -276,6 +286,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 = [
|
||||
|
|
@ -296,6 +309,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 = [
|
||||
|
|
@ -305,6 +319,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 = [
|
||||
|
|
@ -314,6 +329,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 = [
|
||||
|
|
@ -325,6 +341,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 = [
|
||||
|
|
@ -336,6 +355,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 = [
|
||||
|
|
@ -348,6 +370,7 @@ class PrometheusMetricLabels:
|
|||
UserAPIKeyLabelNames.USER.value,
|
||||
UserAPIKeyLabelNames.USER_EMAIL.value,
|
||||
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
|
||||
UserAPIKeyLabelNames.MODEL_ID.value,
|
||||
]
|
||||
|
||||
litellm_total_tokens_metric = [
|
||||
|
|
@ -360,6 +383,7 @@ class PrometheusMetricLabels:
|
|||
UserAPIKeyLabelNames.USER.value,
|
||||
UserAPIKeyLabelNames.USER_EMAIL.value,
|
||||
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
|
||||
UserAPIKeyLabelNames.MODEL_ID.value,
|
||||
]
|
||||
|
||||
litellm_output_tokens_metric = [
|
||||
|
|
@ -372,6 +396,7 @@ class PrometheusMetricLabels:
|
|||
UserAPIKeyLabelNames.USER.value,
|
||||
UserAPIKeyLabelNames.USER_EMAIL.value,
|
||||
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
|
||||
UserAPIKeyLabelNames.MODEL_ID.value,
|
||||
]
|
||||
|
||||
litellm_deployment_state = [
|
||||
|
|
@ -398,6 +423,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
|
||||
|
|
@ -473,6 +499,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 = [
|
||||
|
|
@ -485,10 +513,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] = []
|
||||
|
||||
|
|
@ -509,6 +564,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
|
||||
|
|
@ -601,6 +657,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):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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!")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue