Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/common_utils/config_sync_pubsub.py: 45%

172 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1import asyncio 

2import json 

3import random 

4import time 

5from collections.abc import Awaitable, Callable 

6from dataclasses import asdict, dataclass 

7from typing import TYPE_CHECKING, Final, Protocol, cast # noqa: TID251 # untyped prisma/redis boundary needs cast 

8 

9from litellm._logging import verbose_proxy_logger 

10from litellm.repositories.prisma_protocols import RowT_co, TableActions 

11 

12if TYPE_CHECKING: 12 ↛ 13line 12 didn't jump to line 13 because the condition on line 12 was never true

13 from litellm.caching.redis_cache import RedisCache 

14 

15 

16class _ConfigSyncPubSub(Protocol): 

17 def subscribe(self, *channels: str) -> Awaitable[object]: ... 17 ↛ exitline 17 didn't return from function 'subscribe' because

18 

19 def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> Awaitable[object]: ... 19 ↛ exitline 19 didn't return from function 'get_message' because

20 

21 def aclose(self) -> Awaitable[object]: ... 21 ↛ exitline 21 didn't return from function 'aclose' because

22 

23 

24class _ConfigSyncPubSubClient(Protocol): 

25 def publish(self, channel: str, message: str) -> Awaitable[int]: ... 25 ↛ exitline 25 didn't return from function 'publish' because

26 

27 def pubsub(self) -> _ConfigSyncPubSub: ... 27 ↛ exitline 27 didn't return from function 'pubsub' because

28 

29 

30CONFIG_SYNC_CHANNEL: Final = "litellm_proxy.config_change" 

31CONFIG_SYNC_DEBOUNCE_SECONDS: Final = 1.0 

32CONFIG_SYNC_JITTER_MAX_SECONDS: Final = 5.0 

33CONFIG_SYNC_MIN_RESYNC_INTERVAL_SECONDS: Final = 10.0 

34_POLL_TIMEOUT_SECONDS: Final = 1.0 

35_BACKOFF_INITIAL_SECONDS: Final = 5.0 

36_BACKOFF_MAX_SECONDS: Final = 60.0 

37 

38_WRITE_ACTION_NAMES: Final[frozenset[str]] = frozenset( 

39 {"create", "create_many", "update", "update_many", "upsert", "delete", "delete_many"} 

40) 

41 

42_CONFIG_SYNCED_TABLE_NAMES: Final[frozenset[str]] = frozenset( 

43 { 

44 "litellm_proxymodeltable", 

45 "litellm_credentialstable", 

46 "litellm_guardrailstable", 

47 "litellm_policytable", 

48 "litellm_policyattachmenttable", 

49 "litellm_managedvectorstorestable", 

50 "litellm_managedvectorstoreindextable", 

51 "litellm_mcpservertable", 

52 "litellm_agentstable", 

53 "litellm_prompttable", 

54 "litellm_searchtoolstable", 

55 "litellm_ssoconfig", 

56 "litellm_cacheconfig", 

57 "litellm_configoverrides", 

58 "litellm_uisettings", 

59 } 

60) 

61 

62_RESYNC_APPLIED_CONFIG_PARAM_NAMES: Final[frozenset[str]] = frozenset( 

63 { 

64 "general_settings", 

65 "router_settings", 

66 "litellm_settings", 

67 "model_cost_map_reload_config", 

68 "anthropic_beta_headers_reload_config", 

69 } 

70) 

71 

72 

73def coordination_redis_cache() -> "RedisCache | None": 

74 from litellm.proxy.proxy_server import redis_usage_cache 

75 

76 return redis_usage_cache 

77 

78 

79def config_sync_channel(redis_cache: "RedisCache") -> str: 

80 if redis_cache.namespace is None: 

81 return CONFIG_SYNC_CHANNEL 

82 return f"{redis_cache.namespace}:{CONFIG_SYNC_CHANNEL}" 

83 

84 

85def _raw_async_client(redis_cache: "RedisCache") -> object: 

86 return cast( # cast-ok: redis-py generics leave the client type partially unknown 

87 object, 

88 redis_cache.init_async_client(), # pyright: ignore[reportUnknownMemberType] # redis generics 

89 ) 

90 

91 

92def _pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient | None: 

93 from redis.asyncio import Redis 

94 

95 client: Final = _raw_async_client(redis_cache) 

96 if isinstance(client, Redis): 

97 return cast(_ConfigSyncPubSubClient, client) # cast-ok: protocol view of the standalone redis client 

