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

143 statements  

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

1import asyncio 

2import json 

3from collections.abc import Sequence 

4from dataclasses import asdict, dataclass 

5from typing import TYPE_CHECKING, Final 

6 

7from litellm._logging import verbose_proxy_logger 

8from litellm.proxy.common_utils.config_sync_pubsub import ( 

9 _ConfigSyncPubSub, 

10 _pubsub_capable_client, 

11 coordination_redis_cache, 

12) 

13 

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

15 from litellm.caching.in_memory_cache import InMemoryCache 

16 from litellm.caching.redis_cache import RedisCache 

17 from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache 

18 

19AUTH_CACHE_INVALIDATION_CHANNEL: Final = "litellm_proxy.auth_cache_invalidation" 

20_POLL_TIMEOUT_SECONDS: Final = 1.0 

21_MAX_PENDING_PUBLISHES: Final = 1024 

22_MAX_IN_FLIGHT_PUBLISHES: Final = 16 

23_pending_publishes: Final[set[asyncio.Task[None]]] = set() # mutable-ok: strong refs keep background publishes alive 

24_in_flight_publishes: Final = asyncio.Semaphore(_MAX_IN_FLIGHT_PUBLISHES) 

25_BACKOFF_INITIAL_SECONDS: Final = 5.0 

26_BACKOFF_MAX_SECONDS: Final = 60.0 

27 

28 

29def auth_cache_invalidation_channel(redis_cache: "RedisCache") -> str: 

30 if redis_cache.namespace is None: 

31 return AUTH_CACHE_INVALIDATION_CHANNEL 

32 return f"{redis_cache.namespace}:{AUTH_CACHE_INVALIDATION_CHANNEL}" 

33 

34 

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

36class _CacheInvalidationMessage: 

37 cache_key: str 

38 new_value: float | None = None 

39 ttl: float | None = None 

40 

41 

42def _cache_invalidation_message_json(cache_key: str, new_value: float | None = None, ttl: float | None = None) -> str: 

43 message: Final = asdict(_CacheInvalidationMessage(cache_key=cache_key, new_value=new_value, ttl=ttl)) 

44 return json.dumps({field: value for field, value in message.items() if value is not None}) 

45 

46 

47def _finite_number_or_none(value: object) -> float | None: 

48 if isinstance(value, bool) or not isinstance(value, (int, float)): 

49 return None 

50 return float(value) 

51 

52 

53def _message_from_data(data: object) -> _CacheInvalidationMessage | None: 

54 if isinstance(data, bytes): 

55 data = data.decode("utf-8", errors="replace") # rebind-ok: normalizing the wire payload to str 

56 if not isinstance(data, str): 

57 return None 

58 try: 

59 parsed: Final = json.loads(data) 

60 except json.JSONDecodeError: 

61 return None 

62 if not isinstance(parsed, dict): 

63 return None 

64 cache_key: Final = parsed.get("cache_key") 

65 if not isinstance(cache_key, str): 

66 return None 

67 return _CacheInvalidationMessage( 

68 cache_key=cache_key, 

69 new_value=_finite_number_or_none(parsed.get("new_value")), 

70 ttl=_finite_number_or_none(parsed.get("ttl")), 

71 ) 

72 

73 

74async def _publish_to_redis(redis_cache: "RedisCache", cache_key: str, message: str) -> None: 

75 try: 

76 client: Final = _pubsub_capable_client(redis_cache) 

77 if client is None: 

78 verbose_proxy_logger.debug( 

79 "auth cache invalidation publish for %s skipped: cluster redis client has no pub/sub support", 

80 cache_key, 

81 ) 

82 return 

83 async with _in_flight_publishes: 

84 await client.publish(auth_cache_invalidation_channel(redis_cache), message) 

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

86 verbose_proxy_logger.warning("auth cache invalidation publish for %s failed: %s", cache_key, e) 

87 

88 

89async def publish_auth_cache_invalidation( 

90 cache_key: str, new_value: float | None = None, ttl: float | None = None 

91) -> None: 

92 """ 

93 Best-effort broadcast so every worker drops its local in-memory copy of a 

94 mutated management object; without this, only the handling worker and Redis 

95 are evicted and other workers keep serving the stale object until its TTL. 

96 

97 Passing ``new_value`` broadcasts a SET instead of a delete: every subscriber 

98 (including the publishing worker's own, which receives its own message) 

99 writes the value into its additional in-memory caches rather than deleting 

100 the key. A spend reset uses this so the handler's self-delivered message 

101 cannot erase the freshly-written post-reset counter or floor marker. 

102 

103 The Redis round trip runs as a background task: this call returns once the 

104 publish has been handed to the event loop, so a Redis that accepts 

105 connections but never replies costs the caller nothing. The DB write has 

106 already committed and the local eviction already happened, so the caller 

107 has nothing to do with the publish result. At most 16 publishes hold a 

108 Redis connection at once; the rest wait in the task set, so a wedge cannot 

109 drain the shared connection pool. 

110 """ 

111 redis_cache: Final = coordination_redis_cache() 

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

