From 3526d14be91292fce6ed6ca2ae4ac6df6b549df5 Mon Sep 17 00:00:00 2001 From: Paulo Edgar Castro Date: Mon, 1 Jun 2026 11:21:01 +0100 Subject: [PATCH] feat(prometheus): surface provider cached metrics which are independent of LiteLLM cache. (#27660) --- litellm/integrations/prometheus.py | 70 ++++++++++- litellm/types/integrations/prometheus.py | 8 +- .../test_prometheus_cache_metrics.py | 109 ++++++++++++++++++ 3 files changed, 182 insertions(+), 5 deletions(-) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 5f052842122..9fc09807369 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -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], diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 827d10985cf..55f4fc96504 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -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, diff --git a/tests/test_litellm/integrations/test_prometheus_cache_metrics.py b/tests/test_litellm/integrations/test_prometheus_cache_metrics.py index 88148ce1372..6c9923322fd 100644 --- a/tests/test_litellm/integrations/test_prometheus_cache_metrics.py +++ b/tests/test_litellm/integrations/test_prometheus_cache_metrics.py @@ -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"])