98 return None 

99 

100 

101@dataclass(frozen=True, slots=True) 

102class _ConfigChangeMessage: 

103 object_type: str 

104 

105 

106def _config_change_message_json(object_type: str) -> str: 

107 return json.dumps(asdict(_ConfigChangeMessage(object_type=object_type))) 

108 

109 

110async def publish_config_change(redis_cache: "RedisCache | None", object_type: str) -> None: 

111 if redis_cache is None: 111 ↛ 113line 111 didn't jump to line 113 because the condition on line 111 was always true

112 return 

113 try: 

114 client: Final = _pubsub_capable_client(redis_cache) 

115 if client is None: 

116 verbose_proxy_logger.debug( 

117 "config sync publish for %s skipped: cluster redis client has no pub/sub support", 

118 object_type, 

119 ) 

120 return 

121 await client.publish(config_sync_channel(redis_cache), _config_change_message_json(object_type)) 

122 except Exception as e: # noqa: BLE001 # best-effort publish; writes must never fail on redis errors 

123 verbose_proxy_logger.warning("config sync publish for %s failed: %s", object_type, e) 

124 

125 

126async def publish_config_change_for_object_type(object_type: str) -> None: 

127 await publish_config_change(redis_cache=coordination_redis_cache(), object_type=object_type) 

128 

129 

130async def publish_config_param_change(param_name: str) -> None: 

131 if param_name not in _RESYNC_APPLIED_CONFIG_PARAM_NAMES: 

132 verbose_proxy_logger.debug( 

133 "config sync publish for %s skipped: no resync callback applies this param outside proxy startup", 

134 param_name, 

135 ) 

136 return 

137 await publish_config_change_for_object_type(param_name) 

138 

139 

140class _PublishOnWriteActions: 

141 __slots__ = ("_actions", "_object_type", "_publish") 

142 

143 def __init__(self, actions: object, object_type: str, publish: Callable[[str], Awaitable[None]]) -> None: 

144 self._actions = actions 

145 self._object_type = object_type 

146 self._publish = publish 

147 

148 def __getattr__(self, name: str) -> object: 

149 attribute: Final = cast(object, getattr(self._actions, name)) # cast-ok: getattr on dynamic prisma actions 

150 if name not in _WRITE_ACTION_NAMES: 

151 return attribute 

152 write_action: Final = cast(Callable[..., Awaitable[object]], attribute) # cast-ok: prisma actions are untyped 

153 object_type: Final = self._object_type 

154 publish: Final = self._publish 

155 

156 async def _write_then_publish( 

157 *args: object, 

158 **kwargs: object, # kwargs-ok: transparent passthrough to untyped prisma action 

159 ) -> object: 

160 result: Final = await write_action(*args, **kwargs) 

161 await publish(object_type) 

162 return result 

163 

164 return _write_then_publish 

165 

166 

167def wrap_table_actions_for_config_sync( 

168 actions: "TableActions[RowT_co]", 

169 table_name: str, 

170 publish: Callable[[str], Awaitable[None]] = publish_config_change_for_object_type, 

171) -> "TableActions[RowT_co]": 

172 if table_name not in _CONFIG_SYNCED_TABLE_NAMES: 

173 return actions 

174 wrapped: Final = _PublishOnWriteActions(actions=actions, object_type=table_name, publish=publish) 

175 return cast("TableActions[RowT_co]", wrapped) # cast-ok: dynamic write-through proxy keeps the wrapped row type 

176 

177 

178class ConfigSyncSubscriber: 

179 __slots__ = ( 

180 "_backoff_initial_seconds", 

181 "_backoff_max_seconds", 

182 "_debounce_seconds", 

183 "_jitter_max_seconds", 

184 "_last_resync_at", 

185 "_min_resync_interval_seconds", 

186 "_monotonic", 

187 "_redis_cache", 

188 "_resync_callbacks", 

189 "_rng", 

190 "_sleep", 

191 "_task", 

192 ) 

193 

194 def __init__( 

195 self, 

196 redis_cache: "RedisCache", 

197 resync_callbacks: tuple[Callable[[], Awaitable[None]], ...], 

198 debounce_seconds: float = CONFIG_SYNC_DEBOUNCE_SECONDS, 

199 jitter_max_seconds: float = CONFIG_SYNC_JITTER_MAX_SECONDS, 

200 min_resync_interval_seconds: float = CONFIG_SYNC_MIN_RESYNC_INTERVAL_SECONDS, 

201 backoff_initial_seconds: float = _BACKOFF_INITIAL_SECONDS, 

202 backoff_max_seconds: float = _BACKOFF_MAX_SECONDS, 

203 rng: random.Random | None = None, 

204 sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, 

205 monotonic: Callable[[], float] = time.monotonic, 

206 ) -> None: 