113 return 

114 _pending_publishes.difference_update({task for task in _pending_publishes if task.done()}) 

115 if len(_pending_publishes) >= _MAX_PENDING_PUBLISHES: 

116 verbose_proxy_logger.warning( 

117 "auth cache invalidation publish for %s dropped: %d publishes already waiting on redis; " 

118 "other workers keep their cached copy until its TTL expires", 

119 cache_key, 

120 len(_pending_publishes), 

121 ) 

122 return 

123 task: Final = asyncio.create_task( 

124 _publish_to_redis( 

125 redis_cache, cache_key, _cache_invalidation_message_json(cache_key, new_value=new_value, ttl=ttl) 

126 ) 

127 ) 

128 _pending_publishes.add(task) 

129 await asyncio.sleep(0) 

130 

131 

132async def evict_and_broadcast(cache_keys: Sequence[str], user_api_key_cache: "UserApiKeyCache") -> None: 

133 """ 

134 Drop cached management objects here and on every other worker. 

135 

136 Every endpoint that mutates a cached object must call this: auth serves those objects 

137 cache-first with no freshness check, so a mutation that leaves the entry in place keeps the 

138 stale object enforced until its TTL expires (LIT-3803). Best-effort: the DB write has already 

139 committed, so a cache backend error must not fail the endpoint. 

140 """ 

141 for cache_key in cache_keys: 

142 try: 

143 await user_api_key_cache.async_delete_cache(key=cache_key) 

144 except Exception as e: # noqa: BLE001 # best-effort eviction: any cache backend error must not fail the mutation 

145 verbose_proxy_logger.warning( 

146 "Failed to evict cached entry %s; a stale object may be served until its TTL expires: %s", 

147 cache_key, 

148 e, 

149 ) 

150 await publish_auth_cache_invalidation(cache_key=cache_key) 

151 

152 

153class AuthCacheInvalidationSubscriber: 

154 __slots__ = ("_additional_in_memory_caches", "_redis_cache", "_task", "_user_api_key_cache") 

155 

156 def __init__( 

157 self, 

158 redis_cache: "RedisCache", 

159 user_api_key_cache: "UserApiKeyCache", 

160 additional_in_memory_caches: Sequence["InMemoryCache"] = (), 

161 ) -> None: 

162 self._redis_cache = redis_cache 

163 self._user_api_key_cache = user_api_key_cache 

164 self._additional_in_memory_caches = tuple(additional_in_memory_caches) 

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

166 

167 def start(self) -> None: 

168 if self._task is not None: 

169 return 

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

171 

172 async def stop(self) -> None: 

173 task: Final = self._task 

174 if task is None: 

175 return 

176 self._task = None 

177 _ = task.cancel() 

178 try: 

179 await task 

180 except asyncio.CancelledError: 

181 pass 

182 

183 async def _run(self) -> None: 

184 backoff_seconds = _BACKOFF_INITIAL_SECONDS # rebind-ok: exponential backoff accumulator across reconnects 

185 while True: 

186 try: 

187 client = _pubsub_capable_client(self._redis_cache) 

188 if client is None: 

189 verbose_proxy_logger.warning( 

190 "auth cache invalidation subscriber disabled: cluster redis client has no pub/sub support; " 

191 "cross-worker eviction falls back to the local cache TTL" 

192 ) 

193 return 

194 pubsub = client.pubsub() 

195 try: 

196 await pubsub.subscribe(auth_cache_invalidation_channel(self._redis_cache)) 

197 backoff_seconds = _BACKOFF_INITIAL_SECONDS 

198 await self._consume(pubsub) 

199 finally: 

200 await self._close_pubsub(pubsub) 

201 except asyncio.CancelledError: 

202 raise 

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

204 verbose_proxy_logger.warning( 

205 "auth cache invalidation subscriber redis error: %s; reconnecting in %.0fs", 

206 e, 

207 backoff_seconds, 

208 ) 

209 await asyncio.sleep(backoff_seconds) 

210 backoff_seconds = min(backoff_seconds * 2, _BACKOFF_MAX_SECONDS) 

211 

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

213 while True: 

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

215 if message is None: 

216 continue 

217 self._apply_message(message) 

218 

219 def _apply_message(self, message: object) -> None: 

220 data: Final = message.get("data") if isinstance(message, dict) else None 

221 parsed: Final = _message_from_data(data) 

222 if parsed is None: 

223 return 

224 if parsed.new_value is not None: 

225 for additional_cache in self._additional_in_memory_caches: 

226 additional_cache.set_cache(parsed.cache_key, parsed.new_value, ttl=parsed.ttl) 

227 return 

228 self._user_api_key_cache.in_memory_cache_for(parsed.cache_key).delete_cache(parsed.cache_key) 

229 for additional_cache in self._additional_in_memory_caches: 

230 additional_cache.delete_cache(parsed.cache_key) 

231 

232 @staticmethod 

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

234 try: 

235 await pubsub.aclose() 

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

237 verbose_proxy_logger.debug("auth cache invalidation pubsub close failed: %s", e)