fix(prometheus): populate cache write token metrics for OpenAI-style usage (#34803)

litellm_provider_cache_creation_input_tokens_metric only read the
Anthropic-style top-level usage.cache_creation_input_tokens and had no
prompt_tokens_details fallback, unlike its cache-read twin. OpenAI models
that bill prompt cache writes report them only in
prompt_tokens_details.cache_write_tokens, so the counter never fired for
them. Resolve provider cache read/write tokens through a shared helper
that falls back to prompt_tokens_details.cache_write_tokens (canonical)
then cache_creation_tokens when the explicit top-level field is absent,
and give litellm_input_cache_creation_tokens_metric the same fallback for
raw usage dicts that only carry cache_write_tokens
This commit is contained in:
yucheng-berri 2026-07-27 12:28:19 -07:00 committed by GitHub
parent a7e665620b
commit bb6bb664b1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 245 additions and 20 deletions

View file

@ -16,6 +16,7 @@ from typing import (
Dict,
List,
Literal,
Mapping,
Optional,
Sequence,
Tuple,
@ -1449,6 +1450,8 @@ class PrometheusLogger(CustomLogger):
prompt_details = usage_object.get("prompt_tokens_details") or {}
completion_details = usage_object.get("completion_tokens_details") or {}
cache_creation_detail_tokens = PrometheusLogger._resolve_cache_write_tokens(prompt_details)
detail_metrics: List[Tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]] = [
(
self.litellm_input_cached_tokens_metric,
@ -1458,7 +1461,7 @@ class PrometheusLogger(CustomLogger):
(
self.litellm_input_cache_creation_tokens_metric,
"litellm_input_cache_creation_tokens_metric",
(prompt_details.get("cache_creation_tokens") if isinstance(prompt_details, dict) else None),
cache_creation_detail_tokens,
),
(
self.litellm_input_audio_tokens_metric,
@ -1597,27 +1600,12 @@ class PrometheusLogger(CustomLogger):
)
# 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
(
provider_cache_read_tokens,
provider_cache_creation_tokens,
) = PrometheusLogger._resolve_provider_cache_tokens(usage_obj)
if provider_cache_read_tokens > 0:
PrometheusLogger._inc_labeled_counter(
@ -1639,6 +1627,40 @@ class PrometheusLogger(CustomLogger):
amount=float(provider_cache_creation_tokens),
)
@staticmethod
def _resolve_provider_cache_tokens(usage_obj: Mapping[str, object]) -> tuple[int, int]:
# Prefer explicit provider cache fields when available.
_read = usage_obj.get("cache_read_input_tokens")
_write = usage_obj.get("cache_creation_input_tokens")
provider_cache_read_tokens = _read if isinstance(_read, int) else 0
provider_cache_creation_tokens = _write if isinstance(_write, int) else 0
# Fallback to prompt_tokens_details (common normalization point).
# Only fallback when the explicit field is genuinely absent (None).
prompt_details = usage_obj.get("prompt_tokens_details")
if _read is None and isinstance(prompt_details, dict):
cached_tokens = prompt_details.get("cached_tokens")
if isinstance(cached_tokens, int):
provider_cache_read_tokens = cached_tokens
if _write is None:
write_tokens = PrometheusLogger._resolve_cache_write_tokens(prompt_details)
if write_tokens is not None:
provider_cache_creation_tokens = write_tokens
return provider_cache_read_tokens, provider_cache_creation_tokens
@staticmethod
def _resolve_cache_write_tokens(prompt_details: object) -> int | None:
if not isinstance(prompt_details, dict):
return None
for key in ("cache_write_tokens", "cache_creation_tokens"):
value = prompt_details.get(key)
if isinstance(value, int) and not isinstance(value, bool):
return value
return None
def _increment_mcp_tool_call_metrics(
self,
standard_logging_payload: StandardLoggingPayload,

View file

@ -258,6 +258,158 @@ class TestPrometheusCacheMetrics:
# 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_provider_cache_creation_fallback_to_cache_write_tokens(
self, sample_enum_values
):
"""OpenAI-style usage (prompt_tokens_details.cache_write_tokens, no top-level
cache_creation_input_tokens) must populate the provider cache creation metric."""
mock_logger = MagicMock()
from litellm.integrations.prometheus import PrometheusLogger
standard_logging_payload = {
"cache_hit": False,
"total_tokens": 12100,
"prompt_tokens": 12000,
"completion_tokens": 100,
"model_group": "openai",
"request_tags": [],
"metadata": {
"usage_object": {
"prompt_tokens_details": {
"cached_tokens": 0,
"cache_write_tokens": 800,
},
}
},
}
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,
)
mock_logger.litellm_provider_cache_creation_input_tokens_metric.labels().inc.assert_called_once_with(
800
)
def test_provider_cache_creation_fallback_to_cache_creation_tokens(
self, sample_enum_values
):
"""Normalized litellm usage dumps carry cache_creation_tokens in
prompt_tokens_details; the fallback must read it when cache_write_tokens is absent."""
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": {
"prompt_tokens_details": {"cache_creation_tokens": 42},
}
},
}
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,
)
mock_logger.litellm_provider_cache_creation_input_tokens_metric.labels().inc.assert_called_once_with(
42
)
def test_provider_cache_creation_does_not_fallback_on_explicit_zero(
self, sample_enum_values
):
"""Explicit cache_creation_input_tokens=0 must not trigger fallback to
prompt_tokens_details, mirroring the cache-read semantics."""
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_creation_input_tokens": 0,
"prompt_tokens_details": {"cache_write_tokens": 800},
}
},
}
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,
)
mock_logger.litellm_provider_cache_creation_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

View file

@ -150,6 +150,57 @@ class TestIncrementTokenDetailMetrics:
10.0
)
def test_cache_creation_falls_back_to_cache_write_tokens(self, sample_enum_values):
logger = _make_mock_logger()
payload = {
"metadata": {
"usage_object": {
"prompt_tokens": 12000,
"completion_tokens": 100,
"total_tokens": 12100,
"prompt_tokens_details": {
"cached_tokens": 0,
"cache_write_tokens": 800,
},
}
},
}
PrometheusLogger._increment_token_detail_metrics(
logger,
standard_logging_payload=payload,
enum_values=sample_enum_values,
)
logger.litellm_input_cache_creation_tokens_metric.labels().inc.assert_called_once_with(
800.0
)
def test_cache_write_tokens_takes_precedence_over_cache_creation_tokens(
self, sample_enum_values
):
logger = _make_mock_logger()
payload = {
"metadata": {
"usage_object": {
"prompt_tokens_details": {
"cache_creation_tokens": 25,
"cache_write_tokens": 800,
},
}
},
}
PrometheusLogger._increment_token_detail_metrics(
logger,
standard_logging_payload=payload,
enum_values=sample_enum_values,
)
logger.litellm_input_cache_creation_tokens_metric.labels().inc.assert_called_once_with(
800.0
)
def test_skips_metrics_when_value_is_zero(self, sample_enum_values):
logger = _make_mock_logger()
payload = {