mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
b46733e5b9
commit
f8b07a089f
2 changed files with 209 additions and 0 deletions
|
|
@ -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()
|
||||
|
|
@ -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"
|
||||
Loading…
Add table
Reference in a new issue