fix(prometheus): isolate input buckets and preserve missing usage

Keep built-in buckets on latency histograms, preserve unrelated custom labels, and classify raw incomplete usage and upstream total-only headers as unknown. Cover count conservation, failure callbacks, explicit zero, startup snapshots, and direct caller compatibility. Drop earlier branch budget changes.
This commit is contained in:
Yucheng He 2026-09-07 23:24:59 -07:00
parent 87cce63a43
commit a80945756a
5 changed files with 272 additions and 58 deletions

View file

@ -117,7 +117,7 @@
"limit": 110
},
"reportUnnecessaryComparison": {
"limit": 686
"limit": 687
},
"reportUnnecessaryContains": {
"limit": 4

View file

@ -1417,9 +1417,6 @@ class PrometheusLogger(CustomLogger):
f"inside track_prometheus_metrics, model {model}, response_cost {response_cost}, tokens_used {tokens_used}, end_user_id {end_user_id}, user_api_key {user_api_key}"
)
reported_usage: Final = (
response_obj.get("usage") if isinstance(response_obj, dict) else getattr(response_obj, "usage", None)
)
enum_values: Final = UserAPIKeyLabelValues(
end_user=end_user_id,
hashed_api_key=user_api_key,
@ -1447,17 +1444,6 @@ class PrometheusLogger(CustomLogger):
user_agent=standard_logging_payload["metadata"].get("user_agent"),
stream=(str(standard_logging_payload.get("stream")) if litellm.prometheus_emit_stream_label else None),
service_tier=get_service_tier_from_standard_logging_payload(standard_logging_payload),
input_sequence_length=(
get_input_sequence_length_bucket(
standard_logging_payload.get("prompt_tokens")
if standard_logging_payload.get("prompt_tokens")
or reported_usage is not None
or kwargs.get("combined_usage_object") is not None
else None
)
if self._emit_input_sequence_length_label
else None
),
)
if user_api_key is not None and isinstance(user_api_key, str) and user_api_key.startswith("sk-"):
@ -1537,6 +1523,11 @@ class PrometheusLogger(CustomLogger):
# 2. Pyright does not allow us to run isinstance(standard_logging_payload, StandardLoggingPayload) <- this would be ideal
enum_values=enum_values,
label_context=label_context,
input_sequence_length=(
self._get_input_sequence_length(standard_logging_payload, kwargs, response_obj)
if self._emit_input_sequence_length_label
else None
),
)
# set x-ratelimit headers
@ -2207,6 +2198,36 @@ class PrometheusLogger(CustomLogger):
)
self.litellm_remaining_api_key_tokens_for_model.labels(**tokens_labels).set(remaining_tokens)
@staticmethod
def _get_input_sequence_length(
standard_logging_payload: StandardLoggingPayload,
kwargs: Mapping[str, object],
response_obj: object,
) -> str:
prompt_tokens: Final = standard_logging_payload.get("prompt_tokens")
if prompt_tokens:
return get_input_sequence_length_bucket(prompt_tokens)
combined_usage: Final = kwargs.get("combined_usage_object")
if (
combined_usage is not None
and getattr(kwargs.get("_litellm_upstream_reported_usage"), "total_tokens", None) is not None
):
return get_input_sequence_length_bucket(None)
reported_usage: Final = (
response_obj.get("usage") if isinstance(response_obj, dict) else getattr(response_obj, "usage", None)
)
if reported_usage is None and combined_usage is None:
return get_input_sequence_length_bucket(None)
usage_metadata: Final = standard_logging_payload["metadata"].get("usage_object")
if isinstance(usage_metadata, Mapping):
return get_input_sequence_length_bucket(usage_metadata.get("prompt_tokens"))
if combined_usage is None and isinstance(response_obj, dict):
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
normalized_usage: Final[Mapping[str, object]] = StandardLoggingPayloadSetup.get_usage_as_dict(response_obj)
return get_input_sequence_length_bucket(normalized_usage.get("prompt_tokens"))
return get_input_sequence_length_bucket(prompt_tokens)
def _set_latency_metrics(
self,
kwargs: dict,
@ -2217,7 +2238,16 @@ class PrometheusLogger(CustomLogger):
user_api_team_alias: str | None,
enum_values: UserAPIKeyLabelValues,
label_context: PrometheusLabelFactoryContext | None = None,
input_sequence_length: str | None = None,
):
latency_enum_values: Final = (
replace(enum_values, input_sequence_length=input_sequence_length)
if input_sequence_length is not None
else enum_values
)
latency_label_context: Final = (
PrometheusLabelFactoryContext(latency_enum_values) if input_sequence_length is not None else label_context
)
# latency metrics
end_time: Final[datetime] = kwargs.get("end_time") or datetime.now()
start_time: Final[datetime | None] = kwargs.get("start_time")
@ -2235,8 +2265,8 @@ class PrometheusLogger(CustomLogger):
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_llm_api_time_to_first_token_metric"
),
enum_values=enum_values,
label_context=label_context,
enum_values=latency_enum_values,
label_context=latency_label_context,
)
self.litellm_llm_api_time_to_first_token_metric.labels(**_ttft_labels).observe(time_to_first_token_seconds)
self._track_end_user_metric_series(
@ -2256,8 +2286,8 @@ class PrometheusLogger(CustomLogger):
if api_call_total_time_seconds is not None:
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(metric_name="litellm_llm_api_latency_metric"),
enum_values=enum_values,
label_context=label_context,
enum_values=latency_enum_values,
label_context=latency_label_context,
)
self.litellm_llm_api_latency_metric.labels(**_labels).observe(api_call_total_time_seconds)
self._track_end_user_metric_series(
@ -2287,8 +2317,8 @@ class PrometheusLogger(CustomLogger):
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(metric_name="litellm_request_total_latency_metric"),
enum_values=enum_values,
label_context=label_context,
enum_values=latency_enum_values,
label_context=latency_label_context,
)
self.litellm_request_total_latency_metric.labels(**_labels).observe(_observed_total_time_seconds)
self._track_end_user_metric_series(

View file

@ -21,7 +21,7 @@
"limit": 112
},
"ANN206": {
"limit": 132
"limit": 133
},
"ANN401": {
"limit": 119
@ -57,7 +57,7 @@
"limit": 3
},
"BLE001": {
"limit": 2915
"limit": 2916
},
"C401": {
"limit": 8
@ -156,7 +156,7 @@
"limit": 215
},
"PLW0603": {
"limit": 189
"limit": 190
},
"PLW1508": {
"limit": 190

View file

@ -1,5 +1,7 @@
import asyncio
import datetime
from collections.abc import Mapping
from copy import deepcopy
from typing import Final, cast
import pytest
@ -80,12 +82,30 @@ def test_user_api_key_label_values_carries_input_sequence_length():
assert values.model_dump()["input_sequence_length"] == "4k-16k"
def _latency_bucket_samples() -> tuple[Sample, ...]:
def _assert_latency_metrics(expected: str | None, stream: bool = True, queue_time: float = 0) -> None:
samples: Final = tuple(sample for metric in REGISTRY.collect() for sample in metric.samples)
for metric, duration in zip(LATENCY_METRICS, (2, 1, 3 + queue_time)):
counts: Final = tuple(sample for sample in samples if sample.name == f"{metric}_count")
sums: Final = tuple(sample for sample in samples if sample.name == f"{metric}_sum")
buckets: Final = tuple(sample for sample in samples if sample.name == f"{metric}_bucket")
if not stream and metric == "litellm_llm_api_time_to_first_token_metric":
assert not counts and not sums and not buckets
continue
assert len(counts) == len(sums) == 1
assert counts[0].value == 1
assert sums[0].value == pytest.approx(duration)
assert buckets and any(sample.labels["le"] == "+Inf" for sample in buckets)
assert all(sample.value == int(float(sample.labels["le"]) >= duration) for sample in buckets)
assert all(sample.labels.get("input_sequence_length") == expected for sample in (*counts, *sums, *buckets))
def _non_target_samples() -> tuple[Sample, ...]:
return tuple(
sample
for metric in REGISTRY.collect()
if metric.name not in LATENCY_METRICS
for sample in metric.samples
if sample.name.endswith("_bucket") and any(name in sample.name for name in LATENCY_METRICS)
if "input_sequence_length" in sample.labels and not sample.name.endswith("_created")
)
@ -129,7 +149,7 @@ def _standard_logging_payload(now: datetime.datetime, prompt_tokens: int) -> Sta
def _success_kwargs(
now: datetime.datetime, prompt_tokens: int, requester_metadata: Mapping[str, str] | None = None
now: datetime.datetime, prompt_tokens: int, requester_metadata: Mapping[str, object] | None = None
) -> Mapping[str, object]:
payload: Final = _standard_logging_payload(now, prompt_tokens)
return {
@ -159,9 +179,7 @@ async def test_logger_emits_bucket_from_its_startup_label_set(
await logger.async_log_success_event(dict(_success_kwargs(now, prompt_tokens=4_000)), None, now, now)
samples: Final = _latency_bucket_samples()
assert samples
assert all(sample.labels["input_sequence_length"] == "4k-16k" for sample in samples)
_assert_latency_metrics("4k-16k")
@pytest.mark.asyncio
@ -170,13 +188,24 @@ async def test_logger_emits_bucket_from_its_startup_label_set(
(
({"id": "moderation", "results": []}, None, "unknown"),
({"usage": None}, None, "unknown"),
({"usage": {}}, None, "unknown"),
({"usage": {"completion_tokens": 3}}, None, "unknown"),
({"usage": {"total_tokens": 5}}, None, "unknown"),
({"usage": {"prompt_tokens": 0}}, None, "0-1k"),
({"usage": {"prompt_tokens": 4_000}}, None, "4k-16k"),
({"usage": {"input_tokens": 0, "output_tokens": 3, "total_tokens": 3}}, None, "0-1k"),
({"usage": {"input_tokens": 4_000, "output_tokens": 3, "total_tokens": 4_003}}, None, "4k-16k"),
(litellm.ModelResponse(usage=litellm.Usage(prompt_tokens=0)), None, "0-1k"),
(None, litellm.Usage(prompt_tokens=0), "0-1k"),
),
)
@pytest.mark.parametrize("include_usage_metadata", (True, False))
async def test_logger_distinguishes_missing_usage_from_reported_zero(
monkeypatch: pytest.MonkeyPatch, response: object, combined_usage: object, expected: str
monkeypatch: pytest.MonkeyPatch,
response: object,
combined_usage: object,
expected: str,
include_usage_metadata: bool,
):
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
@ -187,16 +216,72 @@ async def test_logger_distinguishes_missing_usage_from_reported_zero(
response_obj=response if isinstance(response, dict) else None
)
payload: Final = _standard_logging_payload(now, usage.get("prompt_tokens", 0))
await logger.async_log_success_event(
{**_success_kwargs(now, prompt_tokens=usage["prompt_tokens"]), "combined_usage_object": combined_usage},
{
**_success_kwargs(now, prompt_tokens=usage.get("prompt_tokens", 0)),
"combined_usage_object": combined_usage,
"standard_logging_object": {
**payload,
"metadata": {**payload["metadata"], "usage_object": usage if include_usage_metadata else None},
},
},
response,
now,
now,
)
samples: Final = _latency_bucket_samples()
assert samples
assert all(sample.labels["input_sequence_length"] == expected for sample in samples)
_assert_latency_metrics(expected)
@pytest.mark.asyncio
@pytest.mark.parametrize("total_tokens", (None, 0, 5_000))
@pytest.mark.parametrize("prompt_tokens", (0, 4_000))
async def test_upstream_total_only_usage_has_unknown_input_length(
monkeypatch: pytest.MonkeyPatch, total_tokens: int | None, prompt_tokens: int
):
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging, StandardLoggingPayloadSetup
from litellm.proxy.pass_through_endpoints.upstream_usage_headers import apply_upstream_reported_usage
now: Final = datetime.datetime.now()
monkeypatch.setattr(litellm, FLAG, True)
logger: Final = PrometheusLogger()
logging_obj: Final = Logging(
model="gpt-4o-mini",
messages=[],
stream=True,
call_type="pass_through_endpoint",
start_time=now,
litellm_call_id="test-call-id",
function_id="1",
)
headers: Final = httpx.Headers(
{
"x-litellm-response-cost": "0.001",
**({"x-litellm-total-tokens": str(total_tokens)} if total_tokens is not None else {}),
}
)
reported: Final = apply_upstream_reported_usage(logging_obj=logging_obj, headers=headers)
assert reported is not None
combined_usage: Final = logging_obj.model_call_details.get("combined_usage_object")
response: Final = {"usage": {"prompt_tokens": prompt_tokens}}
usage: Final = StandardLoggingPayloadSetup.get_usage_as_dict(response, combined_usage)
payload: Final = _standard_logging_payload(now, usage.get("prompt_tokens", 0))
await logger.async_log_success_event(
{
**logging_obj.model_call_details,
**_success_kwargs(now, usage.get("prompt_tokens", 0)),
"standard_logging_object": {**payload, "metadata": {**payload["metadata"], "usage_object": usage}},
},
response,
now,
now,
)
_assert_latency_metrics("unknown" if total_tokens is not None else get_input_sequence_length_bucket(prompt_tokens))
@pytest.mark.asyncio
@ -207,38 +292,137 @@ async def test_logger_built_with_flag_off_emits_no_bucket_label(monkeypatch: pyt
await logger.async_log_success_event(dict(_success_kwargs(now, prompt_tokens=4_000)), None, now, now)
samples: Final = _latency_bucket_samples()
assert samples
assert all("input_sequence_length" not in sample.labels for sample in samples)
_assert_latency_metrics(None)
@pytest.mark.asyncio
@pytest.mark.parametrize("flag_at_startup", (True, False))
@pytest.mark.parametrize("with_metadata", (True, False))
async def test_custom_input_length_label_preserves_values_when_flag_off(
monkeypatch: pytest.MonkeyPatch, flag_at_startup: bool, with_metadata: bool
@pytest.mark.parametrize("stream", (True, False))
@pytest.mark.parametrize(
"metadata",
(
None,
{},
{"input_sequence_length": None},
{"input_sequence_length": False},
{"input_sequence_length": True},
{"input_sequence_length": 0},
{"input_sequence_length": []},
{"input_sequence_length": {}},
{"input_sequence_length": ""},
{"input_sequence_length": "from-metadata"},
),
)
async def test_custom_input_length_label_is_scoped_to_target_histograms(
monkeypatch: pytest.MonkeyPatch, flag_at_startup: bool, stream: bool, metadata: Mapping[str, object] | None
):
now: Final = datetime.datetime.now()
monkeypatch.setattr(litellm, "custom_prometheus_metadata_labels", ["input_sequence_length"])
kwargs: Final = {
**_success_kwargs(now, prompt_tokens=4_000, requester_metadata=metadata),
"stream": stream,
"litellm_params": {"metadata": {"queue_time_seconds": 0.25}},
}
original_kwargs: Final = deepcopy(kwargs)
baseline_logger: Final = PrometheusLogger()
await baseline_logger.async_log_success_event(kwargs, None, now, now)
baseline_samples: Final = _non_target_samples()
_clear_prometheus_registry()
monkeypatch.setattr(litellm, FLAG, flag_at_startup)
logger: Final = PrometheusLogger()
monkeypatch.setattr(litellm, FLAG, not flag_at_startup)
await logger.async_log_success_event(
dict(
_success_kwargs(
now,
prompt_tokens=4_000,
requester_metadata={"input_sequence_length": "from-metadata"} if with_metadata else None,
)
),
None,
now,
now,
)
await logger.async_log_success_event(kwargs, None, now, now)
samples: Final = _latency_bucket_samples()
expected: Final = "from-metadata" if with_metadata else ("4k-16k" if flag_at_startup else "None")
assert samples
assert all(sample.labels["input_sequence_length"] == expected for sample in samples)
assert kwargs == original_kwargs
custom_value: Final = (metadata or {}).get("input_sequence_length")
expected: Final = custom_value if isinstance(custom_value, str) else ("4k-16k" if flag_at_startup else "None")
_assert_latency_metrics(expected, stream=stream, queue_time=0.25)
assert all(logger.get_labels_for_metric(metric).count("input_sequence_length") == 1 for metric in LATENCY_METRICS)
non_target_samples: Final = _non_target_samples()
assert {
"litellm_requests_metric_total",
"litellm_spend_metric_total",
"litellm_total_tokens_metric_total",
"litellm_request_queue_time_seconds_count",
"litellm_deployment_success_responses_total",
}.issubset({sample.name for sample in non_target_samples})
assert non_target_samples == baseline_samples
queue_sum: Final = tuple(
sample for sample in non_target_samples if sample.name == "litellm_request_queue_time_seconds_sum"
)
assert len(queue_sum) == 1 and queue_sum[0].value == 0.25
@pytest.mark.asyncio
@pytest.mark.parametrize("enabled", (True, False))
async def test_concurrent_requests_keep_independent_buckets(monkeypatch: pytest.MonkeyPatch, enabled: bool):
now: Final = datetime.datetime.now()
monkeypatch.setattr(litellm, FLAG, enabled)
logger: Final = PrometheusLogger()
cases: Final = (
(None, "unknown"),
(0, "0-1k"),
(1_000, "1k-4k"),
(4_000, "4k-16k"),
(16_000, "16k-64k"),
(64_000, "64k+"),
)
calls: Final = tuple(
(
{**_success_kwargs(now, prompt_tokens=tokens or 0), "stream": stream},
{"usage": {"prompt_tokens": tokens}} if tokens is not None else None,
)
for tokens, _ in cases
for stream in (True, False)
for _ in range(2)
)
original_calls: Final = deepcopy(calls)
await asyncio.gather(*(logger.async_log_success_event(kwargs, response, now, now) for kwargs, response in calls))
assert calls == original_calls
samples: Final = tuple(sample for metric in REGISTRY.collect() for sample in metric.samples)
for metric, duration in zip(LATENCY_METRICS, (2, 1, 3)):
expected_count: Final = 2 if metric == "litellm_llm_api_time_to_first_token_metric" else 4
counts: Final = tuple(sample for sample in samples if sample.name == f"{metric}_count")
sums: Final = tuple(sample for sample in samples if sample.name == f"{metric}_sum")
buckets: Final = tuple(sample for sample in samples if sample.name == f"{metric}_bucket")
expected: Final = (
{bucket: expected_count for _, bucket in cases} if enabled else {None: expected_count * len(cases)}
)
assert len(counts) == len(sums) == len(expected)
assert {sample.labels.get("input_sequence_length"): sample.value for sample in counts} == expected
assert {sample.labels.get("input_sequence_length"): sample.value for sample in sums} == {
bucket: count * duration for bucket, count in expected.items()
}
assert sum(sample.value for sample in buckets if sample.labels["le"] == "+Inf") == expected_count * len(cases)
assert all(
sample.value
== expected[sample.labels.get("input_sequence_length")] * int(float(sample.labels["le"]) >= duration)
for sample in buckets
)
@pytest.mark.asyncio
@pytest.mark.parametrize("enabled", (True, False))
async def test_failed_request_does_not_observe_latency(monkeypatch: pytest.MonkeyPatch, enabled: bool):
now: Final = datetime.datetime.now()
monkeypatch.setattr(litellm, FLAG, enabled)
monkeypatch.setattr(litellm, "custom_prometheus_metadata_labels", ["input_sequence_length"])
logger: Final = PrometheusLogger()
kwargs: Final = {
**_success_kwargs(now, prompt_tokens=4_000),
"standard_logging_object": {**_standard_logging_payload(now, 4_000), "status": "failure"},
"exception": RuntimeError("upstream request failed"),
}
await logger.async_log_failure_event(kwargs, None, now, now)
samples: Final = tuple(sample for metric in REGISTRY.collect() for sample in metric.samples)
assert not any(sample.name.startswith(LATENCY_METRICS) for sample in samples)
for metric in ("litellm_llm_api_failed_requests_metric_total", "litellm_deployment_failure_responses_total"):
counts: Final = tuple(sample for sample in samples if sample.name == metric)
assert len(counts) == 1
assert counts[0].value == 1
assert counts[0].labels["input_sequence_length"] == "None"

View file

@ -27,7 +27,7 @@
"limit": 0
},
"LIT010": {
"limit": 16424
"limit": 16426
},
"LIT011": {
"limit": 5506