mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(proxy): defer Prometheus alerting until stored settings load
Reuse successful startup storage resolution and preserve callback deduplication across alerting reloads.
This commit is contained in:
parent
a80945756a
commit
e8868d24b1
4 changed files with 157 additions and 3 deletions
|
|
@ -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
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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``.
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue