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
« 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.
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).
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"""
20from __future__ import annotations
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
31from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
32 OAuthToken,
33)
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``."""
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
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)."""
53 async def acquire(self, key: str, token: str, ttl_seconds: float) -> LockAcquisition: ... 53 ↛ exitline 53 didn't return from function 'acquire' because
55 async def extend(self, key: str, token: str, ttl_seconds: float) -> bool: ... 55 ↛ exitline 55 didn't return from function 'extend' because
57 async def release(self, key: str, token: str) -> None: ... 57 ↛ exitline 57 didn't return from function 'release' because
59 async def is_held(self, key: str) -> bool: ... 59 ↛ exitline 59 didn't return from function 'is_held' because
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
82 def _key(self, user_id: str, server_id: str) -> str:
83 return f"{self.key_prefix}{user_id}:{server_id}"
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()
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)
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