207 self._redis_cache = redis_cache 

208 self._resync_callbacks = resync_callbacks 

209 self._debounce_seconds = debounce_seconds 

210 self._jitter_max_seconds = jitter_max_seconds 

211 self._min_resync_interval_seconds = min_resync_interval_seconds 

212 self._backoff_initial_seconds = backoff_initial_seconds 

213 self._backoff_max_seconds = backoff_max_seconds 

214 self._rng = rng if rng is not None else random.Random() 

215 self._sleep = sleep 

216 self._monotonic = monotonic 

217 self._task: asyncio.Task[None] | None = None 

218 self._last_resync_at: float | None = None 

219 

220 def start(self) -> None: 

221 if self._task is not None: 

222 return 

223 self._task = asyncio.create_task(self._run()) 

224 

225 async def stop(self) -> None: 

226 task: Final = self._task 

227 if task is None: 

228 return 

229 self._task = None 

230 _ = task.cancel() 

231 try: 

232 await task 

233 except asyncio.CancelledError: 

234 pass 

235 

236 async def _run(self) -> None: 

237 backoff_seconds = self._backoff_initial_seconds 

238 while True: 

239 try: 

240 client = _pubsub_capable_client(self._redis_cache) 

241 if client is None: 

242 verbose_proxy_logger.warning( 

243 "config sync subscriber disabled: cluster redis client has no pub/sub support; " 

244 "interval polling remains the only sync mechanism" 

245 ) 

246 return 

247 pubsub = client.pubsub() 

248 try: 

249 await pubsub.subscribe(config_sync_channel(self._redis_cache)) 

250 backoff_seconds = self._backoff_initial_seconds 

251 await self._consume(pubsub) 

252 finally: 

253 await self._close_pubsub(pubsub) 

254 except asyncio.CancelledError: 

255 raise 

256 except Exception as e: # noqa: BLE001 # any redis failure falls through to backoff and reconnect 

257 verbose_proxy_logger.warning( 

258 "config sync subscriber redis error: %s; reconnecting in %.0fs", 

259 e, 

260 backoff_seconds, 

261 ) 

262 await self._sleep(backoff_seconds) 

263 backoff_seconds = min(backoff_seconds * 2, self._backoff_max_seconds) 

264 

265 async def _consume(self, pubsub: _ConfigSyncPubSub) -> None: 

266 while True: 

267 message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=_POLL_TIMEOUT_SECONDS) 

268 if message is None: 

269 continue 

270 await self._sleep(self._debounce_seconds + self._rng.uniform(0.0, self._jitter_max_seconds)) 

271 await self._wait_for_min_resync_interval() 

272 await self._drain_pending(pubsub) 

273 await self._run_resync_callbacks() 

274 self._last_resync_at = self._monotonic() 

275 

276 async def _wait_for_min_resync_interval(self) -> None: 

277 if self._last_resync_at is None: 

278 return 

279 seconds_until_next_resync = self._min_resync_interval_seconds - (self._monotonic() - self._last_resync_at) 

280 if seconds_until_next_resync <= 0: 

281 return 

282 verbose_proxy_logger.debug( 

283 "config sync resync throttled for %.1fs to cap fleet-wide reload rate", 

284 seconds_until_next_resync, 

285 ) 

286 await self._sleep(seconds_until_next_resync) 

287 

288 @staticmethod 

289 async def _drain_pending(pubsub: _ConfigSyncPubSub) -> None: 

290 while await pubsub.get_message(ignore_subscribe_messages=True, timeout=0) is not None: 

291 pass 

292 

293 async def _run_resync_callbacks(self) -> None: 

294 for callback in self._resync_callbacks: 

295 try: 

296 await callback() 

297 except Exception as e: # noqa: BLE001 # one failing resync callback must not kill the subscriber 

298 verbose_proxy_logger.warning("config sync resync callback failed: %s", e) 

299 

300 @staticmethod 

301 async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None: 

302 try: 

303 await pubsub.aclose() 

304 except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection 

305 verbose_proxy_logger.debug("config sync pubsub close failed: %s", e)