Merge pull request #25934 from BerriAI/litellm_prometheus_improvements

[WIP][Perf] Litellm prometheus improvements
This commit is contained in:
ishaan-berri
2026-04-17 13:28:12 -07:00
committed by GitHub
5 changed files with 450 additions and 278 deletions
+205 -150
View File
@@ -1,6 +1,8 @@
# used for /metrics endpoint on LiteLLM Proxy
#### What this does ####
# On success, log events to Prometheus
from __future__ import annotations
import asyncio
import os
import sys
@@ -14,6 +16,7 @@ from typing import (
List,
Literal,
Optional,
Sequence,
Tuple,
Union,
cast,
@@ -22,6 +25,10 @@ from typing import (
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.prometheus_helpers import (
PrometheusLabelFactoryContext,
_get_cached_end_user_id_for_cost_tracking,
)
from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
get_metadata_variable_name_from_kwargs,
@@ -44,25 +51,6 @@ if TYPE_CHECKING:
else:
AsyncIOScheduler = Any
# Cached lazy import for get_end_user_id_for_cost_tracking
# Module-level cache to avoid repeated imports while preserving memory benefits
_get_end_user_id_for_cost_tracking = None
def _get_cached_end_user_id_for_cost_tracking():
"""
Get cached get_end_user_id_for_cost_tracking function.
Lazy imports on first call to avoid loading utils.py at import time (60MB saved).
Subsequent calls use cached function for better performance.
"""
global _get_end_user_id_for_cost_tracking
if _get_end_user_id_for_cost_tracking is None:
from litellm.utils import get_end_user_id_for_cost_tracking
_get_end_user_id_for_cost_tracking = get_end_user_id_for_cost_tracking
return _get_end_user_id_for_cost_tracking
class PrometheusLogger(CustomLogger):
# Class variables or attributes
@@ -573,7 +561,6 @@ class PrometheusLogger(CustomLogger):
self.enabled_metrics = set()
for group_config in config:
# Validate configuration using Pydantic
if isinstance(group_config, dict):
parsed_config = PrometheusMetricsConfig(**group_config)
else:
@@ -993,12 +980,28 @@ class PrometheusLogger(CustomLogger):
return filtered_labels
def _inc_labeled_counter(
self,
counter: Any,
metric_name: DEFINED_PROMETHEUS_METRICS,
enum_values: UserAPIKeyLabelValues,
label_context: Optional[PrometheusLabelFactoryContext] = None,
amount: float = 1.0,
) -> None:
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name=metric_name
),
enum_values=enum_values,
label_context=label_context,
)
counter.labels(**_labels).inc(amount)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
# Define prometheus client
from litellm.types.utils import StandardLoggingPayload
verbose_logger.debug(
f"prometheus Logging - Enters success logging function for kwargs {kwargs}"
"prometheus Logging - Enters success logging function (kwargs keys: %s)",
list(kwargs.keys()) if isinstance(kwargs, dict) else type(kwargs).__name__,
)
# unpack kwargs
@@ -1111,6 +1114,8 @@ class PrometheusLogger(CustomLogger):
user_api_key = hash_token(user_api_key)
label_context = PrometheusLabelFactoryContext(enum_values) #amortized per request.
# increment total LLM requests and spend metric
self._increment_top_level_request_and_spend_metrics(
end_user_id=end_user_id,
@@ -1122,6 +1127,7 @@ class PrometheusLogger(CustomLogger):
user_id=user_id,
response_cost=response_cost,
enum_values=enum_values,
label_context=label_context,
)
# input, output, total token metrics
@@ -1138,6 +1144,7 @@ class PrometheusLogger(CustomLogger):
user_api_team_alias=user_api_team_alias,
user_id=user_id,
enum_values=enum_values,
label_context=label_context,
)
# remaining budget metrics
@@ -1173,29 +1180,36 @@ class PrometheusLogger(CustomLogger):
# 1. We just checked if isinstance(standard_logging_payload, dict). Pyright complains.
# 2. Pyright does not allow us to run isinstance(standard_logging_payload, StandardLoggingPayload) <- this would be ideal
enum_values=enum_values,
label_context=label_context,
)
# set x-ratelimit headers
self.set_llm_deployment_success_metrics(
kwargs, start_time, end_time, enum_values, output_tokens
kwargs,
start_time,
end_time,
enum_values,
output_tokens,
label_context=label_context,
)
# cache metrics
self._increment_cache_metrics(
standard_logging_payload=standard_logging_payload, # type: ignore
enum_values=enum_values,
label_context=label_context,
)
# increment litellm_proxy_total_requests_metric for all successful requests
# (both streaming and non-streaming) in this single location to prevent
# double-counting that occurs when async_post_call_success_hook also increments
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_proxy_total_requests_metric"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_proxy_total_requests_metric,
"litellm_proxy_total_requests_metric",
enum_values,
label_context=label_context,
)
self.litellm_proxy_total_requests_metric.labels(**_labels).inc()
def _increment_token_metrics(
self,
@@ -1208,6 +1222,7 @@ class PrometheusLogger(CustomLogger):
user_api_team_alias: Optional[str],
user_id: Optional[str],
enum_values: UserAPIKeyLabelValues,
label_context: Optional[PrometheusLabelFactoryContext] = None,
):
verbose_logger.debug("prometheus Logging - Enters token metrics function")
# token metrics
@@ -1217,41 +1232,36 @@ class PrometheusLogger(CustomLogger):
):
_tags = standard_logging_payload["request_tags"]
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_total_tokens_metric"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_tokens_metric,
"litellm_total_tokens_metric",
enum_values,
label_context=label_context,
amount=float(standard_logging_payload["total_tokens"]),
)
self.litellm_tokens_metric.labels(**_labels).inc(
standard_logging_payload["total_tokens"]
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_input_tokens_metric,
"litellm_input_tokens_metric",
enum_values,
label_context=label_context,
amount=float(standard_logging_payload["prompt_tokens"]),
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_input_tokens_metric"
),
enum_values=enum_values,
)
self.litellm_input_tokens_metric.labels(**_labels).inc(
standard_logging_payload["prompt_tokens"]
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_output_tokens_metric"
),
enum_values=enum_values,
)
self.litellm_output_tokens_metric.labels(**_labels).inc(
standard_logging_payload["completion_tokens"]
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_output_tokens_metric,
"litellm_output_tokens_metric",
enum_values,
label_context=label_context,
amount=float(standard_logging_payload["completion_tokens"]),
)
def _increment_cache_metrics(
self,
standard_logging_payload: StandardLoggingPayload,
enum_values: UserAPIKeyLabelValues,
label_context: Optional[PrometheusLabelFactoryContext] = None,
):
"""
Increment cache-related Prometheus metrics based on cache hit/miss status.
@@ -1268,33 +1278,34 @@ class PrometheusLogger(CustomLogger):
if cache_hit is True:
# Increment cache hits counter
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_cache_hits_metric"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_cache_hits_metric,
"litellm_cache_hits_metric",
enum_values,
label_context=label_context,
)
self.litellm_cache_hits_metric.labels(**_labels).inc()
# Increment cached tokens counter
total_tokens = standard_logging_payload.get("total_tokens", 0)
if total_tokens > 0:
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_cached_tokens_metric"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_cached_tokens_metric,
"litellm_cached_tokens_metric",
enum_values,
label_context=label_context,
amount=float(total_tokens),
)
self.litellm_cached_tokens_metric.labels(**_labels).inc(total_tokens)
else:
# cache_hit is False - increment cache misses counter
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_cache_misses_metric"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_cache_misses_metric,
"litellm_cache_misses_metric",
enum_values,
label_context=label_context,
)
self.litellm_cache_misses_metric.labels(**_labels).inc()
async def _increment_remaining_budget_metrics(
self,
@@ -1361,25 +1372,24 @@ class PrometheusLogger(CustomLogger):
user_id: Optional[str],
response_cost: float,
enum_values: UserAPIKeyLabelValues,
label_context: Optional[PrometheusLabelFactoryContext] = None,
):
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_requests_metric"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_requests_metric,
"litellm_requests_metric",
enum_values,
label_context=label_context,
)
self.litellm_requests_metric.labels(**_labels).inc()
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_spend_metric"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_spend_metric,
"litellm_spend_metric",
enum_values,
label_context=label_context,
amount=float(response_cost),
)
self.litellm_spend_metric.labels(**_labels).inc(response_cost)
def _set_virtual_key_rate_limit_metrics(
self,
user_api_key: Optional[str],
@@ -1430,6 +1440,7 @@ class PrometheusLogger(CustomLogger):
user_api_team: Optional[str],
user_api_team_alias: Optional[str],
enum_values: UserAPIKeyLabelValues,
label_context: Optional[PrometheusLabelFactoryContext] = None,
):
# latency metrics
end_time: datetime = kwargs.get("end_time") or datetime.now()
@@ -1449,6 +1460,7 @@ class PrometheusLogger(CustomLogger):
metric_name="litellm_llm_api_time_to_first_token_metric"
),
enum_values=enum_values,
label_context=label_context,
)
self.litellm_llm_api_time_to_first_token_metric.labels(
**_ttft_labels
@@ -1468,6 +1480,7 @@ class PrometheusLogger(CustomLogger):
metric_name="litellm_llm_api_latency_metric"
),
enum_values=enum_values,
label_context=label_context,
)
self.litellm_llm_api_latency_metric.labels(**_labels).observe(
api_call_total_time_seconds
@@ -1484,6 +1497,7 @@ class PrometheusLogger(CustomLogger):
metric_name="litellm_request_total_latency_metric"
),
enum_values=enum_values,
label_context=label_context,
)
self.litellm_request_total_latency_metric.labels(**_labels).observe(
total_time_seconds
@@ -1500,16 +1514,16 @@ class PrometheusLogger(CustomLogger):
metric_name="litellm_request_queue_time_seconds"
),
enum_values=enum_values,
label_context=label_context,
)
self.litellm_request_queue_time_metric.labels(**_labels).observe(
queue_time_seconds
)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
from litellm.types.utils import StandardLoggingPayload
verbose_logger.debug(
f"prometheus Logging - Enters failure logging function for kwargs {kwargs}"
"prometheus Logging - Enters failure logging function (kwargs keys: %s)",
list(kwargs.keys()) if isinstance(kwargs, dict) else type(kwargs).__name__,
)
standard_logging_payload: StandardLoggingPayload = kwargs.get(
@@ -1771,21 +1785,21 @@ class PrometheusLogger(CustomLogger):
if litellm.prometheus_emit_stream_label
else None,
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_proxy_failed_requests_metric"
),
enum_values=enum_values,
_label_ctx = PrometheusLabelFactoryContext(enum_values)
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_proxy_failed_requests_metric,
"litellm_proxy_failed_requests_metric",
enum_values,
label_context=_label_ctx,
)
self.litellm_proxy_failed_requests_metric.labels(**_labels).inc()
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_proxy_total_requests_metric"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_proxy_total_requests_metric,
"litellm_proxy_total_requests_metric",
enum_values,
label_context=_label_ctx,
)
self.litellm_proxy_total_requests_metric.labels(**_labels).inc()
except Exception as e:
verbose_logger.exception(
@@ -2015,22 +2029,23 @@ class PrometheusLogger(CustomLogger):
api_base=api_base,
api_provider=llm_provider or "",
)
_deployment_label_ctx = PrometheusLabelFactoryContext(enum_values)
if exception is not None:
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_deployment_failure_responses"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_deployment_failure_responses,
"litellm_deployment_failure_responses",
enum_values,
label_context=_deployment_label_ctx,
)
self.litellm_deployment_failure_responses.labels(**_labels).inc()
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_deployment_total_requests"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_deployment_total_requests,
"litellm_deployment_total_requests",
enum_values,
label_context=_deployment_label_ctx,
)
self.litellm_deployment_total_requests.labels(**_labels).inc()
pass
except Exception as e:
@@ -2090,6 +2105,7 @@ class PrometheusLogger(CustomLogger):
end_time,
enum_values: UserAPIKeyLabelValues,
output_tokens: float = 1.0,
label_context: Optional[PrometheusLabelFactoryContext] = None,
):
try:
verbose_logger.debug("setting remaining tokens requests metric")
@@ -2147,6 +2163,7 @@ class PrometheusLogger(CustomLogger):
metric_name="litellm_overhead_latency_metric"
),
enum_values=enum_values,
label_context=label_context,
)
self.litellm_overhead_latency_metric.labels(**_labels).observe(
litellm_overhead_time_ms / 1000
@@ -2164,6 +2181,7 @@ class PrometheusLogger(CustomLogger):
metric_name="litellm_remaining_requests_metric"
),
enum_values=enum_values,
label_context=label_context,
)
self.litellm_remaining_requests_metric.labels(**_labels).set(
remaining_requests
@@ -2175,6 +2193,7 @@ class PrometheusLogger(CustomLogger):
metric_name="litellm_remaining_tokens_metric"
),
enum_values=enum_values,
label_context=label_context,
)
self.litellm_remaining_tokens_metric.labels(**_labels).set(
remaining_tokens
@@ -2191,21 +2210,20 @@ class PrometheusLogger(CustomLogger):
api_provider=llm_provider or "",
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_deployment_success_responses"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_deployment_success_responses,
"litellm_deployment_success_responses",
enum_values,
label_context=label_context,
)
self.litellm_deployment_success_responses.labels(**_labels).inc()
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_deployment_total_requests"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_deployment_total_requests,
"litellm_deployment_total_requests",
enum_values,
label_context=label_context,
)
self.litellm_deployment_total_requests.labels(**_labels).inc()
# Track deployment Latency
response_ms: timedelta = end_time - start_time
@@ -2235,6 +2253,7 @@ class PrometheusLogger(CustomLogger):
metric_name="litellm_deployment_latency_per_output_token"
),
enum_values=enum_values,
label_context=label_context,
)
self.litellm_deployment_latency_per_output_token.labels(
**_labels
@@ -2468,13 +2487,13 @@ class PrometheusLogger(CustomLogger):
exception_class=self._get_exception_class_name(original_exception),
tags=_tags,
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_deployment_successful_fallbacks"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_deployment_successful_fallbacks,
"litellm_deployment_successful_fallbacks",
enum_values,
label_context=PrometheusLabelFactoryContext(enum_values),
)
self.litellm_deployment_successful_fallbacks.labels(**_labels).inc()
async def log_failure_fallback_event(
self, original_model_group: str, kwargs: dict, original_exception: Exception
@@ -2514,13 +2533,13 @@ class PrometheusLogger(CustomLogger):
tags=_tags,
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_deployment_failed_fallbacks"
),
enum_values=enum_values,
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_deployment_failed_fallbacks,
"litellm_deployment_failed_fallbacks",
enum_values,
label_context=PrometheusLabelFactoryContext(enum_values),
)
self.litellm_deployment_failed_fallbacks.labels(**_labels).inc()
def set_litellm_deployment_state(
self,
@@ -2638,7 +2657,7 @@ class PrometheusLogger(CustomLogger):
self,
data_fetch_function: Callable[..., Awaitable[Tuple[List[Any], Optional[int]]]],
set_metrics_function: Callable[[List[Any]], Awaitable[None]],
data_type: Literal["teams", "keys", "users"],
data_type: Literal["teams", "keys", "users", "orgs"],
):
"""
Generic method to initialize budget metrics for teams or API keys.
@@ -2714,8 +2733,6 @@ class PrometheusLogger(CustomLogger):
"""
Initialize API key budget metrics by reusing the generic pagination logic.
"""
from typing import Union
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
from litellm.proxy.management_endpoints.key_management_endpoints import (
_list_key_helper,
@@ -2762,7 +2779,6 @@ class PrometheusLogger(CustomLogger):
"""
Initialize user budget metrics by reusing the generic pagination logic.
"""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
@@ -3403,7 +3419,6 @@ class PrometheusLogger(CustomLogger):
It emits the current remaining budget metrics for all Keys and Teams.
"""
from litellm.constants import PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES
from litellm.integrations.custom_logger import CustomLogger
prometheus_loggers: List[
CustomLogger
@@ -3458,16 +3473,56 @@ class PrometheusLogger(CustomLogger):
)
def _prometheus_labels_from_context(
supported_enum_labels: List[str],
ctx: PrometheusLabelFactoryContext,
) -> Dict[str, Optional[str]]:
filtered_labels: Dict[str, Optional[str]] = {
label: ctx._sanitized_enum[label]
for label in supported_enum_labels
if label in ctx._sanitized_enum
}
if UserAPIKeyLabelNames.END_USER.value in filtered_labels:
filtered_labels[UserAPIKeyLabelNames.END_USER.value] = ctx.get_resolved_end_user()
for sk, val in ctx._custom_by_sanitized_key.items():
if sk in supported_enum_labels:
filtered_labels[sk] = val
for k, v in ctx._tag_labels.items():
if k in supported_enum_labels:
filtered_labels[k] = v
for label in supported_enum_labels:
if label not in filtered_labels:
filtered_labels[label] = None
return filtered_labels
def prometheus_label_factory(
supported_enum_labels: List[str],
enum_values: UserAPIKeyLabelValues,
tag: Optional[str] = None,
*,
label_context: Optional[PrometheusLabelFactoryContext] = None,
) -> dict:
"""
Returns a dictionary of label + values for prometheus.
Ensures end_user param is not sent to prometheus if it is not supported.
When ``label_context`` is provided, it must have been built from the same
``enum_values`` object; work is amortized (single model_dump, tag map, etc.).
"""
if label_context is not None:
if label_context.enum_values is not enum_values:
raise ValueError(
"label_context.enum_values must be the same object as enum_values"
)
return _prometheus_labels_from_context(supported_enum_labels, label_context)
# Extract dictionary from Pydantic object
enum_dict = enum_values.model_dump()
@@ -3541,7 +3596,7 @@ def get_custom_labels_from_metadata(metadata: dict) -> Dict[str, str]:
def _tag_matches_wildcard_configured_pattern(
tags: List[str], configured_tag: str
tags: Sequence[str], configured_tag: str
) -> bool:
"""
Check if any of the request tags matches a wildcard configured pattern
@@ -3573,7 +3628,7 @@ def _tag_matches_wildcard_configured_pattern(
return any(re.match(pattern=regex_pattern, string=tag) for tag in tags)
def get_custom_labels_from_tags(tags: List[str]) -> Dict[str, str]:
def get_custom_labels_from_tags(tags: Sequence[str]) -> Dict[str, str]:
"""
Get custom labels from tags based on admin configuration.
@@ -0,0 +1,81 @@
"""
Helpers for the Prometheus integration (extracted to keep ``prometheus.py`` smaller).
``PrometheusLabelFactoryContext`` lives here so it has a dedicated module.
"""
from __future__ import annotations
from typing import Any, Dict, Optional, cast
from litellm.types.integrations.prometheus import (
UserAPIKeyLabelValues,
_sanitize_prometheus_label_name,
_sanitize_prometheus_label_value,
)
_get_end_user_id_for_cost_tracking = None
def _get_cached_end_user_id_for_cost_tracking():
"""
Get cached get_end_user_id_for_cost_tracking function.
Lazy imports on first call to avoid loading utils.py at import time (60MB saved).
Subsequent calls use cached function for better performance.
"""
global _get_end_user_id_for_cost_tracking
if _get_end_user_id_for_cost_tracking is None:
from litellm.utils import get_end_user_id_for_cost_tracking
_get_end_user_id_for_cost_tracking = get_end_user_id_for_cost_tracking
return _get_end_user_id_for_cost_tracking
class PrometheusLabelFactoryContext:
"""
Precomputes per-request label inputs so prometheus_label_factory can subset
per metric without repeated model_dump / tag / metadata work.
"""
__slots__ = (
"enum_values",
"_sanitized_enum",
"_custom_by_sanitized_key",
"_tag_labels",
"_resolved_end_user",
)
_END_USER_NOT_COMPUTED = object()
def __init__(self, enum_values: UserAPIKeyLabelValues) -> None:
self.enum_values = enum_values
enum_dict = enum_values.model_dump()
self._sanitized_enum: Dict[str, Optional[str]] = {
k: _sanitize_prometheus_label_value(v)
for k, v in enum_dict.items()
}
self._custom_by_sanitized_key: Dict[str, Optional[str]] = {}
if enum_values.custom_metadata_labels is not None:
for key, value in enum_values.custom_metadata_labels.items():
sk = _sanitize_prometheus_label_name(key)
self._custom_by_sanitized_key[sk] = _sanitize_prometheus_label_value(
value
)
self._tag_labels: Dict[str, Optional[str]] = {}
if enum_values.tags is not None:
# Late import avoids circular import: ``prometheus`` imports this module.
from litellm.integrations.prometheus import get_custom_labels_from_tags
for k, v in get_custom_labels_from_tags(enum_values.tags).items():
self._tag_labels[k] = _sanitize_prometheus_label_value(v)
# Use a dedicated sentinel so `None` can be cached as a computed result.
self._resolved_end_user: Any = self._END_USER_NOT_COMPUTED
def get_resolved_end_user(self) -> Optional[str]:
if self._resolved_end_user is self._END_USER_NOT_COMPUTED:
fn = _get_cached_end_user_id_for_cost_tracking()
self._resolved_end_user = fn(
litellm_params={"user_api_key_end_user_id": self.enum_values.end_user},
service_type="prometheus",
)
return cast(Optional[str], self._resolved_end_user)
+125 -125
View File
@@ -1,10 +1,8 @@
import re
from dataclasses import dataclass
from dataclasses import MISSING, dataclass, field, fields
from enum import Enum
from typing import Any, ClassVar, Dict, List, Literal, Optional, Tuple
from pydantic import BaseModel, Field, field_validator
from typing_extensions import Annotated
from types import MappingProxyType
from typing import Any, ClassVar, Dict, List, Literal, Mapping, Optional, Tuple, Union
import litellm
@@ -41,35 +39,34 @@ def _sanitize_prometheus_label_name(label: str) -> str:
return sanitized
# v1: single translate pass + escape loop (avoids chained str.replace allocations).
_PROMETHEUS_LABEL_VALUE_TRANSLATE_V1 = str.maketrans("\n", " ", "\r\u2028\u2029")
def _sanitize_prometheus_label_value(value: Optional[Any]) -> Optional[str]:
"""
Sanitize a label value for Prometheus text format compatibility.
Removes or replaces characters that break the Prometheus exposition format:
- U+2028 (Line Separator) and U+2029 (Paragraph Separator) are removed
- Carriage returns are removed
- Newlines are replaced with spaces
- Backslashes and double quotes are escaped per Prometheus spec
Same semantics as :func:`_sanitize_prometheus_label_value`, implemented with
``str.translate`` plus a single escape pass instead of chained ``replace``.
"""
if value is None:
return None
# Coerce non-string values (int, bool, etc.) to str before sanitizing
str_value: str = value if isinstance(value, str) else str(value)
# Remove Unicode line/paragraph separators that break text format
str_value = str_value.replace("\u2028", "").replace("\u2029", "")
cleaned = str_value.translate(_PROMETHEUS_LABEL_VALUE_TRANSLATE_V1)
if "\\" not in cleaned and '"' not in cleaned:
return cleaned
# Remove carriage returns
str_value = str_value.replace("\r", "")
# Replace newlines with spaces
str_value = str_value.replace("\n", " ")
# Escape backslashes and double quotes per Prometheus exposition format
str_value = str_value.replace("\\", "\\\\").replace('"', '\\"')
return str_value
parts: List[str] = []
append = parts.append
for ch in cleaned:
if ch == "\\":
append("\\\\")
elif ch == '"':
append('\\"')
else:
append(ch)
return "".join(parts)
@dataclass
@@ -744,112 +741,115 @@ class PrometheusMetricLabels:
return default_labels + custom_labels
class UserAPIKeyLabelValues(BaseModel):
end_user: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.END_USER.value)
] = None
user: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.USER.value)
] = None
user_email: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.USER_EMAIL.value)
] = None
hashed_api_key: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.API_KEY_HASH.value)
] = None
api_key_alias: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.API_KEY_ALIAS.value)
] = None
team: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.TEAM.value)
] = None
team_alias: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.TEAM_ALIAS.value)
] = None
model_group: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.MODEL_GROUP.value)
] = None
requested_model: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.REQUESTED_MODEL.value)
] = None
model: Annotated[
Optional[str],
Field(..., alias=UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value),
] = None
litellm_model_name: Annotated[
Optional[str],
Field(..., alias=UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value),
] = None
tags: List[str] = []
custom_metadata_labels: Dict[str, str] = {}
model_id: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.MODEL_ID.value)
] = None
api_base: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.API_BASE.value)
] = None
api_provider: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.API_PROVIDER.value)
] = None
exception_status: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.EXCEPTION_STATUS.value)
] = None
exception_class: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.EXCEPTION_CLASS.value)
] = None
status_code: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.STATUS_CODE.value)
] = None
fallback_model: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.FALLBACK_MODEL.value)
] = None
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
stream: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.STREAM.value)
] = None
org_id: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.ORG_ID.value)
] = None
org_alias: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.ORG_ALIAS.value)
] = None
@field_validator("stream", mode="before")
@classmethod
def coerce_stream_to_str(cls, v: Any) -> Optional[str]:
if v is None:
return None
return str(v)
_USER_API_KEY_LABEL_VALUE_INIT_ALIASES: Dict[str, str] = {
# Some tests / call sites use ``api_key_hash``; Prometheus field is ``hashed_api_key``.
"api_key_hash": "hashed_api_key",
}
class PrometheusMetricsConfig(BaseModel):
"""Configuration for filtering Prometheus metrics"""
@dataclass(frozen=True, init=False)
class UserAPIKeyLabelValues:
"""
Prometheus metric label inputs (Python field names match historical Pydantic ``model_dump`` keys).
group: str = Field(..., description="Group name for this set of metrics")
metrics: List[str] = Field(
..., description="List of metric names to include in this group"
)
include_labels: Optional[List[str]] = Field(
None,
description="List of labels to include for these metrics. If None, includes all default labels.",
)
Immutable value object: use ``dataclasses.replace()`` to derive a new instance.
``model_dump()`` is provided for call sites that still expect a Pydantic-like dict.
"""
end_user: Optional[str] = None
user: Optional[str] = None
user_email: Optional[str] = None
hashed_api_key: Optional[str] = None
api_key_alias: Optional[str] = None
team: Optional[str] = None
team_alias: Optional[str] = None
model_group: Optional[str] = None
requested_model: Optional[str] = None
model: Optional[str] = None
litellm_model_name: Optional[str] = None
# Accept list/tuple at construction time; normalize to tuple in __post_init__.
tags: Union[Tuple[str, ...], List[str]] = ()
custom_metadata_labels: Mapping[str, str] = field(default_factory=dict)
model_id: Optional[str] = None
api_base: Optional[str] = None
api_provider: Optional[str] = None
exception_status: Optional[str] = None
exception_class: Optional[str] = None
status_code: Optional[str] = None
fallback_model: Optional[str] = None
route: Optional[str] = None
client_ip: Optional[str] = None
user_agent: Optional[str] = None
stream: Optional[str] = None
org_id: Optional[str] = None
org_alias: Optional[str] = None
#Added for test compatibility.
def __init__(self, **kwargs: Any) -> None:
"""
Match former Pydantic behavior: unknown keys are ignored; ``api_key_hash`` maps to
``hashed_api_key``. This supports ``**standard_logging_payload`` in tests.
"""
field_names = {f.name for f in fields(self)}
merged: Dict[str, Any] = {}
for f in fields(self):
if f.default_factory is not MISSING:
merged[f.name] = f.default_factory()
else:
merged[f.name] = f.default
for k, v in kwargs.items():
if k in field_names:
merged[k] = v
continue
canon = _USER_API_KEY_LABEL_VALUE_INIT_ALIASES.get(k)
if canon is not None and canon in field_names:
merged[canon] = v
for f in fields(self):
object.__setattr__(self, f.name, merged[f.name])
self.__post_init__()
def __post_init__(self) -> None:
object.__setattr__(self, "tags", tuple(self.tags))
if self.stream is not None:
object.__setattr__(self, "stream", str(self.stream))
_cmd = dict(self.custom_metadata_labels)
object.__setattr__(
self,
"custom_metadata_labels",
MappingProxyType(_cmd),
)
def __repr__(self) -> str:
# Perf: this object is constructed on every Prometheus logging path; verbose
# dataclass/Pydantic-style repr is expensive and often pulled in accidentally
# via f-strings / debug logging. Return empty so accidental stringification
# stays cheap. (Dataclass default `str()` delegates to `__repr__`.)
return ""
def model_dump(self) -> Dict[str, Any]:
"""Same shape as the former Pydantic ``model_dump()`` (plain dict, list tags)."""
d: Dict[str, Any] = {f.name: getattr(self, f.name) for f in fields(self)}
d["tags"] = list(self.tags)
d["custom_metadata_labels"] = dict(self.custom_metadata_labels)
return d
class PrometheusSettings(BaseModel):
"""Settings for Prometheus metrics configuration"""
@dataclass
class PrometheusMetricsConfig:
"""Configuration for filtering Prometheus metrics (parsed once from proxy config)."""
prometheus_metrics_config: Optional[List[PrometheusMetricsConfig]] = Field(
None,
description="Configuration for filtering Prometheus metrics by groups and labels",
)
group: str
metrics: List[str]
include_labels: Optional[List[str]] = None
@dataclass
class PrometheusSettings:
"""Settings for Prometheus metrics configuration."""
prometheus_metrics_config: Optional[List[PrometheusMetricsConfig]] = None
class NoOpMetric:
@@ -660,7 +660,7 @@ async def test_async_log_failure_event(prometheus_logger):
)
# litellm_llm_api_failed_requests_metric incremented
# Labels: end_user, api_key_hash, api_key_alias, model, team, team_alias, user, model_id
# Labels: end_user, hashed_api_key, api_key_alias, model, team, team_alias, user, model_id
prometheus_logger.litellm_llm_api_failed_requests_metric.labels.assert_called_once_with(
None, # end_user_id
"test_hash",
@@ -1150,10 +1150,10 @@ def test_prometheus_factory(monkeypatch, enable_end_user_cost_tracking_prometheu
enum_values = UserAPIKeyLabelValues(
end_user="test_end_user",
api_key_hash="test_hash",
hashed_api_key="test_hash",
api_key_alias="test_alias",
)
supported_labels = ["end_user", "api_key_hash", "api_key_alias"]
supported_labels = ["end_user", "hashed_api_key", "api_key_alias"]
returned_dict = prometheus_label_factory(
supported_enum_labels=supported_labels, enum_values=enum_values
)
@@ -1162,6 +1162,8 @@ def test_prometheus_factory(monkeypatch, enable_end_user_cost_tracking_prometheu
assert returned_dict["end_user"] == "test_end_user"
else:
assert returned_dict["end_user"] == None
assert returned_dict["hashed_api_key"] == "test_hash"
assert returned_dict["api_key_alias"] == "test_alias"
def test_get_custom_labels_from_metadata(monkeypatch):
@@ -0,0 +1,34 @@
import pytest
from litellm.types.integrations.prometheus import (
_sanitize_prometheus_label_value,
)
@pytest.mark.parametrize(
"value,expected",
[
(None, None),
("", ""),
("plain", "plain"),
# Newlines -> spaces, carriage returns removed
("a\nb", "a b"),
("a\rb", "ab"),
("a\r\nb", "a b"),
# Unicode line/paragraph separators removed
("a\u2028b", "ab"),
("a\u2029b", "ab"),
("a\u2028b\u2029c", "abc"),
# Escapes per Prometheus text format
('he said "hi"', 'he said \\"hi\\"'),
(r"path\to\file", r"path\\to\\file"),
(r'quote\"slash\\', r'quote\\\"slash\\\\'),
# Non-string inputs get coerced to str first
(123, "123"),
(True, "True"),
(False, "False"),
],
)
def test_sanitize_prometheus_label_value_expected_outputs(value, expected):
assert _sanitize_prometheus_label_value(value) == expected