mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
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:
parent
a7e665620b
commit
bb6bb664b1
3 changed files with 245 additions and 20 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue