From c833b0c362caf9bcfbe9bac4362e16ab5a0d8ecf Mon Sep 17 00:00:00 2001 From: FERNANDO IZAR Date: Thu, 2 Jul 2026 00:50:22 -0300 Subject: [PATCH] fix(prometheus): bound per-request budget metric emission with a timeout (#31632) * fix(prometheus): bound per-request budget metric emission with a timeout Wrap the per-request budget-metric gather in asyncio.wait_for so a slow Redis or DB lookup cannot consume the whole LoggingWorker watchdog and get the success-logging event cancelled. On timeout the emission is skipped in isolation; budget gauges are still refreshed by the periodic cron. The timeout is configurable via PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT and defaults to 5.0 seconds, falling back to the default on an invalid value instead of raising * fix(prometheus): reject non-finite and non-positive budget-metrics timeout env float() accepts 0, negatives, nan and inf, which bypass the fallback: a value <= 0 makes asyncio.wait_for time out immediately and skip every per-request emission, and inf reintroduces the unbounded wait the timeout was meant to bound. Validate the parsed value is finite and greater than zero before using it, otherwise fall back to the default --- litellm/integrations/prometheus.py | 41 ++++- .../test_prometheus_budget_metrics_timeout.py | 168 ++++++++++++++++++ 2 files changed, 208 insertions(+), 1 deletion(-) create mode 100644 tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 4ebe312e301..0f89e21abda 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -4,6 +4,7 @@ from __future__ import annotations import asyncio +import math import os import sys from datetime import datetime, timedelta @@ -65,6 +66,26 @@ if TYPE_CHECKING: else: AsyncIOScheduler = Any +_DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT = 5.0 + + +def _get_budget_metrics_per_request_timeout() -> float: + raw = os.getenv("PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT") + if raw is None: + return _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT + try: + parsed = float(raw) + except ValueError: + parsed = None + if parsed is None or not math.isfinite(parsed) or parsed <= 0: + verbose_logger.debug( + "[Non-Blocking] Prometheus: invalid PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT=%r; using default %ss.", + raw, + _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT, + ) + return _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT + return parsed + class PrometheusLogger(CustomLogger): # Class variables or attributes @@ -1607,7 +1628,15 @@ class PrometheusLogger(CustomLogger): _user_spend = _metadata.get("user_api_key_user_spend", None) _user_max_budget = _metadata.get("user_api_key_user_max_budget", None) - results = await asyncio.gather( + # Bound the per-request budget-metric emission so that slow Redis/DB + # lookups under load cannot consume the whole LoggingWorker watchdog + # (LOGGING_WORKER_MAX_TIME_PER_COROUTINE, default 20s) and get the entire + # success-logging event cancelled. Budget gauges are also refreshed by the + # periodic cron every PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES, + # so dropping one slow per-request emission only loses sub-cron real-time + # detail, not correctness. + budget_metrics_timeout = _get_budget_metrics_per_request_timeout() + gather_coro = asyncio.gather( self._set_api_key_budget_metrics_after_api_request( user_api_key=user_api_key, user_api_key_alias=user_api_key_alias, @@ -1634,6 +1663,16 @@ class PrometheusLogger(CustomLogger): ), return_exceptions=True, ) + try: + results = await asyncio.wait_for(gather_coro, timeout=budget_metrics_timeout) + except asyncio.TimeoutError: + verbose_logger.debug( + "[Non-Blocking] Prometheus: per-request budget metric emission " + "exceeded %ss under load; skipping (values are refreshed by the " + "periodic budget-metrics cron job).", + budget_metrics_timeout, + ) + return for i, r in enumerate(results): if isinstance(r, Exception): verbose_logger.debug( diff --git a/tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py b/tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py new file mode 100644 index 00000000000..a4d245e9dc0 --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py @@ -0,0 +1,168 @@ +""" +Unit tests for the per-request budget-metric emission timeout in +PrometheusLogger._increment_remaining_budget_metrics. + +A slow Redis/DB lookup in one of the budget branches must not let the gather run +unbounded; it is wrapped in asyncio.wait_for so the success-logging coroutine +cannot exceed the LoggingWorker watchdog and get the whole event cancelled. +""" + +import asyncio +from unittest.mock import AsyncMock, patch + +import pytest +from prometheus_client import REGISTRY + +from litellm.integrations.prometheus import ( + PrometheusLogger, + _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT, + _get_budget_metrics_per_request_timeout, +) + +TIMEOUT_ENV = "PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT" + + +@pytest.fixture(autouse=True) +def cleanup_prometheus_registry(): + collectors = list(REGISTRY._collector_to_names.keys()) + for collector in collectors: + try: + REGISTRY.unregister(collector) + except Exception: + pass + yield + collectors = list(REGISTRY._collector_to_names.keys()) + for collector in collectors: + try: + REGISTRY.unregister(collector) + except Exception: + pass + + +@pytest.fixture +def prometheus_logger(): + return PrometheusLogger() + + +def _call_increment(logger: PrometheusLogger): + return logger._increment_remaining_budget_metrics( + user_api_team="team-1", + user_api_team_alias="team-alias", + user_api_key="key-1", + user_api_key_alias="key-alias", + litellm_params={"metadata": {}}, + response_cost=0.01, + user_id="user-1", + user_api_key_org_id="org-1", + ) + + +def _skip_logged(debug_mock) -> bool: + return any("skipping" in str(call.args[0]) for call in debug_mock.call_args_list if call.args) + + +@pytest.mark.asyncio +async def test_budget_metric_emission_skips_on_timeout(prometheus_logger, monkeypatch): + """A branch slower than the timeout is skipped without propagating, and the + skip is logged instead of cancelling the success-logging event.""" + monkeypatch.setenv(TIMEOUT_ENV, "0.05") + + async def _slow_branch(**kwargs): + await asyncio.sleep(30) + + prometheus_logger._set_api_key_budget_metrics_after_api_request = _slow_branch + prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock() + + with patch("litellm.integrations.prometheus.verbose_logger") as mock_logger: + await _call_increment(prometheus_logger) + + assert _skip_logged(mock_logger.debug) + + +@pytest.mark.asyncio +async def test_budget_metric_emission_completes_within_timeout(prometheus_logger, monkeypatch): + """With a generous timeout every branch is awaited and no skip is logged.""" + monkeypatch.setenv(TIMEOUT_ENV, "5.0") + + prometheus_logger._set_api_key_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock() + + with patch("litellm.integrations.prometheus.verbose_logger") as mock_logger: + await _call_increment(prometheus_logger) + + assert prometheus_logger._set_api_key_budget_metrics_after_api_request.await_count == 1 + assert prometheus_logger._set_team_budget_metrics_after_api_request.await_count == 1 + assert prometheus_logger._set_user_budget_metrics_after_api_request.await_count == 1 + assert prometheus_logger._set_org_budget_metrics_after_api_request.await_count == 1 + assert not _skip_logged(mock_logger.debug) + + +@pytest.mark.asyncio +async def test_invalid_timeout_env_falls_back_to_default(prometheus_logger, monkeypatch): + """A malformed timeout env value must not raise (which would recreate the + failure mode); it falls back to the default and every branch still runs.""" + monkeypatch.setenv(TIMEOUT_ENV, "not-a-number") + + prometheus_logger._set_api_key_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock() + + await _call_increment(prometheus_logger) + + assert prometheus_logger._set_api_key_budget_metrics_after_api_request.await_count == 1 + assert prometheus_logger._set_org_budget_metrics_after_api_request.await_count == 1 + + +@pytest.mark.parametrize("value", ["not-a-number", "0", "-1", "nan", "inf", "-inf"]) +def test_unusable_timeout_env_falls_back_to_default(value, monkeypatch): + """Values that parse but disable or unbound the timeout (0, negative, nan, + inf) must fall back to the default instead of being used; otherwise they + either skip every emission or recreate the unbounded-wait failure mode.""" + monkeypatch.setenv(TIMEOUT_ENV, value) + + assert _get_budget_metrics_per_request_timeout() == _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT + + +@pytest.mark.parametrize("value,expected", [("0.05", 0.05), ("5.0", 5.0), ("30", 30.0)]) +def test_valid_timeout_env_is_used(value, expected, monkeypatch): + """A finite positive value is parsed and returned unchanged.""" + monkeypatch.setenv(TIMEOUT_ENV, value) + + assert _get_budget_metrics_per_request_timeout() == expected + + +def test_missing_timeout_env_uses_default(monkeypatch): + """With the env unset the default is returned.""" + monkeypatch.delenv(TIMEOUT_ENV, raising=False) + + assert _get_budget_metrics_per_request_timeout() == _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT + + +@pytest.mark.asyncio +async def test_outer_cancellation_still_propagates(prometheus_logger, monkeypatch): + """Only asyncio.TimeoutError is swallowed; an outer cancellation (cooperative + shutdown / watchdog) injected while awaiting must still propagate.""" + monkeypatch.setenv(TIMEOUT_ENV, "30") + + started = asyncio.Event() + + async def _slow_branch(**kwargs): + started.set() + await asyncio.sleep(30) + + prometheus_logger._set_api_key_budget_metrics_after_api_request = _slow_branch + prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock() + prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock() + + task = asyncio.create_task(_call_increment(prometheus_logger)) + await started.wait() + task.cancel() + + with pytest.raises(asyncio.CancelledError): + await task