Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
62 changes: 42 additions & 20 deletions litellm/integrations/prometheus.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
Dict,
List,
Literal,
Mapping,
Optional,
Sequence,
Tuple,
Expand Down Expand Up @@ -1449,6 +1450,8 @@ def _increment_token_detail_metrics(
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,
Expand All @@ -1458,7 +1461,7 @@ def _increment_token_detail_metrics(
(
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,
Expand Down Expand Up @@ -1597,27 +1600,12 @@ def _increment_cache_metrics(
)

# 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(
Expand All @@ -1639,6 +1627,40 @@ def _increment_cache_metrics(
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,
Expand Down
152 changes: 152 additions & 0 deletions tests/test_litellm/integrations/test_prometheus_cache_metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -258,6 +258,158 @@ def test_provider_cache_read_does_not_fallback_on_explicit_zero(
# 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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -150,6 +150,57 @@ def test_increments_all_present_token_types(self, sample_enum_values):
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 = {
Expand Down
Loading