diff --git a/litellm/proxy/common_utils/config_sync_pubsub.py b/litellm/proxy/common_utils/config_sync_pubsub.py index 4292c67c8c2..e4521c397b7 100644 --- a/litellm/proxy/common_utils/config_sync_pubsub.py +++ b/litellm/proxy/common_utils/config_sync_pubsub.py @@ -1,6 +1,7 @@ import asyncio import json import random +import time from collections.abc import Awaitable, Callable from dataclasses import asdict, dataclass from typing import TYPE_CHECKING, Protocol, cast # noqa: TID251 # untyped prisma/redis boundary needs cast @@ -28,6 +29,7 @@ class _ConfigSyncPubSubClient(Protocol): CONFIG_SYNC_CHANNEL = "litellm_proxy.config_change" CONFIG_SYNC_DEBOUNCE_SECONDS = 1.0 CONFIG_SYNC_JITTER_MAX_SECONDS = 5.0 +CONFIG_SYNC_MIN_RESYNC_INTERVAL_SECONDS = 10.0 _POLL_TIMEOUT_SECONDS = 1.0 _BACKOFF_INITIAL_SECONDS = 5.0 _BACKOFF_MAX_SECONDS = 60.0 @@ -55,6 +57,15 @@ _CONFIG_SYNCED_TABLE_NAMES: frozenset[str] = frozenset( } ) +_RESYNC_APPLIED_CONFIG_PARAM_NAMES: frozenset[str] = frozenset( + { + "general_settings", + "litellm_settings", + "model_cost_map_reload_config", + "anthropic_beta_headers_reload_config", + } +) + def coordination_redis_cache() -> "RedisCache | None": from litellm.proxy.proxy_server import redis_usage_cache @@ -113,6 +124,16 @@ async def publish_config_change_for_object_type(object_type: str) -> None: await publish_config_change(redis_cache=coordination_redis_cache(), object_type=object_type) +async def publish_config_param_change(param_name: str) -> None: + if param_name not in _RESYNC_APPLIED_CONFIG_PARAM_NAMES: + verbose_proxy_logger.debug( + "config sync publish for %s skipped: no resync callback applies this param outside proxy startup", + param_name, + ) + return + await publish_config_change_for_object_type(param_name) + + class _PublishOnWriteActions: __slots__ = ("_actions", "_object_type", "_publish") @@ -156,6 +177,9 @@ class ConfigSyncSubscriber: "_backoff_max_seconds", "_debounce_seconds", "_jitter_max_seconds", + "_last_resync_at", + "_min_resync_interval_seconds", + "_monotonic", "_redis_cache", "_resync_callbacks", "_rng", @@ -169,20 +193,25 @@ class ConfigSyncSubscriber: resync_callbacks: tuple[Callable[[], Awaitable[None]], ...], debounce_seconds: float = CONFIG_SYNC_DEBOUNCE_SECONDS, jitter_max_seconds: float = CONFIG_SYNC_JITTER_MAX_SECONDS, + min_resync_interval_seconds: float = CONFIG_SYNC_MIN_RESYNC_INTERVAL_SECONDS, backoff_initial_seconds: float = _BACKOFF_INITIAL_SECONDS, backoff_max_seconds: float = _BACKOFF_MAX_SECONDS, rng: random.Random | None = None, sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, + monotonic: Callable[[], float] = time.monotonic, ) -> None: self._redis_cache = redis_cache self._resync_callbacks = resync_callbacks self._debounce_seconds = debounce_seconds self._jitter_max_seconds = jitter_max_seconds + self._min_resync_interval_seconds = min_resync_interval_seconds self._backoff_initial_seconds = backoff_initial_seconds self._backoff_max_seconds = backoff_max_seconds self._rng = rng if rng is not None else random.Random() self._sleep = sleep + self._monotonic = monotonic self._task: asyncio.Task[None] | None = None + self._last_resync_at: float | None = None def start(self) -> None: if self._task is not None: @@ -235,8 +264,22 @@ class ConfigSyncSubscriber: if message is None: continue await self._sleep(self._debounce_seconds + self._rng.uniform(0.0, self._jitter_max_seconds)) + await self._wait_for_min_resync_interval() await self._drain_pending(pubsub) await self._run_resync_callbacks() + self._last_resync_at = self._monotonic() + + async def _wait_for_min_resync_interval(self) -> None: + if self._last_resync_at is None: + return + seconds_until_next_resync = self._min_resync_interval_seconds - (self._monotonic() - self._last_resync_at) + if seconds_until_next_resync <= 0: + return + verbose_proxy_logger.debug( + "config sync resync throttled for %.1fs to cap fleet-wide reload rate", + seconds_until_next_resync, + ) + await self._sleep(seconds_until_next_resync) @staticmethod async def _drain_pending(pubsub: _ConfigSyncPubSub) -> None: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6bb0032b38c..7d3e58bfe52 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1153,11 +1153,7 @@ async def proxy_startup_event(app: FastAPI): except Exception as e: verbose_proxy_logger.error(f"Error stopping DB health watchdog task: {e}") - if proxy_config.config_sync_subscriber is not None: - try: - await proxy_config.config_sync_subscriber.stop() - except Exception as e: - verbose_proxy_logger.error(f"Error stopping config sync subscriber: {e}") + await proxy_config.stop_config_sync_subscriber() await proxy_shutdown_event() # type: ignore[reportGeneralTypeIssues] @@ -6213,6 +6209,38 @@ class ProxyConfig: "litellm.proxy.proxy_server.py::ProxyConfig:add_deployment - {}".format(str(e)) ) + def start_config_sync_subscriber( + self, + prisma_client: PrismaClient, + proxy_logging_obj: ProxyLogging, + redis_cache: Optional[RedisCache], + ) -> None: + if redis_cache is None or self.config_sync_subscriber is not None: + return + + async def _resync_config_from_db() -> None: + await self.add_deployment(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj) + + async def _resync_credentials_from_db() -> None: + await self.get_credentials(prisma_client=prisma_client) + + subscriber = ConfigSyncSubscriber( + redis_cache=redis_cache, + resync_callbacks=(_resync_config_from_db, _resync_credentials_from_db), + ) + self.config_sync_subscriber = subscriber + subscriber.start() + + async def stop_config_sync_subscriber(self) -> None: + subscriber = self.config_sync_subscriber + if subscriber is None: + return + self.config_sync_subscriber = None + try: + await subscriber.stop() + except Exception as e: + verbose_proxy_logger.error(f"Error stopping config sync subscriber: {e}") + async def _init_non_llm_objects_in_db(self, prisma_client: PrismaClient): """ Use this to read non-llm objects from the db and initialize them @@ -8174,19 +8202,11 @@ class ProxyStartupEvent: ) await proxy_config.get_credentials(prisma_client=prisma_client) - if redis_usage_cache is not None and proxy_config.config_sync_subscriber is None: - - async def _resync_config_from_db() -> None: - await proxy_config.add_deployment(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj) - - async def _resync_credentials_from_db() -> None: - await proxy_config.get_credentials(prisma_client=prisma_client) - - proxy_config.config_sync_subscriber = ConfigSyncSubscriber( - redis_cache=redis_usage_cache, - resync_callbacks=(_resync_config_from_db, _resync_credentials_from_db), - ) - proxy_config.config_sync_subscriber.start() + proxy_config.start_config_sync_subscriber( + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + redis_cache=redis_usage_cache, + ) if store_model_in_db is not True: await proxy_config.init_mcp_servers_from_db() diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index fbcca73e779..5d2d29efa12 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -119,10 +119,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.route_checks import RouteChecks -from litellm.proxy.common_utils.config_sync_pubsub import ( - coordination_redis_cache, - publish_config_change, -) +from litellm.proxy.common_utils.config_sync_pubsub import publish_config_param_change from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.db.create_views import ( create_missing_views, @@ -2982,7 +2979,7 @@ async def evict_config_param(param_name: str) -> None: async def invalidate_config_param(param_name: str) -> None: """Evict from both cache layers; call after every LiteLLM_Config write.""" await evict_config_param(param_name) - await publish_config_change(redis_cache=coordination_redis_cache(), object_type=param_name) + await publish_config_param_change(param_name) async def prefetch_config_params(prisma_client: Any, param_names: List[str]) -> None: diff --git a/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py b/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py index 8a8ced8bc41..f50eef4f1cc 100644 --- a/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py +++ b/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py @@ -11,9 +11,11 @@ import litellm from litellm.proxy.common_utils.config_sync_pubsub import ( CONFIG_SYNC_CHANNEL, CONFIG_SYNC_JITTER_MAX_SECONDS, + CONFIG_SYNC_MIN_RESYNC_INTERVAL_SECONDS, ConfigSyncSubscriber, _CONFIG_SYNCED_TABLE_NAMES, _PublishOnWriteActions, + _RESYNC_APPLIED_CONFIG_PARAM_NAMES, _WRITE_ACTION_NAMES, publish_config_change, wrap_table_actions_for_config_sync, @@ -48,6 +50,17 @@ _EXPECTED_CONFIG_SYNCED_TABLE_NAMES = frozenset( } ) +_EXPECTED_RESYNC_APPLIED_CONFIG_PARAM_NAMES = frozenset( + { + "anthropic_beta_headers_reload_config", + "general_settings", + "litellm_settings", + "model_cost_map_reload_config", + } +) + +_STARTUP_ONLY_CONFIG_PARAM_NAMES = ("environment_variables", "router_settings") + class _RecordingRedisClient(Redis): def __init__(self) -> None: @@ -106,6 +119,31 @@ class _BrokenPubSub(_QueuePubSub): raise ConnectionError("connection lost") +class _CloseFailingBrokenPubSub(_BrokenPubSub): + async def aclose(self) -> None: + raise ConnectionError("close failed") + + +class _EmptyPollsThenMessagePubSub(_QueuePubSub): + def __init__(self, empty_polls: int, initial_messages: Iterable[str] = ()) -> None: + super().__init__(initial_messages=initial_messages) + self.remaining_empty_polls = empty_polls + + async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> Optional[str]: + if timeout != 0 and self.remaining_empty_polls > 0: + self.remaining_empty_polls -= 1 + return None + return await super().get_message(ignore_subscribe_messages=ignore_subscribe_messages, timeout=timeout) + + +class _FakeClock: + def __init__(self, now: float = 1000.0) -> None: + self.now = now + + def __call__(self) -> float: + return self.now + + class _ScriptedPubSubRedisClient(Redis): def __init__(self, pubsubs: Iterable[_QueuePubSub]) -> None: self._scripted_pubsubs = iter(pubsubs) @@ -286,6 +324,158 @@ def test_default_jitter_window_is_nonzero() -> None: assert CONFIG_SYNC_JITTER_MAX_SECONDS > 0 +def test_default_min_resync_interval_caps_reload_rate() -> None: + assert CONFIG_SYNC_MIN_RESYNC_INTERVAL_SECONDS > CONFIG_SYNC_JITTER_MAX_SECONDS + + +def _throttled_subscriber( + cache: object, + events: List[str], + fired: asyncio.Event, + clock: _FakeClock, + min_resync_interval_seconds: float = 10.0, +) -> ConfigSyncSubscriber: + async def recording_sleep(seconds: float) -> None: + events.append(f"sleep:{seconds}") + await asyncio.sleep(0) + + async def resync() -> None: + events.append("resync") + fired.set() + + return ConfigSyncSubscriber( + redis_cache=cache, + resync_callbacks=(resync,), + debounce_seconds=0.0, + jitter_max_seconds=0.0, + min_resync_interval_seconds=min_resync_interval_seconds, + sleep=recording_sleep, + monotonic=clock, + ) + + +async def test_resync_arriving_inside_min_interval_waits_out_the_remainder() -> None: + pubsub = _QueuePubSub() + cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) + events: List[str] = [] + fired = asyncio.Event() + clock = _FakeClock() + subscriber = _throttled_subscriber(cache=cache, events=events, fired=fired, clock=clock) + + subscriber.start() + pubsub.queue.put_nowait("change") + await asyncio.wait_for(fired.wait(), timeout=5) + fired.clear() + clock.now += 4.0 + pubsub.queue.put_nowait("change") + await asyncio.wait_for(fired.wait(), timeout=5) + await subscriber.stop() + + assert events == ["sleep:0.0", "resync", "sleep:0.0", "sleep:6.0", "resync"] + + +async def test_resync_after_min_interval_elapsed_is_not_throttled() -> None: + pubsub = _QueuePubSub() + cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) + events: List[str] = [] + fired = asyncio.Event() + clock = _FakeClock() + subscriber = _throttled_subscriber(cache=cache, events=events, fired=fired, clock=clock) + + subscriber.start() + pubsub.queue.put_nowait("change") + await asyncio.wait_for(fired.wait(), timeout=5) + fired.clear() + clock.now += 30.0 + pubsub.queue.put_nowait("change") + await asyncio.wait_for(fired.wait(), timeout=5) + await subscriber.stop() + + assert events == ["sleep:0.0", "resync", "sleep:0.0", "resync"] + + +async def test_writes_during_the_throttle_wait_collapse_into_the_next_resync() -> None: + pubsub = _QueuePubSub() + cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) + events: List[str] = [] + fired = asyncio.Event() + clock = _FakeClock() + subscriber = _throttled_subscriber(cache=cache, events=events, fired=fired, clock=clock) + + subscriber.start() + pubsub.queue.put_nowait("change") + await asyncio.wait_for(fired.wait(), timeout=5) + fired.clear() + for _ in range(5): + pubsub.queue.put_nowait("change") + await asyncio.wait_for(fired.wait(), timeout=5) + await asyncio.sleep(0.1) + await subscriber.stop() + + assert events.count("resync") == 2 + assert pubsub.queue.empty() + + +async def test_polls_without_messages_do_not_trigger_resyncs() -> None: + pubsub = _EmptyPollsThenMessagePubSub(empty_polls=3, initial_messages=["change"]) + cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) + resyncs: List[str] = [] + fired = asyncio.Event() + subscriber = ConfigSyncSubscriber( + redis_cache=cache, + resync_callbacks=(_recording_callback(resyncs, "resync", fired),), + debounce_seconds=0.01, + jitter_max_seconds=0.0, + ) + + subscriber.start() + await asyncio.wait_for(fired.wait(), timeout=5) + await asyncio.sleep(0.1) + await subscriber.stop() + + assert pubsub.remaining_empty_polls == 0 + assert resyncs == ["resync"] + + +async def test_failing_pubsub_close_still_reconnects() -> None: + broken = _CloseFailingBrokenPubSub() + healthy = _QueuePubSub(initial_messages=["change"]) + cache = _FakeRedisCache(_ScriptedPubSubRedisClient([broken, healthy])) + resyncs: List[str] = [] + fired = asyncio.Event() + subscriber = ConfigSyncSubscriber( + redis_cache=cache, + resync_callbacks=(_recording_callback(resyncs, "resync", fired),), + debounce_seconds=0.01, + jitter_max_seconds=0.0, + backoff_initial_seconds=0.02, + backoff_max_seconds=0.05, + ) + + subscriber.start() + await asyncio.wait_for(fired.wait(), timeout=5) + await subscriber.stop() + + assert healthy.subscribed_channels == [CONFIG_SYNC_CHANNEL] + assert resyncs == ["resync"] + + +async def test_second_start_does_not_open_a_second_subscription() -> None: + pubsub = _QueuePubSub() + cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) + subscriber = ConfigSyncSubscriber(redis_cache=cache, resync_callbacks=(), debounce_seconds=0.01) + + subscriber.start() + task = subscriber._task + subscriber.start() + assert task is not None + assert subscriber._task is task + await asyncio.sleep(0.05) + await subscriber.stop() + + assert pubsub.subscribed_channels == [CONFIG_SYNC_CHANNEL] + + async def test_redis_error_leads_to_backoff_and_resubscribe() -> None: broken = _BrokenPubSub() healthy = _QueuePubSub(initial_messages=[json.dumps({"object_type": "litellm_credentialstable"})]) @@ -328,6 +518,7 @@ async def test_failing_resync_callback_does_not_kill_subscriber() -> None: resync_callbacks=(failing_callback, _recording_callback(resyncs, "resync", fired)), debounce_seconds=0.01, jitter_max_seconds=0.0, + min_resync_interval_seconds=0.0, ) subscriber.start() @@ -510,7 +701,7 @@ async def test_model_repository_write_publishes_via_live_coordination_cache() -> assert json.loads(message) == {"object_type": "litellm_proxymodeltable"} -async def test_invalidate_config_param_publishes_param_name() -> None: +async def _publish_calls_for_invalidated_param(param_name: str) -> List[Tuple[str, str]]: from litellm.proxy import proxy_server from litellm.proxy.proxy_server import _set_redis_usage_cache from litellm.proxy.utils import invalidate_config_param @@ -519,14 +710,30 @@ async def test_invalidate_config_param_publishes_param_name() -> None: previous_cache = proxy_server.redis_usage_cache _set_redis_usage_cache(_FakeRedisCache(client)) try: - await invalidate_config_param("environment_variables") + await invalidate_config_param(param_name) finally: _set_redis_usage_cache(previous_cache) + return client.published - assert len(client.published) == 1 - channel, message = client.published[0] + +async def test_invalidate_config_param_publishes_params_a_resync_applies() -> None: + published = await _publish_calls_for_invalidated_param("general_settings") + + assert len(published) == 1 + channel, message = published[0] assert channel == CONFIG_SYNC_CHANNEL - assert json.loads(message) == {"object_type": "environment_variables"} + assert json.loads(message) == {"object_type": "general_settings"} + + +@pytest.mark.parametrize("param_name", _STARTUP_ONLY_CONFIG_PARAM_NAMES) +async def test_invalidate_config_param_does_not_publish_startup_only_params(param_name: str) -> None: + published = await _publish_calls_for_invalidated_param(param_name) + + assert published == [] + + +def test_resync_applied_config_param_membership_is_pinned() -> None: + assert _RESYNC_APPLIED_CONFIG_PARAM_NAMES == _EXPECTED_RESYNC_APPLIED_CONFIG_PARAM_NAMES async def test_evict_config_param_does_not_publish() -> None: @@ -598,3 +805,88 @@ async def test_anthropic_beta_headers_reload_does_not_publish_config_change() -> prisma_client.db.litellm_config.upsert.assert_awaited_once() assert client.published == [] + + +class _StopFailingSubscriber(ConfigSyncSubscriber): + async def stop(self) -> None: + raise RuntimeError("stop failed") + + +async def test_proxy_config_subscriber_resyncs_deployments_and_credentials() -> None: + from litellm.proxy.proxy_server import ProxyConfig + + cache = _FakeRedisCache(_ScriptedPubSubRedisClient([_QueuePubSub()])) + config = ProxyConfig() + prisma_client = MagicMock() + proxy_logging_obj = MagicMock() + calls: List[Tuple[str, object, object]] = [] + + async def fake_add_deployment(prisma_client: object, proxy_logging_obj: object) -> None: + calls.append(("add_deployment", prisma_client, proxy_logging_obj)) + + async def fake_get_credentials(prisma_client: object) -> None: + calls.append(("get_credentials", prisma_client, None)) + + config.add_deployment = fake_add_deployment + config.get_credentials = fake_get_credentials + config.start_config_sync_subscriber( + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + redis_cache=cache, + ) + subscriber = config.config_sync_subscriber + assert subscriber is not None + for callback in subscriber._resync_callbacks: + await callback() + await config.stop_config_sync_subscriber() + + assert calls == [ + ("add_deployment", prisma_client, proxy_logging_obj), + ("get_credentials", prisma_client, None), + ] + assert config.config_sync_subscriber is None + assert subscriber._task is None + + +async def test_proxy_config_does_not_start_subscriber_without_coordination_redis() -> None: + from litellm.proxy.proxy_server import ProxyConfig + + config = ProxyConfig() + + config.start_config_sync_subscriber( + prisma_client=MagicMock(), + proxy_logging_obj=MagicMock(), + redis_cache=None, + ) + + assert config.config_sync_subscriber is None + + +async def test_proxy_config_keeps_the_first_subscriber_on_repeat_start() -> None: + from litellm.proxy.proxy_server import ProxyConfig + + cache = _FakeRedisCache(_ScriptedPubSubRedisClient([_QueuePubSub()])) + config = ProxyConfig() + + config.start_config_sync_subscriber(prisma_client=MagicMock(), proxy_logging_obj=MagicMock(), redis_cache=cache) + first = config.config_sync_subscriber + config.start_config_sync_subscriber(prisma_client=MagicMock(), proxy_logging_obj=MagicMock(), redis_cache=cache) + second = config.config_sync_subscriber + await config.stop_config_sync_subscriber() + + assert first is not None + assert second is first + + +async def test_proxy_config_shutdown_survives_a_failing_subscriber_stop() -> None: + from litellm.proxy.proxy_server import ProxyConfig + + config = ProxyConfig() + config.config_sync_subscriber = _StopFailingSubscriber( + redis_cache=_FakeRedisCache(_ScriptedPubSubRedisClient([])), + resync_callbacks=(), + ) + + await config.stop_config_sync_subscriber() + + assert config.config_sync_subscriber is None