From a80945756ad1c6c8f95d7e9a346d1c3da822e962 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Mon, 7 Sep 2026 23:24:59 -0700 Subject: [PATCH] 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. --- basedpyright-code-budget.json | 2 +- litellm/integrations/prometheus.py | 70 +++-- ruff-strict-budget.json | 6 +- ..._prometheus_input_sequence_length_label.py | 250 +++++++++++++++--- type-discipline-budget.json | 2 +- 5 files changed, 272 insertions(+), 58 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 931ebc46f51..0b0a61192e6 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -117,7 +117,7 @@ "limit": 110 }, "reportUnnecessaryComparison": { - "limit": 686 + "limit": 687 }, "reportUnnecessaryContains": { "limit": 4 diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index d8a7027f62f..540ce6738fc 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -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( diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 3ec8eea9807..fd7b30bc314 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -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 diff --git a/tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py b/tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py index a3527cf1016..bc922061544 100644 --- a/tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py +++ b/tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py @@ -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" diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 970dc38072a..e7186dfe186 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16424 + "limit": 16426 }, "LIT011": { "limit": 5506