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