Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_refresh_coordinator.py: 48%

61 statements  

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

1"""Cross-replica ``RefreshCoordinator``: one refresh per ``(user, server)`` across all workers. 

2 

3Plugs into the foundation's ``RefreshingTokenStore`` via the ``RefreshCoordinator`` seam. A ``SET NX 

4PX`` lock elects one worker to run the refresh while the rest wait for it and re-read the token it 

5persisted - so a rotating refresh_token is used once across the fleet, not once per worker. The holder 

6renews the ``PX`` lease while refresh runs (up to a refresh budget, so a hung endpoint can't hold the 

7lock forever), and a loser waits longer than that budget - so a loser only re-reads once the holder has 

8finished or its bounded lease has lapsed, never mid-refresh, and the surrounding store re-checks expiry 

9on the next fetch, so a crash self-heals rather than serving stale forever. Reading needs no lock, so 

10losers don't serialize behind each other. The lock is injected (a thin Redis wrapper in production, a 

11fake in tests). 

12 

13The lock is a single-flight optimization, not a correctness mutex, so it fails open: when the lock 

14backend is unreachable, ``acquire`` reports ``ERROR`` (distinct from ``HELD``) and this coordinator 

15refreshes anyway rather than wait on a holder that may not exist and then serve a still-expired token. 

16That degrades a Redis outage to the no-coordinator behavior (each worker may refresh), never a stale 

17bearer the upstream would 401. 

18""" 

19 

20from __future__ import annotations 

21 

22import asyncio 

23import time 

24import uuid 

25from collections.abc import Awaitable, Callable 

26from contextlib import suppress 

27from dataclasses import KW_ONLY, dataclass 

28from enum import Enum 

29from typing import Final, Protocol 

30 

31from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( 

32 OAuthToken, 

33) 

34 

35 

36class LockAcquisition(Enum): 

37 """Outcome of a best-effort ``acquire``. ``ERROR`` is kept distinct from ``HELD`` so a caller can 

38 tell "someone else is refreshing" (wait and re-read) from "the lock backend is down" (no election 

39 happened, so refresh anyway) instead of conflating both into a single ``False``.""" 

40 

41 ACQUIRED = "acquired" # won the election; this worker refreshes 

42 HELD = "held" # another worker holds it; wait then re-read 

43 ERROR = "error" # lock backend unreachable; holder unknown, so refresh anyway 

44 

45 

46class DistributedLock(Protocol): 

47 """A best-effort cross-replica lock. ``acquire`` is ``SET key token NX PX ttl`` reported as a 

48 ``LockAcquisition`` (won / held by another / backend error); ``release`` deletes the key only if 

49 it still holds this caller's ``token`` (so it cannot delete a lock another worker re-acquired 

50 after PX-expiry); ``extend`` refreshes the ``PX`` lease only for the owner; ``is_held`` is 

51 ``EXISTS`` (so a waiter can poll without taking the lock).""" 

52 

53 async def acquire(self, key: str, token: str, ttl_seconds: float) -> LockAcquisition: ... 53 ↛ exitline 53 didn't return from function 'acquire' because

54 

55 async def extend(self, key: str, token: str, ttl_seconds: float) -> bool: ... 55 ↛ exitline 55 didn't return from function 'extend' because

56 

57 async def release(self, key: str, token: str) -> None: ... 57 ↛ exitline 57 didn't return from function 'release' because

58 

59 async def is_held(self, key: str) -> bool: ... 59 ↛ exitline 59 didn't return from function 'is_held' because

60 

61 

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

63class RedisRefreshCoordinator: 

64 lock: DistributedLock 

65 _: KW_ONLY 

66 key_prefix: str = "mcp:refresh_lock:" 

67 lock_ttl_seconds: float = 10.0 

68 # The holder renews its lease while a slow token endpoint runs, but only up to this budget; past it 

69 # it stops renewing and the lock lapses, so a hung refresh degrades to "maybe an extra refresh" 

70 # rather than holding every loser behind it indefinitely. 

71 refresh_budget_seconds: float = 20.0 

72 # How long a loser waits for the holder before giving up and re-reading. It MUST outlast the 

73 # holder's max lock-hold (refresh_budget_seconds + one lock_ttl_seconds tail); otherwise a loser 

74 # bails while the holder is still legitimately refreshing, re-reads the still-expired token, and 

75 # challenges the user mid-refresh. 

76 wait_timeout_seconds: float = 35.0 

77 poll_interval_seconds: float = 0.05 

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

79 clock: Callable[[], float] = time.monotonic 

80 new_token: Callable[[], str] = lambda: uuid.uuid4().hex 

81 

82 def _key(self, user_id: str, server_id: str) -> str: 

83 return f"{self.key_prefix}{user_id}:{server_id}" 

84 

85 async def run( 

86 self, 

87 user_id: str, 

88 server_id: str, 

89 refresh: Callable[[], Awaitable[OAuthToken | None]], 

90 reread: Callable[[], Awaitable[OAuthToken | None]], 

91 ) -> OAuthToken | None: 

92 key: Final = self._key(user_id, server_id) 

93 token: Final = self.new_token() 

94 match await self.lock.acquire(key, token, self.lock_ttl_seconds): 

95 case LockAcquisition.ACQUIRED: 

96 return await self._refresh_with_lease_renewal(key, token, refresh) 

97 case LockAcquisition.ERROR: 

98 # No election happened (lock backend down), so waiting would just re-read the 

99 # still-expired token. Refresh anyway; worst case is an extra refresh, not a stale bearer. 

100 return await refresh() 

101 case LockAcquisition.HELD: 

102 # Another worker holds the lock; wait for it to finish (release or PX-expiry), then read 

103 # the token it persisted - the winner wrote the fresh token to the store, so a plain 

104 # re-read sees it without us refreshing again. 

105 deadline: Final = self.clock() + self.wait_timeout_seconds 

106 while self.clock() < deadline and await self.lock.is_held(key): 

107 await self.sleep(self.poll_interval_seconds) 

108 return await reread() 

109 

110 async def _refresh_with_lease_renewal( 

111 self, 

112 key: str, 

113 token: str, 

114 refresh: Callable[[], Awaitable[OAuthToken | None]], 

115 ) -> OAuthToken | None: 

116 refresh_task: Final = asyncio.ensure_future(refresh()) 

117 renewal_task: Final = asyncio.create_task(self._renew_lease_until_done(key, token, refresh_task)) 

118 try: 

119 return await refresh_task 

120 finally: 

121 renewal_task.cancel() 

122 with suppress(asyncio.CancelledError): 

123 await renewal_task 

124 await self.lock.release(key, token) 

125 

126 async def _renew_lease_until_done( 

127 self, 

128 key: str, 

129 token: str, 

130 refresh_task: asyncio.Future[OAuthToken | None], 

131 ) -> None: 

132 budget_deadline: Final = self.clock() + self.refresh_budget_seconds 

133 while not refresh_task.done() and self.clock() < budget_deadline: 

134 await self.sleep(self.lock_ttl_seconds / 2) 

135 if not refresh_task.done() and not await self.lock.extend(key, token, self.lock_ttl_seconds): 

136 return