diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_refresh_coordinator.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_refresh_coordinator.py new file mode 100644 index 00000000000..ef5b137bfc3 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_refresh_coordinator.py @@ -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() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_refresh_coordinator.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_refresh_coordinator.py new file mode 100644 index 00000000000..bf0690f97f8 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_refresh_coordinator.py @@ -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"