From e8868d24b1cab68e6c411b983443a085ac459831 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 8 Sep 2026 00:01:47 -0700 Subject: [PATCH] fix(proxy): defer Prometheus alerting until stored settings load Reuse successful startup storage resolution and preserve callback deduplication across alerting reloads. --- litellm/proxy/proxy_server.py | 9 +- .../proxy/proxy_server/test_lifecycle.py | 110 ++++++++++++++++++ .../proxy/proxy_server/test_proxy_config.py | 28 +++++ tests/test_litellm/proxy/test_proxy_server.py | 13 ++- 4 files changed, 157 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8599c6fe1f6..597dd40d1b8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -270,6 +270,7 @@ from litellm.constants import ( from litellm.exceptions import RejectedRequestError from litellm.integrations.custom_guardrail import ModifyResponseException from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.prometheus import PrometheusLogger from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting from litellm.litellm_core_utils.agentic_loop_settings import ( validated_max_agentic_loops, @@ -1289,6 +1290,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: proxy_budget_rescheduler_max_time=proxy_budget_rescheduler_max_time, proxy_batch_write_at=proxy_batch_write_at, proxy_logging_obj=proxy_logging_obj, + resolved_store_model_in_db=should_load_db_litellm_settings, ) if prisma_client is not None else None @@ -6041,6 +6043,9 @@ class ProxyConfig: if _alert == "slack": # [OLD] v0 implementation - already handled by update_values above pass + elif _alert == "prometheus": + if PrometheusLogger.get_instance() is None: + litellm.logging_callback_manager.add_litellm_callback("prometheus") else: # [NEW] v1 implementation - init as a custom logger if _alert in litellm._known_custom_logger_compatible_callbacks: @@ -9462,6 +9467,7 @@ class ProxyStartupEvent: proxy_budget_rescheduler_max_time: int, proxy_batch_write_at: int, proxy_logging_obj: ProxyLogging, + resolved_store_model_in_db: bool = False, ) -> ProxyWorkerHeartbeat: """Initializes scheduled background jobs""" global heuristic_v1_tuning_baselines, store_model_in_db, scheduler # rebind-ok: startup publishes the one read-only baseline snapshot @@ -9595,7 +9601,8 @@ class ProxyStartupEvent: ### ADD NEW MODELS ### store_model_in_db = ( # rebind-ok: preserve legacy YAML values unless env or DB explicitly enables storage - await cls.resolve_store_model_in_db(prisma_client=prisma_client, configured=store_model_in_db) + resolved_store_model_in_db + or await cls.resolve_store_model_in_db(prisma_client=prisma_client, configured=store_model_in_db) or store_model_in_db ) diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index e82047d40d7..52650c8d8e5 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -17,6 +17,7 @@ Pins covered: from __future__ import annotations +import datetime import inspect import json import logging @@ -26,10 +27,15 @@ from typing import Final, Optional, Union from unittest.mock import AsyncMock, MagicMock, call import pytest +from prometheus_client import REGISTRY from pydantic import BaseModel from typing_extensions import TypedDict +import litellm +import litellm.litellm_core_utils.litellm_logging as logging_module import litellm.proxy.proxy_server as ps +import litellm.proxy.utils as proxy_utils +from litellm.integrations.prometheus import PrometheusLogger from litellm.proxy.proxy_server import ( ProxyStartupEvent, _initialize_shared_aiohttp_session, @@ -858,6 +864,110 @@ async def test_startup_logging_applies_db_settings_before_callback_init(monkeypa assert events.method_calls == [call.db_settings(), call.callback_init()] +@pytest.mark.asyncio +@pytest.mark.parametrize(("yaml_value", "db_value"), ((False, True), (True, False))) +@pytest.mark.parametrize("also_callback", (False, True)) +async def test_prometheus_alerting_uses_persisted_settings_without_duplicate_callbacks( + monkeypatch: pytest.MonkeyPatch, yaml_value: bool, db_value: bool, also_callback: bool +): + collectors: Final = tuple(REGISTRY._collector_to_names) # pyright: ignore[reportPrivateUsage] # isolate global registry + for collector in collectors: + REGISTRY.unregister(collector) + monkeypatch.setattr(logging_module, "_in_memory_loggers", []) + monkeypatch.setattr(litellm, "prometheus_emit_input_sequence_length_label", yaml_value) + monkeypatch.setattr(litellm, "callbacks", ["prometheus"] if also_callback else []) + monkeypatch.setattr(proxy_utils, "PROXY_HOOKS", ()) + monkeypatch.setattr(ps, "cost_tracking", MagicMock()) + proxy_logging: Final = ps.ProxyLogging(user_api_key_cache=ps.user_api_key_cache) + proxy_logging.deprecation_check_started = True + monkeypatch.setattr(ps, "proxy_logging_obj", proxy_logging) + config: Final = ps.ProxyConfig() + settings: Final = {"alerting": ["prometheus"], "alert_types": []} + prisma_client: Final = MagicMock() + monkeypatch.setattr(proxy_utils, "litellm_config_cache", proxy_utils.DualCache()) + prisma_client.get_generic_data = AsyncMock( + return_value=MagicMock(param_value={"prometheus_emit_input_sequence_length_label": db_value}) + ) + try: + config._load_alerting_settings(settings) + config._load_alerting_settings(settings) + await ProxyStartupEvent._initialize_startup_logging( + llm_router=None, + proxy_logging_obj=proxy_logging, + redis_usage_cache=None, + prisma_client=prisma_client, + should_load_db_litellm_settings=True, + proxy_config_obj=config, + ) + logger: Final = PrometheusLogger.get_instance() + assert logger is not None + assert litellm.prometheus_emit_input_sequence_length_label is db_value + for metric in ( + "litellm_llm_api_latency_metric", + "litellm_request_total_latency_metric", + "litellm_llm_api_time_to_first_token_metric", + ): + assert ("input_sequence_length" in logger.get_labels_for_metric(metric)) is db_value + + config._load_alerting_settings(settings) + config._load_alerting_settings(settings) + assert "prometheus" not in litellm.callbacks + assert sum(isinstance(callback, PrometheusLogger) for callback in litellm.callbacks) == 1 + assert sum(isinstance(callback, PrometheusLogger) for callback in litellm._async_success_callback) == 1 + now: Final = datetime.datetime.now() + await logger.async_log_success_event( + { + "model": "test-model", + "litellm_params": {"metadata": {}}, + "start_time": now, + "end_time": now, + "standard_logging_object": { + "id": "alerting-startup", + "call_type": "completion", + "status": "success", + "model": "test-model", + "model_group": "test-model", + "model_id": "test-model", + "api_base": "https://api.openai.com", + "custom_llm_provider": "openai", + "request_tags": [], + "prompt_tokens": 4_000, + "completion_tokens": 20, + "total_tokens": 4_020, + "response_cost": 0, + "startTime": now, + "endTime": now, + "metadata": { + "user_api_key_user_id": None, + "user_api_key_hash": None, + "user_api_key_alias": None, + "user_api_key_team_id": None, + "user_api_key_team_alias": None, + "user_api_key_user_email": None, + }, + "hidden_params": {}, + }, + }, + None, + now, + now, + ) + samples: Final = tuple( + sample + for metric in REGISTRY.collect() + for sample in metric.samples + if sample.name == "litellm_request_total_latency_metric_count" + ) + assert len(samples) == 1 + assert samples[0].value == 1 + assert samples[0].labels.get("input_sequence_length") == ("4k-16k" if db_value else None) + finally: + for collector in tuple(REGISTRY._collector_to_names): # pyright: ignore[reportPrivateUsage] # restore test registry + REGISTRY.unregister(collector) + for collector in collectors: + REGISTRY.register(collector) + + def test_otel_global_provider_published_after_callback_init(): """The OTel V2 global-provider publish must run after callback initialization in ``proxy_startup_event``. diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 9b27f8e3ceb..fde272bc192 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -19,6 +19,8 @@ from unittest.mock import AsyncMock, MagicMock import pytest import litellm +import litellm.proxy.proxy_server as ps +from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import CommonProxyErrors from litellm.proxy.proxy_server import ( ProxyConfig, @@ -1993,6 +1995,32 @@ def test_ProxyConfig__load_alerting_settings_noop_when_no_alerting(): } +def test_ProxyConfig__load_alerting_settings_preserves_other_logger_arguments(monkeypatch: pytest.MonkeyPatch): + logger: Final = CustomLogger() + factory: Final = MagicMock(return_value=logger) + proxy_logging: Final = MagicMock() + monkeypatch.setattr(ps, "_init_custom_logger_compatible_class", factory) + monkeypatch.setattr(ps, "proxy_logging_obj", proxy_logging) + settings: Final = { + "alerting": ["slack", "pagerduty", "prometheus"], + "alerting_args": {"routing_key": "test-routing-key"}, + "alerting_threshold": 15, + } + + ProxyConfig()._load_alerting_settings(settings) + + factory.assert_called_once_with( + logging_integration="pagerduty", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={"alerting_args": {"routing_key": "test-routing-key"}}, + ) + assert logger in ps.litellm.callbacks + assert "prometheus" in ps.litellm.callbacks + assert proxy_logging.update_values.call_args.kwargs["alerting"] == settings["alerting"] + assert proxy_logging.update_values.call_args.kwargs["alerting_threshold"] == 15 + + def test_ProxyConfig__load_alerting_settings_invalid_alerting_raises(): pc = ProxyConfig() with pytest.raises(RuntimeError): diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index cdec7acce05..e3da06039a4 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -7407,7 +7407,8 @@ async def test_batch_cost_poller_is_confirmed_before_serving(monkeypatch): @pytest.mark.asyncio -async def test_store_model_in_db_db_override_when_config_false(): +@pytest.mark.parametrize("resolve_before_logging", (False, True)) +async def test_store_model_in_db_db_override_when_config_false(resolve_before_logging: bool): """ Verify the early DB check in initialize_scheduled_background_jobs overrides store_model_in_db=False when DB has True. @@ -7430,8 +7431,11 @@ async def test_store_model_in_db_db_override_when_config_false(): with ( patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config), patch("litellm.proxy.proxy_server.store_model_in_db", False), - patch("litellm.proxy.proxy_server.get_secret_bool", return_value=False), + patch("litellm.proxy.proxy_server.get_secret_bool", return_value=False) as secret_lookup, ): + resolved: Final = resolve_before_logging and await ProxyStartupEvent.resolve_store_model_in_db( + prisma_client=mock_prisma_client, configured=False + ) await ProxyStartupEvent.initialize_scheduled_background_jobs( general_settings={}, prisma_client=mock_prisma_client, @@ -7439,8 +7443,13 @@ async def test_store_model_in_db_db_override_when_config_false(): proxy_budget_rescheduler_max_time=2, proxy_batch_write_at=5, proxy_logging_obj=mock_proxy_logging, + resolved_store_model_in_db=resolved, ) + mock_prisma_client.db.litellm_config.find_first.assert_awaited_once_with( + where={"param_name": "general_settings"} + ) + assert sum(args.args[0] == "STORE_MODEL_IN_DB" for args in secret_lookup.call_args_list) == 1 import litellm.proxy.proxy_server as ps # store_model_in_db should now be True (overridden by DB)