From 83df471b058c18a064d02f7d071d8ffe85ca4c9c Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Thu, 25 Jun 2026 22:17:10 -0700 Subject: [PATCH] feat(mcp): inject cache-backend and refresh-coordinator seams (cross-replica token caching) Make CachedOAuthTokenStore's storage and RefreshingTokenStore's single-flight injectable so a cross-replica deployment can back them with Redis without touching the resolver. The defaults preserve today's behavior exactly: InMemoryTokenCacheBackend (the bounded per-process dict) and InProcessRefreshCoordinator (the asyncio single-flight). A distributed deployment injects a shared DualCache-backed backend and a SET NX PX coordinator. invalidate() is now async (the backend may be). The cache stores via the backend with a TTL derived from the token's expiry; the coordinator threads a reread callback for the cross-replica case (losers re-read the persisted token) that the in-process default ignores. --- .../outbound_credentials/oauth_token_store.py | 179 +++++++++++++----- .../test_oauth_token_store.py | 80 +++++++- 2 files changed, 209 insertions(+), 50 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py index e909c99685f..e37aafd6913 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py @@ -16,7 +16,7 @@ from __future__ import annotations import asyncio import time -from collections.abc import Callable +from collections.abc import Awaitable, Callable from dataclasses import dataclass from typing import Protocol @@ -82,6 +82,57 @@ class TokenRefresher(Protocol): ) -> OAuthToken | None: ... +class TokenCacheBackend(Protocol): + """Storage behind ``CachedOAuthTokenStore``: hold a token under ``(user_id, server_id)`` for + ``ttl_seconds``, then forget it. The default ``InMemoryTokenCacheBackend`` is per-process; a + cross-replica deployment injects a shared (Redis) backend so every worker reads one refresh, + matching v1. ``get`` returns ``None`` once the entry's TTL has elapsed. + """ + + async def get(self, user_id: str, server_id: str) -> OAuthToken | None: ... + + async def set( + self, user_id: str, server_id: str, token: OAuthToken, ttl_seconds: float + ) -> None: ... + + async def delete(self, user_id: str, server_id: str) -> None: ... + + +class InMemoryTokenCacheBackend: + """Per-process token cache: a bounded dict with wall-clock TTLs (the default backend).""" + + def __init__( + self, *, max_size: int = 4096, clock: Callable[[], float] = time.time + ) -> None: + self._max_size = max_size + self._clock = clock + self._cache: dict[tuple[str, str], tuple[OAuthToken, float]] = {} + + async def get(self, user_id: str, server_id: str) -> OAuthToken | None: + key = (user_id, server_id) + hit = self._cache.get(key) + if hit is None: + return None + token, valid_until = hit + if self._clock() < valid_until: + return token + self._cache.pop(key, None) + return None + + async def set( + self, user_id: str, server_id: str, token: OAuthToken, ttl_seconds: float + ) -> None: + key = (user_id, server_id) + if key not in self._cache and len(self._cache) >= self._max_size: + # Evict the oldest entry (insertion order), rather than clearing the whole cache and + # forcing every key to re-read the store at once. + self._cache.pop(next(iter(self._cache)), None) + self._cache[key] = (token, self._clock() + ttl_seconds) + + async def delete(self, user_id: str, server_id: str) -> None: + self._cache.pop((user_id, server_id), None) + + class CachedOAuthTokenStore: """Expiry-aware cache over an ``OAuthTokenStore``. Caches positive tokens only. @@ -101,60 +152,102 @@ class CachedOAuthTokenStore: default_ttl_seconds: float, expiry_skew_seconds: float = 60.0, max_size: int = 4096, + backend: TokenCacheBackend | None = None, clock: Callable[[], float] = time.time, ) -> None: self._inner = inner self._default_ttl_seconds = default_ttl_seconds self._expiry_skew_seconds = expiry_skew_seconds - self._max_size = max_size self._clock = clock - self._cache: dict[tuple[str, str], tuple[OAuthToken, float]] = {} + self._backend: TokenCacheBackend = backend or InMemoryTokenCacheBackend( + max_size=max_size, clock=clock + ) - def _valid_until(self, token: OAuthToken) -> float: + def _ttl(self, token: OAuthToken) -> float: if token.expires_at is not None: - return token.expires_at - self._expiry_skew_seconds - return self._clock() + self._default_ttl_seconds + return max( + 0.0, token.expires_at - self._expiry_skew_seconds - self._clock() + ) + return self._default_ttl_seconds async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: - key = (user_id, server_id) - hit = self._cache.get(key) + hit = await self._backend.get(user_id, server_id) if hit is not None: - token, valid_until = hit - if self._clock() < valid_until: - return token + return hit token = await self._inner.fetch(user_id, server_id) if token is None: # Never cache "not authorized": drop any stale entry and re-read on the next call, so # a token stored after the OAuth flow is seen immediately rather than after a TTL. - self._cache.pop(key, None) + await self._backend.delete(user_id, server_id) return token - if key not in self._cache and len(self._cache) >= self._max_size: - # Evict the oldest entry (insertion order) to make room, rather than clearing the - # whole cache and forcing every key to re-read the store at once. - self._cache.pop(next(iter(self._cache)), None) - self._cache[key] = (token, self._valid_until(token)) + await self._backend.set(user_id, server_id, token, self._ttl(token)) return token - def invalidate(self, user_id: str, server_id: str) -> None: + async def invalidate(self, user_id: str, server_id: str) -> None: """Drop a cached entry after the user (re)authorizes or revokes, so a stale token or a stale "not authorized" None cannot mask the change.""" - self._cache.pop((user_id, server_id), None) + await self._backend.delete(user_id, server_id) + + +class RefreshCoordinator(Protocol): + """Ensures one refresh runs per ``(user_id, server_id)`` at a time. Concurrent callers either + share the winner's result (the default ``InProcessRefreshCoordinator``) or, in a cross-replica + coordinator, wait for the holder and ``reread`` the token it persisted - so the IdP sees one + refresh per key across all workers, not one per worker. + """ + + async def run( + self, + user_id: str, + server_id: str, + refresh: Callable[[], Awaitable[OAuthToken | None]], + reread: Callable[[], Awaitable[OAuthToken | None]], + ) -> OAuthToken | None: ... + + +class InProcessRefreshCoordinator: + """Single-flight within one event loop (the default): the first caller per key refreshes while + concurrent callers await the same in-flight task and share its result. ``reread`` is unused here - + the shared task already yields the new token - and exists for the cross-replica coordinator, where + losers re-read the persisted token instead of sharing an in-process future. + """ + + def __init__(self) -> None: + # In-flight refreshes, one task per (user, server); each entry is removed by the task's + # done-callback, so the map is bounded by concurrent refreshes, not by distinct keys seen. + self._inflight: dict[tuple[str, str], asyncio.Future[OAuthToken | None]] = {} + + async def run( + self, + user_id: str, + server_id: str, + refresh: Callable[[], Awaitable[OAuthToken | None]], + reread: Callable[[], Awaitable[OAuthToken | None]], + ) -> OAuthToken | None: + key = (user_id, server_id) + task = self._inflight.get(key) + if task is None: + # The task is detached from the caller, so a cancelled caller does not abort the refresh. + task = asyncio.ensure_future(refresh()) + self._inflight[key] = task + task.add_done_callback(lambda _t, k=key: self._inflight.pop(k, None)) + return await task class RefreshingTokenStore: """An ``OAuthTokenStore`` that proactively refreshes a near-expiry token. Reads from an inner store; if the token is within ``expiry_skew_seconds`` of expiry, it mints a - fresh one via the injected ``TokenRefresher`` under per-(user, server) single-flight: the first - caller refreshes while concurrent callers await the same in-flight future and share its result, - instead of stampeding the IdP. The refresher persists the new token so later requests (and the - surrounding cache) read it without refreshing again. An expired token the refresher cannot renew - (``None``) is surfaced as ``None`` so the arm challenges, never a stale bearer. + fresh one via the injected ``TokenRefresher``, serialized per ``(user, server)`` by the injected + ``RefreshCoordinator`` so callers don't stampede the IdP. The refresher persists the new token so + later requests (and the surrounding cache) read it without refreshing again. An expired token the + refresher cannot renew (``None``) is surfaced as ``None`` so the arm challenges, never a stale + bearer. - Single-flight here is in-process (one event loop). Cross-replica single-flight (Redis SET NX) - and reactive-401 refresh are the later distributed hardening. Composes under - ``CachedOAuthTokenStore`` so the refreshed token is cached until its own expiry. + The default coordinator is in-process; a cross-replica deployment injects a distributed one (Redis + SET NX). Reactive-401 refresh is later hardening (it lives in the egress transport, which sees the + upstream's 401). Composes under ``CachedOAuthTokenStore`` so the refreshed token is cached. """ def __init__( @@ -163,16 +256,16 @@ class RefreshingTokenStore: refresher: TokenRefresher, *, expiry_skew_seconds: float = 60.0, + coordinator: RefreshCoordinator | None = None, clock: Callable[[], float] = time.time, ) -> None: self._inner = inner self._refresher = refresher self._expiry_skew_seconds = expiry_skew_seconds self._clock = clock - # In-flight refreshes, one task per (user, server). Each entry is removed by the task's - # done-callback, so the map is bounded by concurrent refreshes, not by the number of - # distinct users/servers ever seen. - self._inflight: dict[tuple[str, str], asyncio.Future[OAuthToken | None]] = {} + self._coordinator: RefreshCoordinator = ( + coordinator or InProcessRefreshCoordinator() + ) def _is_expired(self, token: OAuthToken) -> bool: return ( @@ -184,21 +277,9 @@ class RefreshingTokenStore: token = await self._inner.fetch(user_id, server_id) if token is None or not self._is_expired(token): return token - return await self._refresh_single_flight(user_id, server_id, token) - - async def _refresh_single_flight( - self, user_id: str, server_id: str, token: OAuthToken - ) -> OAuthToken | None: - key = (user_id, server_id) - task = self._inflight.get(key) - if task is None: - # First caller starts the refresh; concurrent callers await the same task and share its - # result (or exception). The done-callback removes the entry, so the map self-cleans and - # is bounded by in-flight refreshes, not by the number of distinct users/servers. The - # task is detached from the caller, so a cancelled caller does not abort the refresh. - task = asyncio.ensure_future( - self._refresher.refresh(user_id, server_id, token) - ) - self._inflight[key] = task - task.add_done_callback(lambda _t, k=key: self._inflight.pop(k, None)) - return await task + return await self._coordinator.run( + user_id, + server_id, + refresh=lambda: self._refresher.refresh(user_id, server_id, token), + reread=lambda: self._inner.fetch(user_id, server_id), + ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_oauth_token_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_oauth_token_store.py index eb7c31e8dc6..e341f703829 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_oauth_token_store.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_oauth_token_store.py @@ -107,7 +107,7 @@ async def test_invalidate_drops_a_cached_token(): first = await store.fetch("u", "s") assert first is not None and first.access_token == "t1" # cached inner._values[("u", "s")] = OAuthToken(access_token="t2") # rotated - store.invalidate("u", "s") + await store.invalidate("u", "s") second = await store.fetch("u", "s") assert ( second is not None and second.access_token == "t2" @@ -316,3 +316,81 @@ def test_oauth_token_repr_masks_the_secrets(): assert "rt-secret" not in rendered assert "access_token=***" in rendered assert "has_refresh_token=True" in rendered + + +class _RecordingBackend: + """A TokenCacheBackend that records calls, proving CachedOAuthTokenStore delegates storage.""" + + def __init__(self) -> None: + self.sets: List[Tuple[str, str, OAuthToken, float]] = [] + self.deletes: List[Tuple[str, str]] = [] + self._store: Dict[Tuple[str, str], OAuthToken] = {} + + async def get(self, user_id: str, server_id: str) -> Optional[OAuthToken]: + return self._store.get((user_id, server_id)) + + async def set( + self, user_id: str, server_id: str, token: OAuthToken, ttl_seconds: float + ) -> None: + self.sets.append((user_id, server_id, token, ttl_seconds)) + self._store[(user_id, server_id)] = token + + async def delete(self, user_id: str, server_id: str) -> None: + self.deletes.append((user_id, server_id)) + self._store.pop((user_id, server_id), None) + + +async def test_cache_delegates_storage_to_an_injected_backend(): + backend = _RecordingBackend() + inner = _FakeStore({("u", "s"): OAuthToken(access_token="at", expires_at=1100.0)}) + store = CachedOAuthTokenStore( + inner, + default_ttl_seconds=60, + expiry_skew_seconds=30, + backend=backend, + clock=_Clock(1000.0), + ) + + token = await store.fetch("u", "s") + assert token is not None and token.access_token == "at" + # written through to the injected backend, TTL = expires_at - skew - now = 1100 - 30 - 1000 + assert backend.sets == [("u", "s", token, 70.0)] + again = await store.fetch("u", "s") # served by the backend, not the inner store + assert again is not None + assert inner.calls == [("u", "s")] + + +async def test_cache_miss_deletes_from_the_injected_backend(): + backend = _RecordingBackend() + store = CachedOAuthTokenStore( + _FakeStore({}), default_ttl_seconds=60, backend=backend, clock=_Clock() + ) + assert await store.fetch("u", "s") is None + assert backend.deletes == [("u", "s")] # a miss is never cached + + +class _RecordingCoordinator: + """A RefreshCoordinator that records the call and runs the refresh, proving delegation.""" + + def __init__(self) -> None: + self.calls = 0 + + async def run(self, user_id, server_id, refresh, reread): + self.calls += 1 + return await refresh() + + +async def test_refreshing_delegates_single_flight_to_an_injected_coordinator(): + pair = _RefreshablePair(OAuthToken(access_token="old", expires_at=900.0)) + coordinator = _RecordingCoordinator() + store = RefreshingTokenStore( + pair, + pair, + expiry_skew_seconds=30, + coordinator=coordinator, + clock=_Clock(1000.0), + ) + + token = await store.fetch("u", "s") + assert token is not None and token.access_token == "refreshed" + assert coordinator.calls == 1 # the injected coordinator drove the refresh