From b99ce6e1bfa95b0d2f82f6a9a8870b00ff577ec9 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Mon, 7 Sep 2026 10:43:50 -0700 Subject: [PATCH] fix(prometheus): preserve disabled custom input length labels --- litellm/integrations/prometheus.py | 7 +++- ..._prometheus_input_sequence_length_label.py | 42 ++++++++++++++++++- 2 files changed, 46 insertions(+), 3 deletions(-) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index d7ed16bba9a..f2b19815563 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -246,6 +246,7 @@ class PrometheusLogger(CustomLogger): # logger so toggling these flags only takes effect after a # restart, keeping init-time and runtime label sets in sync. self._cached_metric_labels: dict[str, list[str]] = {} + self._emit_input_sequence_length_label = litellm.prometheus_emit_input_sequence_length_label is True _custom_buckets: Final = litellm.prometheus_latency_buckets self.latency_buckets = tuple(_custom_buckets) if _custom_buckets is not None else LATENCY_BUCKETS @@ -1443,7 +1444,11 @@ 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")), + input_sequence_length=( + get_input_sequence_length_bucket(standard_logging_payload.get("prompt_tokens")) + 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-"): 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 b53466685f8..e6f386abc03 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 @@ -128,11 +128,17 @@ def _standard_logging_payload(now: datetime.datetime, prompt_tokens: int) -> Sta ) -def _success_kwargs(now: datetime.datetime, prompt_tokens: int) -> Mapping[str, object]: +def _success_kwargs( + now: datetime.datetime, prompt_tokens: int, requester_metadata: Mapping[str, str] | None = None +) -> Mapping[str, object]: + payload: Final = _standard_logging_payload(now, prompt_tokens) return { "model": "gpt-4o-mini", "litellm_params": {"metadata": {}}, - "standard_logging_object": _standard_logging_payload(now, prompt_tokens), + "standard_logging_object": { + **payload, + "metadata": {**payload["metadata"], "requester_metadata": requester_metadata}, + }, "stream": True, "start_time": now - datetime.timedelta(seconds=3), "api_call_start_time": now - datetime.timedelta(seconds=2), @@ -169,3 +175,35 @@ async def test_logger_built_with_flag_off_emits_no_bucket_label(monkeypatch: pyt samples: Final = _latency_bucket_samples() assert samples assert all("input_sequence_length" not in sample.labels for sample in samples) + + +@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 +): + now: Final = datetime.datetime.now() + monkeypatch.setattr(litellm, "custom_prometheus_metadata_labels", ["input_sequence_length"]) + 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, + ) + + 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 all(logger.get_labels_for_metric(metric).count("input_sequence_length") == 1 for metric in LATENCY_METRICS)