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
This commit is contained in:
FERNANDO IZAR 2026-07-02 00:50:22 -03:00 • committed by Sameer Kankute
parent 6d796d0f1f
commit c833b0c362
No known key found for this signature in database
2 changed files with 208 additions and 1 deletions

View file

@ -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(

View file

@ -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