feat(mcp): Redis SET NX PX refresh coordinator (step 1b §1.5)

The cross-replica RefreshCoordinator that plugs into the foundation's RefreshingTokenStore seam: a SET
NX PX lock elects one worker to refresh per (user, server) while the rest wait for it and re-read the
token it persisted, so a rotating refresh_token is used once across the fleet, not once per worker. The
lock self-expires (PX) so a crashed holder can't wedge refresh; a loser falls back to a bounded re-read
and the surrounding store re-checks expiry next fetch, so a crash self-heals. The lock (a thin Redis
SET NX/DEL/EXISTS wrapper in prod) is injected, so the single-flight logic is testable without Redis.
This commit is contained in:
Tin Chi Lo 2026-06-25 23:14:43 -07:00
parent b46733e5b9
commit f8b07a089f
2 changed files with 209 additions and 0 deletions

View file

@ -0,0 +1,80 @@
"""Cross-replica ``RefreshCoordinator``: one refresh per ``(user, server)`` across all workers.
Plugs into the foundation's ``RefreshingTokenStore`` via the ``RefreshCoordinator`` seam. A ``SET NX
PX`` lock elects one worker to run the refresh while the rest wait for it and re-read the token it
persisted - so a rotating refresh_token is used once across the fleet, not once per worker. The lock
auto-expires (``PX``), so a crashed holder can't wedge refresh; a loser that times out (or whose
holder crashed mid-refresh) falls back to a re-read, and the surrounding store re-checks expiry on the
next fetch, so a crash self-heals rather than serving stale forever. Reading needs no lock, so losers
don't serialize behind each other. The lock is injected (a thin Redis ``SET NX``/``DEL``/``EXISTS``
wrapper in production, a fake in tests).
"""
from __future__ import annotations
import asyncio
import time
from collections.abc import Awaitable, Callable
from typing import Protocol
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
class DistributedLock(Protocol):
"""A best-effort cross-replica lock. ``acquire`` is ``SET key NX PX ttl`` (only the first caller
wins, the entry self-expires); ``release`` is ``DEL``; ``is_held`` is ``EXISTS`` (so a waiter can
poll without taking the lock)."""
async def acquire(self, key: str, ttl_seconds: float) -> bool: ...
async def release(self, key: str) -> None: ...
async def is_held(self, key: str) -> bool: ...
class RedisRefreshCoordinator:
def __init__(
self,
lock: DistributedLock,
*,
key_prefix: str = "mcp:refresh_lock:",
lock_ttl_seconds: float = 10.0,
wait_timeout_seconds: float = 10.0,
poll_interval_seconds: float = 0.05,
sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
clock: Callable[[], float] = time.monotonic,
) -> None:
self._lock = lock
self._key_prefix = key_prefix
self._lock_ttl_seconds = lock_ttl_seconds
self._wait_timeout_seconds = wait_timeout_seconds
self._poll_interval_seconds = poll_interval_seconds
self._sleep = sleep
self._clock = clock
def _key(self, user_id: str, server_id: str) -> str:
return f"{self._key_prefix}{user_id}:{server_id}"
async def run(
self,
user_id: str,
server_id: str,
refresh: Callable[[], Awaitable[OAuthToken | None]],
reread: Callable[[], Awaitable[OAuthToken | None]],
) -> OAuthToken | None:
key = self._key(user_id, server_id)
if await self._lock.acquire(key, self._lock_ttl_seconds):
try:
return await refresh()
finally:
await self._lock.release(key)
# Another worker holds the lock; wait for it to finish (release or PX-expiry), then read the
# token it persisted - the winner wrote the fresh token to the store, so a plain re-read sees
# it without us refreshing again.
deadline = self._clock() + self._wait_timeout_seconds
while self._clock() < deadline and await self._lock.is_held(key):
await self._sleep(self._poll_interval_seconds)
return await reread()

View file

@ -0,0 +1,129 @@
"""Tests for the cross-replica refresh coordinator: winner refreshes, losers wait then re-read."""
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_coordinator import (
RedisRefreshCoordinator,
)
_KEY = "mcp:refresh_lock:u:s"
class _FakeLock:
def __init__(self, acquired, held_sequence=()):
self._acquired = acquired
self._held = list(held_sequence)
self.acquired_keys = []
self.released = []
async def acquire(self, key, ttl_seconds):
self.acquired_keys.append((key, ttl_seconds))
return self._acquired
async def release(self, key):
self.released.append(key)
async def is_held(self, key):
return self._held.pop(0) if self._held else False
class _Clock:
def __init__(self):
self.t = 0.0
def __call__(self):
return self.t
def _advancing_sleep(clock, step=0.5):
async def sleep(_seconds):
clock.t += step
return sleep
@pytest.mark.asyncio
async def test_winner_refreshes_then_releases_and_never_rereads():
lock = _FakeLock(acquired=True)
refreshed = OAuthToken(access_token="new")
reread_calls = []
async def refresh():
return refreshed
async def reread():
reread_calls.append(1)
return None
result = await RedisRefreshCoordinator(lock).run("u", "s", refresh, reread)
assert result is refreshed
assert lock.acquired_keys == [(_KEY, 10.0)]
assert lock.released == [_KEY]
assert reread_calls == []
@pytest.mark.asyncio
async def test_winner_releases_even_when_refresh_raises():
lock = _FakeLock(acquired=True)
async def refresh():
raise RuntimeError("boom")
async def reread():
return None
with pytest.raises(RuntimeError):
await RedisRefreshCoordinator(lock).run("u", "s", refresh, reread)
assert lock.released == [_KEY] # the lock is freed even on failure
@pytest.mark.asyncio
async def test_loser_waits_for_the_holder_then_rereads_persisted_token():
clock = _Clock()
lock = _FakeLock(acquired=False, held_sequence=[True, True, False])
refresh_calls = []
async def refresh():
refresh_calls.append(1)
return None
async def reread():
return OAuthToken(access_token="persisted-by-winner")
coord = RedisRefreshCoordinator(
lock, clock=clock, sleep=_advancing_sleep(clock), wait_timeout_seconds=100.0
)
result = await coord.run("u", "s", refresh, reread)
assert result is not None and result.access_token == "persisted-by-winner"
assert (
refresh_calls == []
) # the loser never refreshes - it reads the winner's result
assert lock.released == [] # ...and never holds the lock
@pytest.mark.asyncio
async def test_loser_rereads_after_timeout_if_holder_never_releases():
clock = _Clock()
lock = _FakeLock(
acquired=False, held_sequence=[True] * 100
) # holder never releases
async def refresh():
return None
async def reread():
return OAuthToken(access_token="whatever-is-there")
coord = RedisRefreshCoordinator(
lock,
clock=clock,
sleep=_advancing_sleep(clock, step=0.5),
wait_timeout_seconds=1.0,
poll_interval_seconds=0.1,
)
result = await coord.run("u", "s", refresh, reread)
# Gave up waiting (bounded) and returned what's persisted rather than blocking forever.
assert result is not None and result.access_token == "whatever-is-there"