feat(prometheus): surface provider cached metrics which are independent of LiteLLM cache. (#27660)

This commit is contained in:
Paulo Edgar Castro 2026-06-01 11:21:01 +01:00 • committed by GitHub
parent 2fe6e6c45e
commit 3526d14be9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 182 additions and 5 deletions

View file

@ -511,6 +511,23 @@ class PrometheusLogger(CustomLogger):
labelnames=self.get_labels_for_metric("litellm_cached_tokens_metric"),
)
# Provider prompt-caching metrics
self.litellm_provider_cache_read_input_tokens_metric = self._counter_factory(
name="litellm_provider_cache_read_input_tokens_metric",
documentation="Total prompt/input tokens read from provider prompt cache (e.g. OpenAI/Anthropic/Gemini/Bedrock)",
labelnames=self.get_labels_for_metric(
"litellm_provider_cache_read_input_tokens_metric"
),
)
self.litellm_provider_cache_creation_input_tokens_metric = self._counter_factory(
name="litellm_provider_cache_creation_input_tokens_metric",
documentation="Total prompt/input tokens written to provider prompt cache (e.g. Anthropic/Bedrock)",
labelnames=self.get_labels_for_metric(
"litellm_provider_cache_creation_input_tokens_metric"
),
)
# User and Team count metrics
self.litellm_total_users_metric = self._gauge_factory(
"litellm_total_users",
@ -1458,11 +1475,11 @@ class PrometheusLogger(CustomLogger):
"""
cache_hit = standard_logging_payload.get("cache_hit")
# Only track if cache_hit has a definite value (True or False)
if cache_hit is None:
return
if cache_hit is True:
# Historically these metrics only tracked LiteLLM caching.
# Provider prompt-caching metrics are still emitted below.
pass
elif cache_hit is True:
# Increment cache hits counter
PrometheusLogger._inc_labeled_counter(
self,
@ -1493,6 +1510,51 @@ class PrometheusLogger(CustomLogger):
label_context=label_context,
)
# Provider prompt caching metrics are independent of LiteLLM cache_hit.
provider_cache_read_tokens = 0
provider_cache_creation_tokens = 0
usage_obj = (standard_logging_payload.get("metadata", {}) or {}).get(
"usage_object"
)
if isinstance(usage_obj, dict):
# Prefer explicit provider cache fields when available.
_read = usage_obj.get("cache_read_input_tokens")
_write = usage_obj.get("cache_creation_input_tokens")
if isinstance(_read, int):
provider_cache_read_tokens = _read
if isinstance(_write, int):
provider_cache_creation_tokens = _write
# Fallback to prompt_tokens_details.cached_tokens (common normalization point).
# Only fallback when the explicit field is genuinely absent (None).
if _read is None:
prompt_details = usage_obj.get("prompt_tokens_details")
if isinstance(prompt_details, dict):
cached_tokens = prompt_details.get("cached_tokens")
if isinstance(cached_tokens, int):
provider_cache_read_tokens = cached_tokens
if provider_cache_read_tokens > 0:
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_provider_cache_read_input_tokens_metric,
"litellm_provider_cache_read_input_tokens_metric",
enum_values,
label_context=label_context,
amount=float(provider_cache_read_tokens),
)
if provider_cache_creation_tokens > 0:
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_provider_cache_creation_input_tokens_metric,
"litellm_provider_cache_creation_input_tokens_metric",
enum_values,
label_context=label_context,
amount=float(provider_cache_creation_tokens),
)
async def _increment_remaining_budget_metrics(
self,
user_api_team: Optional[str],

View file

@ -238,6 +238,9 @@ DEFINED_PROMETHEUS_METRICS = Literal[
"litellm_cache_hits_metric",
"litellm_cache_misses_metric",
"litellm_cached_tokens_metric",
# Provider prompt-caching metrics (e.g. OpenAI/Anthropic/Bedrock/Gemini)
"litellm_provider_cache_read_input_tokens_metric",
"litellm_provider_cache_creation_input_tokens_metric",
"litellm_deployment_tpm_limit",
"litellm_deployment_rpm_limit",
"litellm_remaining_api_key_requests_for_model",
@ -655,6 +658,10 @@ class PrometheusMetricLabels:
litellm_cache_misses_metric = _cache_metric_labels
litellm_cached_tokens_metric = _cache_metric_labels
# Provider prompt-caching metrics - track tokens read/written to provider caches
litellm_provider_cache_read_input_tokens_metric = _cache_metric_labels
litellm_provider_cache_creation_input_tokens_metric = _cache_metric_labels
# Metrics whose emission paths supply org context (used by get_labels)
_org_label_metrics: ClassVar[frozenset] = frozenset(
{
@ -672,7 +679,6 @@ class PrometheusMetricLabels:
"litellm_output_tokens_metric",
}
)
# Managed batch metrics
_batch_user_labels = [
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,

View file

@ -35,6 +35,8 @@ class TestPrometheusCacheMetrics:
assert "litellm_cache_hits_metric" in defined_metrics
assert "litellm_cache_misses_metric" in defined_metrics
assert "litellm_cached_tokens_metric" in defined_metrics
assert "litellm_provider_cache_read_input_tokens_metric" in defined_metrics
assert "litellm_provider_cache_creation_input_tokens_metric" in defined_metrics
def test_cache_metric_labels_defined(self):
"""Test that cache metric labels are properly defined"""
@ -44,6 +46,13 @@ class TestPrometheusCacheMetrics:
assert hasattr(PrometheusMetricLabels, "litellm_cache_hits_metric")
assert hasattr(PrometheusMetricLabels, "litellm_cache_misses_metric")
assert hasattr(PrometheusMetricLabels, "litellm_cached_tokens_metric")
assert hasattr(
PrometheusMetricLabels, "litellm_provider_cache_read_input_tokens_metric"
)
assert hasattr(
PrometheusMetricLabels,
"litellm_provider_cache_creation_input_tokens_metric",
)
# Verify labels include expected keys
expected_labels = [
@ -59,6 +68,14 @@ class TestPrometheusCacheMetrics:
assert label in PrometheusMetricLabels.litellm_cache_hits_metric
assert label in PrometheusMetricLabels.litellm_cache_misses_metric
assert label in PrometheusMetricLabels.litellm_cached_tokens_metric
assert (
label
in PrometheusMetricLabels.litellm_provider_cache_read_input_tokens_metric
)
assert (
label
in PrometheusMetricLabels.litellm_provider_cache_creation_input_tokens_metric
)
def test_increment_cache_metrics_on_cache_hit(self, sample_enum_values):
"""Test that cache hit increments the correct metrics"""
@ -76,12 +93,20 @@ class TestPrometheusCacheMetrics:
"completion_tokens": 50,
"model_group": "openai",
"request_tags": [],
"metadata": {
"usage_object": {
"cache_read_input_tokens": 25,
"cache_creation_input_tokens": 10,
}
},
}
# Create mock metrics
mock_logger.litellm_cache_hits_metric = MagicMock()
mock_logger.litellm_cache_misses_metric = MagicMock()
mock_logger.litellm_cached_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_read_input_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_creation_input_tokens_metric = MagicMock()
mock_logger.get_labels_for_metric = MagicMock(
return_value=[
"model",
@ -114,6 +139,14 @@ class TestPrometheusCacheMetrics:
# Verify cache misses metric was NOT called
mock_logger.litellm_cache_misses_metric.labels.assert_not_called()
# Verify provider prompt caching metrics were incremented
mock_logger.litellm_provider_cache_read_input_tokens_metric.labels().inc.assert_called_once_with(
25
)
mock_logger.litellm_provider_cache_creation_input_tokens_metric.labels().inc.assert_called_once_with(
10
)
def test_increment_cache_metrics_on_cache_miss(self, sample_enum_values):
"""Test that cache miss increments the correct metrics"""
# Create mock for PrometheusLogger instance
@ -129,12 +162,20 @@ class TestPrometheusCacheMetrics:
"completion_tokens": 50,
"model_group": "openai",
"request_tags": [],
"metadata": {
"usage_object": {
# Explicit provider field absent -> fallback should use prompt_tokens_details.cached_tokens
"prompt_tokens_details": {"cached_tokens": 20},
}
},
}
# Create mock metrics
mock_logger.litellm_cache_hits_metric = MagicMock()
mock_logger.litellm_cache_misses_metric = MagicMock()
mock_logger.litellm_cached_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_read_input_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_creation_input_tokens_metric = MagicMock()
mock_logger.get_labels_for_metric = MagicMock(
return_value=[
"model",
@ -162,6 +203,61 @@ class TestPrometheusCacheMetrics:
mock_logger.litellm_cache_hits_metric.labels.assert_not_called()
mock_logger.litellm_cached_tokens_metric.labels.assert_not_called()
# Provider prompt caching metrics should still be emitted
mock_logger.litellm_provider_cache_read_input_tokens_metric.labels().inc.assert_called_once_with(
20
)
mock_logger.litellm_provider_cache_creation_input_tokens_metric.labels.assert_not_called()
def test_provider_cache_read_does_not_fallback_on_explicit_zero(
self, sample_enum_values
):
"""Explicit cache_read_input_tokens=0 must not trigger fallback to cached_tokens."""
mock_logger = MagicMock()
from litellm.integrations.prometheus import PrometheusLogger
standard_logging_payload = {
"cache_hit": False,
"total_tokens": 100,
"prompt_tokens": 50,
"completion_tokens": 50,
"model_group": "openai",
"request_tags": [],
"metadata": {
"usage_object": {
"cache_read_input_tokens": 0,
"prompt_tokens_details": {"cached_tokens": 20},
}
},
}
mock_logger.litellm_cache_hits_metric = MagicMock()
mock_logger.litellm_cache_misses_metric = MagicMock()
mock_logger.litellm_cached_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_read_input_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_creation_input_tokens_metric = MagicMock()
mock_logger.get_labels_for_metric = MagicMock(
return_value=[
"model",
"hashed_api_key",
"api_key_alias",
"team",
"team_alias",
"end_user",
"user",
]
)
PrometheusLogger._increment_cache_metrics(
mock_logger,
standard_logging_payload=standard_logging_payload,
enum_values=sample_enum_values,
)
# Should not emit read metric, because explicit provider value is zero.
mock_logger.litellm_provider_cache_read_input_tokens_metric.labels.assert_not_called()
def test_increment_cache_metrics_when_cache_hit_is_none(self, sample_enum_values):
"""Test that no metrics are incremented when cache_hit is None"""
# Create mock for PrometheusLogger instance
@ -177,12 +273,19 @@ class TestPrometheusCacheMetrics:
"completion_tokens": 50,
"model_group": "openai",
"request_tags": [],
"metadata": {
"usage_object": {
"cache_read_input_tokens": 25,
}
},
}
# Create mock metrics
mock_logger.litellm_cache_hits_metric = MagicMock()
mock_logger.litellm_cache_misses_metric = MagicMock()
mock_logger.litellm_cached_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_read_input_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_creation_input_tokens_metric = MagicMock()
mock_logger.get_labels_for_metric = MagicMock(
return_value=[
"model",
@ -207,6 +310,12 @@ class TestPrometheusCacheMetrics:
mock_logger.litellm_cache_misses_metric.labels.assert_not_called()
mock_logger.litellm_cached_tokens_metric.labels.assert_not_called()
# Provider prompt caching metrics should still be emitted
mock_logger.litellm_provider_cache_read_input_tokens_metric.labels().inc.assert_called_once_with(
25
)
mock_logger.litellm_provider_cache_creation_input_tokens_metric.labels.assert_not_called()
if __name__ == "__main__":
pytest.main([__file__, "-v"])