fix(proxy): throttle pub/sub resyncs and stop publishing startup-only config params

Caps fleet-wide reload rate at one resync per 10s per pod so a burst of
authenticated writes cannot amplify into continuous cross-pod reloads, and
skips publishing config params (environment_variables, router_settings) that
no resync callback applies outside proxy startup
This commit is contained in:
mateo-berri 2026-07-31 21:08:32 -07:00
parent 629d58443e
commit a8018f7500
4 changed files with 380 additions and 28 deletions

View file

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

View file

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

View file

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

View file

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