feat(mcp): cross-replica single-flight refresh for the v2 per-user OAuth store [2/2] (#31493)

* feat(mcp): encrypt+serialize codec for caching OAuth tokens in Redis (step 1b §1.5)

The serialize+encrypt boundary a cross-replica cache needs: a plaintext bearer in Redis is a leak, so
encode() encrypts (NaCl in prod via the injected encrypt, identity in tests). Caches only access_token
and expires_at, never the refresh_token - the hot path needs just the bearer, and the long-lived
refresh_token stays in the DB (the refresh path is always a cache miss), matching v1. A decoded token
always has refresh_token=None. Undecryptable (key rotation) or corrupt entries read as a miss.

* feat(mcp): DualCache-backed token cache backend (step 1b §1.5)

The cross-replica TokenCacheBackend implementation that plugs into the foundation's
CachedOAuthTokenStore seam: encrypts+serializes the token via the codec and stores it in LiteLLM's
shared DualCache under the same per-(user,server) key v1 used, so workers share one refresh and a
token cached by v1 or v2 is readable by the other across the cutover. Cache and codec are injected;
a non-positive TTL (already-expired token) is not cached, and a missing/corrupt entry reads as a miss.

* 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.

* feat(mcp): Redis SET NX PX distributed lock (step 1b §1.5)

The concrete DistributedLock the RedisRefreshCoordinator elects refreshers with: acquire is an atomic
SET key NX PX ttl (first caller wins, entry self-expires so a crashed holder can't wedge refresh),
release is DEL, is_held is EXISTS. The async Redis client is injected (the client from LiteLLM's
RedisCache in prod), so it is unit-testable with a fake. Any Redis error degrades to not-acquired /
not-held so a cache blip causes an extra refresh, never a crash on the resolve path.

* feat(mcp): wire the cross-replica cache + coordinator into the per-user store (step 1b §1.5)

Upgrade the composition root to use the DualCache-backed cache and SET NX PX refresh
coordinator when Redis is wired, falling back to the foundation's in-process defaults on a
single replica. Layers the cross-replica path on top of the single-replica dispatch store.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>

* fix(mcp): refresh on lock-backend error instead of serving a stale token

The cross-replica refresh coordinator elected refreshers with a boolean acquire:
a Redis transport error was caught and returned as False, which is
indistinguishable from "another worker holds the lock". On a total Redis
outage every worker therefore took the wait-then-reread branch and served the
still-expired token upstream (the upstream then 401s), even though the lock and
coordinator docstrings claimed a Redis blip "degrades to an extra refresh".

Make acquire tristate (LockAcquisition: ACQUIRED / HELD / ERROR) so the
coordinator can tell a busy holder from a dead backend, and refresh anyway on
ERROR. This single-flight lock is a load optimization, not a correctness mutex,
so failing open is correct: it degrades a lock-backend outage to the
no-coordinator behavior (an extra refresh), never a stale bearer.

Add a regression test asserting an acquire error refreshes rather than
re-reading the expired token, and update the docstrings to match.

* style(mcp): wrap redis lock signatures at line-length 88 for CI ruff format

* fix(mcp): a refresh loser surfaces None, not a stale token, when the winner failed

The cross-replica coordinator's losers re-read the token the winner persisted.
If the winner's refresh failed, the store still holds the expired token, so the
loser re-read it and RefreshingTokenStore handed that expired bearer to the
caller (the upstream then 401s) instead of the re-auth challenge the winner
returned via None.

Make the loser's re-read expiry-aware, mirroring refresh_latest_token: a
re-read that is still expired surfaces None so the arm challenges. This only
affects the loser path; the winner's freshly refreshed token is returned
directly by the coordinator and is unaffected.

* fix(mcp): log per-user token decrypt failures at debug, matching v1

When a cached blob cannot be decrypted (e.g. after a salt or master-key rotation) the codec logged a full traceback at error level, since decrypt_value_helper defaults to exception_type=error. v1's MCPPerUserTokenCache passed exception_type=debug on the same path. The blob is ciphertext so this is log noise only, but matching v1 avoids error-level traceback spam on stale entries after a key rotation

* fix(mcp): namespace the refresh lock key and fence its release with a token

The Redis lock wrote its key through the raw client from init_async_client(), bypassing RedisCache's namespace, so two deployments sharing one Redis collided on mcp:refresh_lock:<user>:<server> for any overlapping (user, server) and a colliding deployment skipped the refresh and challenged its own users. The lock now runs every key through an injected namespace_key wired to RedisCache.check_and_fix_namespace, matching the namespace its token cache already uses

release() also deleted the key unconditionally, so a holder whose lock PX-expired and was re-acquired by another worker could delete the new holder's lock and let a third worker run a duplicate refresh, recreating the rotating refresh_token race. acquire now writes a unique per-acquisition token generated by the coordinator and release deletes only when the key still holds that token, via a compare-and-delete Lua script

Adds regression tests: release with a stale token is a no-op while the owner's release deletes; keys are namespaced before reaching Redis; the coordinator acquires and releases with the same token

* fix(mcp): fail open when the per-user token cache delete errors

DualCache swallows get/set errors internally but not delete, and the Redis
layer underneath re-raises through its circuit breaker. So a Redis outage on
the delete() path escaped CachedOAuthTokenStore.fetch()'s unauthorized branch
(which deletes before returning None) and invalidate(), turning a cache blip
into a 500 instead of the v1-style fallback. Catch in the backend so delete
degrades to the TTL-bounded stale entry like get/set already do.

* style(mcp): reformat outbound-credentials files to line-length 120

The merge from staging brought in ruff's line-length 120, but these two
PR-authored files were still wrapped at the old width, so the diff-scoped
ruff format --check in CI flagged them. Pure reformatting; no behavior change.

* fix: harden mcp oauth redis refresh coordination

* fix(mcp): make the per-user token cache backend airtight on boundary failures

get/set now degrade a cache or codec failure to the safe value (miss / no-op)
in the backend itself rather than relying on DualCache and decrypt_value_helper
happening to swallow internally, matching delete() and v1's MCPPerUserTokenCache.
This upholds the layer's boundary-failure-is-a-miss contract regardless of the
injected collaborators, so a Redis outage or an undecryptable entry reads as a
cache miss that re-reads the DB instead of a 500. Adds contract tests for the
cache raising on get/set/delete and the codec raising on encode.

* test(mcp): pin per-user cache get() to a miss when decrypt raises

Greptile's out-of-diff repro had the decrypt reject a blob with ValueError
(bad ciphertext after key rotation); cover that exact raise path, not just the
decrypt-returns-None case, so get() is regression-locked to read it as a miss.

* refactor(mcp): use frozen dataclasses for the trivial DI constructors

Replace the hand-written self._<arg> = arg constructors on OAuthTokenCacheCodec,
RedisRefreshCoordinator, RedisDistributedLock, and DualCacheTokenCacheBackend
with frozen slotted dataclasses, matching the rest of this layer. Fields take the
former parameter names so the constructor API (and the tests' keyword args) are
unchanged; KW_ONLY preserves the keyword-only collaborators.

* fix: serialize lazy per-user oauth store rebuild

* fix(mcp): stop losers challenging mid-refresh by decoupling wait from lease TTL

wait_timeout_seconds defaulted to the same 10s as lock_ttl_seconds, but the
holder renews its lease while a slow token endpoint runs, so a loser waiting
past 10s bailed and re-read the still-expired DB token, challenging the user
even though a valid refresh was in flight. Bound the holder's renewal with a
refresh budget so its lock-hold is finite, and set the loser's wait to outlast
that budget (refresh_budget_seconds + one lease tail) so a loser only re-reads
once the holder has finished or its bounded lease has lapsed, never mid-refresh.

* fix: allow concurrent lazy OAuth fetches without Redis

---------

Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Co-authored-by: Cursor Agent <cursoragent@cursor.com>
This commit is contained in:
tin-berri 2026-06-27 16:27:34 -07:00 • committed by GitHub
parent d515e5bf05
commit 5963b9320f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 1307 additions and 74 deletions

View file

@ -0,0 +1,74 @@
"""Cross-replica ``TokenCacheBackend``: stores the token in LiteLLM's shared ``DualCache``.
Plugs into the foundation's ``CachedOAuthTokenStore`` via the ``TokenCacheBackend`` seam. The token is
encrypted + serialized by the injected codec and written under a per-``(user, server)`` key with the
given TTL, so every worker reads one refresh rather than each re-reading and re-refreshing - matching
v1's ``MCPPerUserTokenCache`` (same NaCl encryption and key, so a token cached by either is readable by
the other across the cutover). A missing or undecryptable entry reads as a miss.
"""
from __future__ import annotations
from dataclasses import KW_ONLY, dataclass
from typing import Protocol
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_codec import (
OAuthTokenCacheCodec,
)
class AsyncCache(Protocol):
"""The slice of LiteLLM's ``DualCache`` this backend needs (Redis-backed, shared across workers)."""
async def async_get_cache(self, key: str) -> object | None: ...
async def async_set_cache(self, key: str, value: str, ttl: float | None = None) -> None: ...
async def async_delete_cache(self, key: str) -> None: ...
@dataclass(frozen=True, slots=True)
class DualCacheTokenCacheBackend:
"""Every method degrades a cache or codec failure to its safe value - ``get`` to a miss
(``None``), ``set``/``delete`` to a no-op - so a Redis outage or an undecryptable entry reads as a
cache miss rather than a request error, matching v1 and this layer's "boundary failure = miss"
contract. The guarantee holds here regardless of whether the injected cache/codec also swallow.
"""
cache: AsyncCache
codec: OAuthTokenCacheCodec
_: KW_ONLY
key_prefix: str = "mcp:per_user_token:"
def _key(self, user_id: str, server_id: str) -> str:
return f"{self.key_prefix}{user_id}:{server_id}"
async def get(self, user_id: str, server_id: str) -> OAuthToken | None:
try:
blob = await self.cache.async_get_cache(self._key(user_id, server_id))
return self.codec.decode(blob) if isinstance(blob, str) else None
except Exception as exc: # noqa: BLE001
verbose_logger.debug("MCP per-user token cache get failed (miss): %s", exc)
return None
async def set(self, user_id: str, server_id: str, token: OAuthToken, ttl_seconds: float) -> None:
if ttl_seconds <= 0:
return
try:
await self.cache.async_set_cache(
self._key(user_id, server_id),
self.codec.encode(token),
ttl=ttl_seconds,
)
except Exception as exc: # noqa: BLE001
verbose_logger.debug("MCP per-user token cache set failed (ignored): %s", exc)
async def delete(self, user_id: str, server_id: str) -> None:
try:
await self.cache.async_delete_cache(self._key(user_id, server_id))
except Exception as exc: # noqa: BLE001
verbose_logger.debug("MCP per-user token cache delete failed (ignored): %s", exc)

View file

@ -272,9 +272,18 @@ class RefreshingTokenStore:
return latest_token
return await self._refresher.refresh(user_id, server_id, latest_token)
async def reread_fresh_token() -> OAuthToken | None:
# A loser re-reads what the winner persisted. If the winner's refresh failed, the store
# still holds the expired token; surface None (-> challenge) like the winner did rather
# than the stale bearer the upstream would 401.
latest_token = await self._inner.fetch(user_id, server_id)
if latest_token is None or self._is_expired(latest_token):
return None
return latest_token
return await self._coordinator.run(
user_id,
server_id,
refresh=refresh_latest_token,
reread=lambda: self._inner.fetch(user_id, server_id),
reread=reread_fresh_token,
)

View file

@ -1,15 +1,16 @@
"""Composition root for the v2-native authorization_code per-user OAuth token store (step 1b).
Assembles ``Cached(Refreshing(V2PerUserTokenStore))`` and replaces ``V1PerUserTokenStore`` in the
resolver. The runtime collaborators (DB, HTTP) are LiteLLM globals not ready at import time, so the
chain is built lazily on first use. The cache and refresh coordinator use the foundation's in-process
defaults (correct for a single replica); the cross-replica path is layered on separately. The DB
read/refresh-grant/persist collaborators acquire their globals per call, mirroring v1's lazy-import
pattern.
resolver. The runtime collaborators (DB, HTTP, the shared cache, Redis) are LiteLLM globals not ready
at import time, so the chain is built lazily on first use. When Redis is wired it uses the
cross-replica path (DualCache-backed cache + ``SET NX PX`` coordinator); otherwise it falls back to
the foundation's in-process defaults (correct for a single replica). The DB read/refresh-grant/persist
collaborators acquire their globals per call, mirroring v1's lazy-import pattern.
"""
from __future__ import annotations
import asyncio
from collections.abc import Callable
from typing import TYPE_CHECKING
@ -17,12 +18,28 @@ from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.outbound_credentials.authz_code_refresher import (
AuthorizationCodeRefresher,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.dual_cache_token_backend import (
AsyncCache,
DualCacheTokenCacheBackend,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
CachedOAuthTokenStore,
OAuthToken,
OAuthTokenStore,
RefreshCoordinator,
RefreshingTokenStore,
TokenCacheBackend,
TokenStoreUnavailable,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_distributed_lock import (
RedisDistributedLock,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_coordinator import (
RedisRefreshCoordinator,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_codec import (
OAuthTokenCacheCodec,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.v2_token_store import (
V2PerUserTokenStore,
)
@ -34,6 +51,7 @@ if TYPE_CHECKING:
_DEFAULT_TTL_SECONDS = 300.0
ServerLookup = Callable[[str], "MCPServer | None"]
StoreBuilder = Callable[[ServerLookup], tuple[OAuthTokenStore, bool]]
async def _read_credential(user_id: str, server_id: str) -> dict[str, object] | None:
@ -97,14 +115,56 @@ async def _post_token_endpoint(url: str, form: dict[str, str]) -> dict[str, obje
return body # pyright: ignore
def _redis_cache_is_available() -> bool:
from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415
return user_api_key_cache.redis_cache is not None
def _runtime_backend_and_coordinator() -> tuple[TokenCacheBackend | None, RefreshCoordinator | None, bool]:
"""The cross-replica cache + coordinator when Redis is wired, else ``(None, None, False)`` so the
foundation's in-process defaults are used (a single replica needs no shared cache or lock).
"""
from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: PLC0415
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415
redis_cache = user_api_key_cache.redis_cache
if redis_cache is None:
return None, None, False
codec = OAuthTokenCacheCodec(
encrypt_value_helper,
lambda blob: decrypt_value_helper(blob, "mcp_per_user_token", exception_type="debug"),
)
# user_api_key_cache satisfies the AsyncCache slice (DualCache types ttl via **kwargs) and the
# Redis client from init_async_client() is partially typed - both are untyped-boundary casts.
cache: AsyncCache = user_api_key_cache # pyright: ignore
redis_client = redis_cache.init_async_client() # pyright: ignore
lock = RedisDistributedLock(
redis_client, # pyright: ignore
namespace_key=redis_cache.check_and_fix_namespace,
)
backend = DualCacheTokenCacheBackend(cache, codec)
coordinator = RedisRefreshCoordinator(lock)
return backend, coordinator, True
def _build_per_user_oauth_token_store(
server_lookup: ServerLookup,
) -> tuple[CachedOAuthTokenStore, bool]:
backend, coordinator, uses_redis = _runtime_backend_and_coordinator()
refresher = AuthorizationCodeRefresher(server_lookup, _post_token_endpoint, _persist_credential)
refreshing = RefreshingTokenStore(V2PerUserTokenStore(_read_credential), refresher, coordinator=coordinator)
return CachedOAuthTokenStore(refreshing, default_ttl_seconds=_DEFAULT_TTL_SECONDS, backend=backend), uses_redis
def build_per_user_oauth_token_store(
server_lookup: ServerLookup,
) -> CachedOAuthTokenStore:
refresher = AuthorizationCodeRefresher(server_lookup, _post_token_endpoint, _persist_credential)
# Cache and refresh coordinator use the foundation's in-process defaults (a single replica needs
# no shared cache or lock); the cross-replica path is layered on separately.
refreshing = RefreshingTokenStore(V2PerUserTokenStore(_read_credential), refresher)
return CachedOAuthTokenStore(refreshing, default_ttl_seconds=_DEFAULT_TTL_SECONDS)
store, _uses_redis = _build_per_user_oauth_token_store(server_lookup)
return store
class LazyPerUserOAuthTokenStore:
@ -112,14 +172,54 @@ class LazyPerUserOAuthTokenStore:
The chain's cache/lock collaborators are LiteLLM runtime globals not available when the resolver
is constructed at import time, so construction is deferred to the first request (by when they are
wired). Built once, then reused.
wired). A no-Redis chain is replaced once Redis becomes available.
"""
def __init__(self, server_lookup: ServerLookup) -> None:
def __init__(
self,
server_lookup: ServerLookup,
*,
store_builder: StoreBuilder = _build_per_user_oauth_token_store,
redis_available: Callable[[], bool] = _redis_cache_is_available,
) -> None:
self._server_lookup = server_lookup
self._store: CachedOAuthTokenStore | None = None
self._store_builder = store_builder
self._redis_available = redis_available
self._store: OAuthTokenStore | None = None
self._uses_redis = False
self._fetch_lock = asyncio.Condition()
self._local_fetches = 0
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
if self._store is None:
self._store = build_per_user_oauth_token_store(self._server_lookup)
return await self._store.fetch(user_id, server_id)
if self._uses_redis:
store = self._store
if store is not None:
return await store.fetch(user_id, server_id)
store, uses_redis = await self._store_for_fetch()
try:
return await store.fetch(user_id, server_id)
finally:
if not uses_redis:
await self._finish_local_fetch()
async def _store_for_fetch(self) -> tuple[OAuthTokenStore, bool]:
async with self._fetch_lock:
while (
self._store is not None and not self._uses_redis and self._redis_available() and self._local_fetches > 0
):
await self._fetch_lock.wait()
store = self._store
if store is None or (not self._uses_redis and self._redis_available()):
store, self._uses_redis = self._store_builder(self._server_lookup)
self._store = store
uses_redis = self._uses_redis
if not uses_redis:
self._local_fetches += 1
return store, uses_redis
async def _finish_local_fetch(self) -> None:
async with self._fetch_lock:
self._local_fetches -= 1
if self._local_fetches == 0:
self._fetch_lock.notify_all()

View file

@ -0,0 +1,94 @@
"""Concrete ``DistributedLock`` over a Redis client: ``SET NX PX`` / owner-only renew / delete.
The cross-replica lock the ``RedisRefreshCoordinator`` elects refreshers with. ``acquire`` is an
atomic ``SET key token NX PX ttl`` (only the first caller wins; the entry self-expires so a crashed
holder can't wedge refresh). ``extend`` renews the lease only when the token still matches, and
``release`` deletes the key only when it still holds this caller's token, so a holder whose lock already
PX-expired and was re-acquired by another worker cannot delete the new holder's lock. ``is_held`` is
``EXISTS``. Every key is run through the injected ``namespace_key`` before it reaches Redis, so lock
keys carry the same namespace as cache keys and cannot collide with another deployment sharing Redis.
The Redis client is injected (in production the async client from LiteLLM's ``RedisCache``), so the
lock is unit-testable with a fake. A transport error on ``acquire`` returns ``LockAcquisition.ERROR`` -
distinct from ``HELD`` - so the coordinator refreshes anyway instead of mistaking a dead backend for a
busy holder; a Redis blip degrades to an extra refresh, never a stale bearer.
"""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import KW_ONLY, dataclass
from typing import Protocol
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_coordinator import (
LockAcquisition,
)
# Delete the key only if it still holds this caller's token, so a holder whose lock already expired
# (PX) and was re-acquired by another worker cannot delete the new holder's lock.
_RELEASE_IF_OWNER = "if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('del', KEYS[1]) else return 0 end"
_EXTEND_IF_OWNER = (
"if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('pexpire', KEYS[1], ARGV[2]) else return 0 end"
)
class RedisCommands(Protocol):
"""The slice of the async Redis client this lock needs."""
async def set(self, name: str, value: str, *, nx: bool = False, px: int | None = None) -> object | None: ...
async def eval(self, script: str, numkeys: int, *keys_and_args: str) -> object: ...
async def exists(self, *names: str) -> int: ...
@dataclass(frozen=True, slots=True)
class RedisDistributedLock:
client: RedisCommands
_: KW_ONLY
namespace_key: Callable[[str], str] = lambda key: key
async def acquire(self, key: str, token: str, ttl_seconds: float) -> LockAcquisition:
try:
result = await self.client.set(self.namespace_key(key), token, nx=True, px=int(ttl_seconds * 1000))
# Degrade on any Redis client error: redis.exceptions narrows only via an import that
# is Unknown under basedpyright, and the lock must never crash the resolve path.
except Exception as exc: # noqa: BLE001
verbose_logger.warning("RedisDistributedLock.acquire failed: %s", exc)
return LockAcquisition.ERROR
return LockAcquisition.ACQUIRED if result is not None else LockAcquisition.HELD
async def extend(self, key: str, token: str, ttl_seconds: float) -> bool:
try:
result = await self.client.eval(
_EXTEND_IF_OWNER,
1,
self.namespace_key(key),
token,
str(int(ttl_seconds * 1000)),
)
# Degrade on any Redis client error: redis.exceptions narrows only via an import that
# is Unknown under basedpyright, and the lock must never crash the resolve path.
except Exception as exc: # noqa: BLE001
verbose_logger.warning("RedisDistributedLock.extend failed: %s", exc)
return False
return result == 1
async def release(self, key: str, token: str) -> None:
try:
await self.client.eval(_RELEASE_IF_OWNER, 1, self.namespace_key(key), token)
# Degrade on any Redis client error: redis.exceptions narrows only via an import that
# is Unknown under basedpyright, and the lock must never crash the resolve path.
except Exception as exc: # noqa: BLE001
verbose_logger.warning("RedisDistributedLock.release failed: %s", exc)
async def is_held(self, key: str) -> bool:
try:
return await self.client.exists(self.namespace_key(key)) > 0
# Degrade on any Redis client error: redis.exceptions narrows only via an import that
# is Unknown under basedpyright, and the lock must never crash the resolve path.
except Exception as exc: # noqa: BLE001
# On error, report "not held" so a waiter stops waiting and re-reads rather than blocking.
verbose_logger.warning("RedisDistributedLock.is_held failed: %s", exc)
return False

View file

@ -0,0 +1,136 @@
"""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 holder
renews the ``PX`` lease while refresh runs (up to a refresh budget, so a hung endpoint can't hold the
lock forever), and a loser waits longer than that budget - so a loser only re-reads once the holder has
finished or its bounded lease has lapsed, never mid-refresh, 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 wrapper in production, a
fake in tests).
The lock is a single-flight optimization, not a correctness mutex, so it fails open: when the lock
backend is unreachable, ``acquire`` reports ``ERROR`` (distinct from ``HELD``) and this coordinator
refreshes anyway rather than wait on a holder that may not exist and then serve a still-expired token.
That degrades a Redis outage to the no-coordinator behavior (each worker may refresh), never a stale
bearer the upstream would 401.
"""
from __future__ import annotations
import asyncio
import time
import uuid
from collections.abc import Awaitable, Callable
from contextlib import suppress
from dataclasses import KW_ONLY, dataclass
from enum import Enum
from typing import Protocol
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
class LockAcquisition(Enum):
"""Outcome of a best-effort ``acquire``. ``ERROR`` is kept distinct from ``HELD`` so a caller can
tell "someone else is refreshing" (wait and re-read) from "the lock backend is down" (no election
happened, so refresh anyway) instead of conflating both into a single ``False``."""
ACQUIRED = "acquired" # won the election; this worker refreshes
HELD = "held" # another worker holds it; wait then re-read
ERROR = "error" # lock backend unreachable; holder unknown, so refresh anyway
class DistributedLock(Protocol):
"""A best-effort cross-replica lock. ``acquire`` is ``SET key token NX PX ttl`` reported as a
``LockAcquisition`` (won / held by another / backend error); ``release`` deletes the key only if
it still holds this caller's ``token`` (so it cannot delete a lock another worker re-acquired
after PX-expiry); ``extend`` refreshes the ``PX`` lease only for the owner; ``is_held`` is
``EXISTS`` (so a waiter can poll without taking the lock)."""
async def acquire(self, key: str, token: str, ttl_seconds: float) -> LockAcquisition: ...
async def extend(self, key: str, token: str, ttl_seconds: float) -> bool: ...
async def release(self, key: str, token: str) -> None: ...
async def is_held(self, key: str) -> bool: ...
@dataclass(frozen=True, slots=True)
class RedisRefreshCoordinator:
lock: DistributedLock
_: KW_ONLY
key_prefix: str = "mcp:refresh_lock:"
lock_ttl_seconds: float = 10.0
# The holder renews its lease while a slow token endpoint runs, but only up to this budget; past it
# it stops renewing and the lock lapses, so a hung refresh degrades to "maybe an extra refresh"
# rather than holding every loser behind it indefinitely.
refresh_budget_seconds: float = 20.0
# How long a loser waits for the holder before giving up and re-reading. It MUST outlast the
# holder's max lock-hold (refresh_budget_seconds + one lock_ttl_seconds tail); otherwise a loser
# bails while the holder is still legitimately refreshing, re-reads the still-expired token, and
# challenges the user mid-refresh.
wait_timeout_seconds: float = 35.0
poll_interval_seconds: float = 0.05
sleep: Callable[[float], Awaitable[None]] = asyncio.sleep
clock: Callable[[], float] = time.monotonic
new_token: Callable[[], str] = lambda: uuid.uuid4().hex
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)
token = self.new_token()
match await self.lock.acquire(key, token, self.lock_ttl_seconds):
case LockAcquisition.ACQUIRED:
return await self._refresh_with_lease_renewal(key, token, refresh)
case LockAcquisition.ERROR:
# No election happened (lock backend down), so waiting would just re-read the
# still-expired token. Refresh anyway; worst case is an extra refresh, not a stale bearer.
return await refresh()
case LockAcquisition.HELD:
# 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()
async def _refresh_with_lease_renewal(
self,
key: str,
token: str,
refresh: Callable[[], Awaitable[OAuthToken | None]],
) -> OAuthToken | None:
refresh_task = asyncio.ensure_future(refresh())
renewal_task = asyncio.create_task(self._renew_lease_until_done(key, token, refresh_task))
try:
return await refresh_task
finally:
renewal_task.cancel()
with suppress(asyncio.CancelledError):
await renewal_task
await self.lock.release(key, token)
async def _renew_lease_until_done(
self,
key: str,
token: str,
refresh_task: asyncio.Future[OAuthToken | None],
) -> None:
budget_deadline = self.clock() + self.refresh_budget_seconds
while not refresh_task.done() and self.clock() < budget_deadline:
await self.sleep(self.lock_ttl_seconds / 2)
if not refresh_task.done() and not await self.lock.extend(key, token, self.lock_ttl_seconds):
return

View file

@ -0,0 +1,34 @@
"""Serialize + encrypt boundary for caching an OAuth token in a shared (Redis) cache.
A cross-replica cache must serialize the token, and a plaintext bearer in Redis is a leak, so this
encrypts the value (NaCl in production via the injected ``encrypt``, identity in tests). It caches
**only** the ``access_token``: the hot path needs just the bearer, expiry is carried by the cache
entry's TTL (set from the token's ``expires_at`` by the cache), and the long-lived refresh_token stays
in the DB - the refresh path is always a cache miss that re-reads it - so it never reaches Redis. A
decoded token therefore carries only the bearer (``expires_at`` and ``refresh_token`` both None); the
TTL, not the value, bounds its life. An empty/undecryptable blob (e.g. master-key rotation) is a miss.
"""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
@dataclass(frozen=True, slots=True)
class OAuthTokenCacheCodec:
encrypt: Callable[[str], str]
decrypt: Callable[[str], str | None]
def encode(self, token: OAuthToken) -> str:
return self.encrypt(token.access_token)
def decode(self, blob: str) -> OAuthToken | None:
access_token = self.decrypt(blob)
if not access_token:
return None
return OAuthToken(access_token=access_token, refresh_token=None)

View file

@ -0,0 +1,156 @@
"""Tests for the DualCache-backed token cache backend: encrypted round-trip, key, TTL, miss."""
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.dual_cache_token_backend import (
DualCacheTokenCacheBackend,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_codec import (
OAuthTokenCacheCodec,
)
class _FakeCache:
def __init__(self):
self.values = {}
self.ttls = {}
async def async_get_cache(self, key):
return self.values.get(key)
async def async_set_cache(self, key, value, ttl=None):
self.values[key] = value
self.ttls[key] = ttl
async def async_delete_cache(self, key):
self.values.pop(key, None)
def _backend(cache):
codec = OAuthTokenCacheCodec(
encrypt=lambda s: f"enc:{s}",
decrypt=lambda b: b[4:] if b.startswith("enc:") else None,
)
return DualCacheTokenCacheBackend(cache, codec)
@pytest.mark.asyncio
async def test_set_then_get_round_trips_encrypted_under_the_per_user_key():
cache = _FakeCache()
backend = _backend(cache)
await backend.set("alice", "srv", OAuthToken(access_token="at"), 120.0)
key = "mcp:per_user_token:alice:srv"
assert cache.ttls[key] == 120.0
# Stored via the codec (encrypted), not the bare token; the codec's own test proves real
# NaCl output hides the secret - here the fake encrypt just wraps, so we check it was applied.
assert cache.values[key] == "enc:at"
got = await backend.get("alice", "srv")
assert got is not None and got.access_token == "at"
@pytest.mark.asyncio
async def test_get_missing_key_is_none():
assert await _backend(_FakeCache()).get("alice", "srv") is None
@pytest.mark.asyncio
async def test_non_str_cache_value_is_a_miss():
cache = _FakeCache()
cache.values["mcp:per_user_token:alice:srv"] = 12345 # corrupt / wrong type
assert await _backend(cache).get("alice", "srv") is None
@pytest.mark.asyncio
async def test_undecryptable_blob_is_a_miss():
# e.g. a master-key rotation leaves an entry the codec can't decrypt; it must read as a miss so
# the store re-reads the DB, not raise (the decrypt helper here returns None for unknown blobs).
cache = _FakeCache()
cache.values["mcp:per_user_token:alice:srv"] = "not-our-ciphertext"
assert await _backend(cache).get("alice", "srv") is None
@pytest.mark.asyncio
async def test_non_positive_ttl_is_not_written():
cache = _FakeCache()
await _backend(cache).set("alice", "srv", OAuthToken(access_token="at"), 0.0)
assert cache.values == {} # an already-expired token is not cached
@pytest.mark.asyncio
async def test_delete_removes_the_entry():
cache = _FakeCache()
backend = _backend(cache)
await backend.set("alice", "srv", OAuthToken(access_token="at"), 60.0)
await backend.delete("alice", "srv")
assert await backend.get("alice", "srv") is None
@pytest.mark.asyncio
async def test_keys_isolate_users_and_servers():
cache = _FakeCache()
backend = _backend(cache)
await backend.set("alice", "srv", OAuthToken(access_token="a"), 60.0)
await backend.set("bob", "srv", OAuthToken(access_token="b"), 60.0)
alice = await backend.get("alice", "srv")
bob = await backend.get("bob", "srv")
assert alice is not None and alice.access_token == "a"
assert bob is not None and bob.access_token == "b"
class _RaisingCache(_FakeCache):
"""A cache whose every op raises, e.g. a Redis outage tripping the circuit breaker."""
async def async_get_cache(self, key):
raise ConnectionError("redis down")
async def async_set_cache(self, key, value, ttl=None):
raise ConnectionError("redis down")
async def async_delete_cache(self, key):
raise ConnectionError("redis down")
# A cache or codec failure must degrade to the safe value (miss / no-op), never propagate: otherwise a
# Redis outage turns CachedOAuthTokenStore.fetch() into a 500 instead of a cache miss that re-reads the
# DB (get/set) or issues the OAuth challenge (the unauthorized branch deletes before returning None).
@pytest.mark.asyncio
async def test_get_is_a_miss_when_the_cache_raises():
assert await _backend(_RaisingCache()).get("alice", "srv") is None
@pytest.mark.asyncio
async def test_set_is_swallowed_when_the_cache_raises():
await _backend(_RaisingCache()).set("alice", "srv", OAuthToken(access_token="at"), 60.0)
@pytest.mark.asyncio
async def test_delete_is_swallowed_when_the_cache_raises():
await _backend(_RaisingCache()).delete("alice", "srv")
@pytest.mark.asyncio
async def test_set_is_swallowed_when_the_codec_raises():
def _boom(_: str) -> str:
raise ValueError("encrypt unavailable")
codec = OAuthTokenCacheCodec(encrypt=_boom, decrypt=lambda b: None)
backend = DualCacheTokenCacheBackend(_FakeCache(), codec)
await backend.set("alice", "srv", OAuthToken(access_token="at"), 60.0) # must not raise
@pytest.mark.asyncio
async def test_get_is_a_miss_when_decrypt_raises():
# A blob the decrypt rejects with an exception (e.g. bad ciphertext after key rotation) must read
# as a miss, not propagate the error out of get() and abort fetch().
def _boom(_: str) -> str | None:
raise ValueError("bad ciphertext")
codec = OAuthTokenCacheCodec(encrypt=lambda s: f"enc:{s}", decrypt=_boom)
cache = _FakeCache()
cache.values["mcp:per_user_token:alice:srv"] = "whatever"
assert await DualCacheTokenCacheBackend(cache, codec).get("alice", "srv") is None

View file

@ -37,9 +37,7 @@ async def test_serves_token_until_its_expiry():
token = OAuthToken(access_token="at", expires_at=1100.0)
inner = _FakeStore({("u", "s"): token})
clock = _Clock(1000.0)
store = CachedOAuthTokenStore(
inner, default_ttl_seconds=60, expiry_skew_seconds=30, clock=clock
)
store = CachedOAuthTokenStore(inner, default_ttl_seconds=60, expiry_skew_seconds=30, clock=clock)
assert await store.fetch("u", "s") is token
clock.t = 1060.0 # still before expiry - skew (1100 - 30 = 1070)
@ -51,9 +49,7 @@ async def test_refetches_once_token_has_expired():
token = OAuthToken(access_token="at", expires_at=1100.0)
inner = _FakeStore({("u", "s"): token})
clock = _Clock(1000.0)
store = CachedOAuthTokenStore(
inner, default_ttl_seconds=60, expiry_skew_seconds=30, clock=clock
)
store = CachedOAuthTokenStore(inner, default_ttl_seconds=60, expiry_skew_seconds=30, clock=clock)
await store.fetch("u", "s")
clock.t = 1080.0 # past expiry - skew (1070)
@ -109,9 +105,7 @@ async def test_invalidate_drops_a_cached_token():
inner._values[("u", "s")] = OAuthToken(access_token="t2") # rotated
await store.invalidate("u", "s")
second = await store.fetch("u", "s")
assert (
second is not None and second.access_token == "t2"
) # re-read after invalidate
assert second is not None and second.access_token == "t2" # re-read after invalidate
async def test_store_unavailable_is_not_cached():
@ -152,9 +146,7 @@ async def test_bounded_cache_evicts_oldest_not_everything():
("u3", "s"): OAuthToken(access_token="k3"),
}
)
store = CachedOAuthTokenStore(
inner, default_ttl_seconds=60, max_size=2, clock=_Clock()
)
store = CachedOAuthTokenStore(inner, default_ttl_seconds=60, max_size=2, clock=_Clock())
await store.fetch("u1", "s")
await store.fetch("u2", "s")
@ -162,9 +154,7 @@ async def test_bounded_cache_evicts_oldest_not_everything():
await store.fetch("u2", "s") # still cached
await store.fetch("u1", "s") # was evicted -> re-read
assert (
inner.calls.count(("u2", "s")) == 1
) # only the oldest was evicted, not everything
assert inner.calls.count(("u2", "s")) == 1 # only the oldest was evicted, not everything
assert inner.calls.count(("u1", "s")) == 2
@ -182,23 +172,17 @@ class _RefreshablePair:
self.fetch_calls += 1
return self._current
async def refresh(
self, user_id: str, server_id: str, token: OAuthToken
) -> Optional[OAuthToken]:
async def refresh(self, user_id: str, server_id: str, token: OAuthToken) -> Optional[OAuthToken]:
self.refresh_calls += 1
self.refresh_args.append((user_id, server_id))
await asyncio.sleep(
0
) # yield so other concurrent callers reach the lock and wait
await asyncio.sleep(0) # yield so other concurrent callers reach the lock and wait
self._current = OAuthToken(access_token="refreshed", expires_at=9999.0)
return self._current
async def test_refreshing_passes_through_a_fresh_token():
pair = _RefreshablePair(OAuthToken(access_token="ok", expires_at=9999.0))
store = RefreshingTokenStore(
pair, pair, expiry_skew_seconds=30, clock=_Clock(1000.0)
)
store = RefreshingTokenStore(pair, pair, expiry_skew_seconds=30, clock=_Clock(1000.0))
token = await store.fetch("u", "s")
assert token is not None and token.access_token == "ok"
@ -207,16 +191,12 @@ async def test_refreshing_passes_through_a_fresh_token():
async def test_refreshing_mints_a_fresh_token_when_expired():
pair = _RefreshablePair(OAuthToken(access_token="old", expires_at=900.0))
store = RefreshingTokenStore(
pair, pair, expiry_skew_seconds=30, clock=_Clock(1000.0)
)
store = RefreshingTokenStore(pair, pair, expiry_skew_seconds=30, clock=_Clock(1000.0))
token = await store.fetch("u", "s")
assert token is not None and token.access_token == "refreshed"
assert pair.refresh_calls == 1
assert pair.refresh_args == [
("u", "s")
] # the seam threads the grant/persist key through
assert pair.refresh_args == [("u", "s")] # the seam threads the grant/persist key through
async def test_refreshing_returns_none_when_it_cannot_refresh():
@ -224,9 +204,7 @@ async def test_refreshing_returns_none_when_it_cannot_refresh():
async def fetch(self, user_id: str, server_id: str) -> Optional[OAuthToken]:
return OAuthToken(access_token="old", expires_at=900.0)
async def refresh(
self, user_id: str, server_id: str, token: OAuthToken
) -> Optional[OAuthToken]:
async def refresh(self, user_id: str, server_id: str, token: OAuthToken) -> Optional[OAuthToken]:
return None # e.g. no refresh_token
src = _NoRefresh()
@ -237,9 +215,7 @@ async def test_refreshing_returns_none_when_it_cannot_refresh():
async def test_refreshing_is_single_flight_under_concurrency():
pair = _RefreshablePair(OAuthToken(access_token="old", expires_at=900.0))
store = RefreshingTokenStore(
pair, pair, expiry_skew_seconds=30, clock=_Clock(1000.0)
)
store = RefreshingTokenStore(pair, pair, expiry_skew_seconds=30, clock=_Clock(1000.0))
results = await asyncio.gather(*[store.fetch("u", "s") for _ in range(5)])
assert pair.refresh_calls == 1 # one refresh shared across 5 concurrent callers
@ -265,9 +241,7 @@ class _StaleReadRacePair:
return self._old_token
return self._current
async def refresh(
self, user_id: str, server_id: str, token: OAuthToken
) -> Optional[OAuthToken]:
async def refresh(self, user_id: str, server_id: str, token: OAuthToken) -> Optional[OAuthToken]:
self.refresh_calls += 1
self.refresh_started.set()
await self.finish_refresh.wait()
@ -277,9 +251,7 @@ class _StaleReadRacePair:
async def test_stale_read_after_refresh_rereads_before_starting_new_refresh():
pair = _StaleReadRacePair()
store = RefreshingTokenStore(
pair, pair, expiry_skew_seconds=30, clock=_Clock(1000.0)
)
store = RefreshingTokenStore(pair, pair, expiry_skew_seconds=30, clock=_Clock(1000.0))
first = asyncio.create_task(store.fetch("u", "s"))
await pair.refresh_started.wait()
@ -304,9 +276,7 @@ async def test_refresh_failure_is_shared_by_joiners_not_re_run():
async def fetch(self, user_id: str, server_id: str) -> Optional[OAuthToken]:
return OAuthToken(access_token="old", expires_at=900.0)
async def refresh(
self, user_id: str, server_id: str, token: OAuthToken
) -> Optional[OAuthToken]:
async def refresh(self, user_id: str, server_id: str, token: OAuthToken) -> Optional[OAuthToken]:
self.calls += 1
await asyncio.sleep(0) # let the concurrent callers join the same task
raise RuntimeError("refresh boom")
@ -314,9 +284,7 @@ async def test_refresh_failure_is_shared_by_joiners_not_re_run():
src = _FailingRefresher()
store = RefreshingTokenStore(src, src, expiry_skew_seconds=30, clock=_Clock(1000.0))
results = await asyncio.gather(
*[store.fetch("u", "s") for _ in range(3)], return_exceptions=True
)
results = await asyncio.gather(*[store.fetch("u", "s") for _ in range(3)], return_exceptions=True)
assert src.calls == 1 # single-flight: one attempt, the failure is shared
assert all(isinstance(r, RuntimeError) for r in results)
@ -358,9 +326,7 @@ async def test_refreshing_default_skew_is_60_seconds():
def test_oauth_token_repr_masks_the_secrets():
token = OAuthToken(
access_token="super-secret", expires_at=123.0, refresh_token="rt-secret"
)
token = OAuthToken(access_token="super-secret", expires_at=123.0, refresh_token="rt-secret")
rendered = repr(token)
assert "super-secret" not in rendered
assert "rt-secret" not in rendered
@ -379,9 +345,7 @@ class _RecordingBackend:
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:
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
@ -412,9 +376,7 @@ async def test_cache_delegates_storage_to_an_injected_backend():
async def test_cache_miss_deletes_from_the_injected_backend():
backend = _RecordingBackend()
store = CachedOAuthTokenStore(
_FakeStore({}), default_ttl_seconds=60, backend=backend, clock=_Clock()
)
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
@ -444,3 +406,61 @@ async def test_refreshing_delegates_single_flight_to_an_injected_coordinator():
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
class _LoserCoordinator:
"""Simulates losing the cross-replica election: the loser only re-reads what the winner persisted,
it never refreshes itself."""
async def run(self, user_id, server_id, refresh, reread):
return await reread()
async def test_loser_surfaces_none_not_a_stale_token_when_the_winner_refresh_failed():
# The winner's refresh failed, so the store still holds the expired token. A loser must surface
# None (-> re-auth challenge) like the winner did, never the still-expired bearer (the upstream
# would 401 it). _RefreshablePair only updates on refresh, so a loser that never refreshes keeps
# re-reading the expired token.
pair = _RefreshablePair(OAuthToken(access_token="expired", expires_at=900.0))
store = RefreshingTokenStore(
pair,
pair,
expiry_skew_seconds=30,
coordinator=_LoserCoordinator(),
clock=_Clock(1000.0),
)
assert await store.fetch("u", "s") is None
assert pair.refresh_calls == 0 # the loser never refreshes; it only re-reads
async def test_loser_rereads_the_fresh_token_a_successful_winner_persisted():
class _ExpiredThenWinnerPersisted:
# First read sees the expired token (triggers the refresh path); the re-read sees the fresh
# token a winner persisted in between.
def __init__(self) -> None:
self._reads = 0
self.refresh_calls = 0
async def fetch(self, user_id: str, server_id: str) -> Optional[OAuthToken]:
self._reads += 1
if self._reads == 1:
return OAuthToken(access_token="expired", expires_at=900.0)
return OAuthToken(access_token="winner-fresh", expires_at=9999.0)
async def refresh(self, user_id: str, server_id: str, token: OAuthToken) -> Optional[OAuthToken]:
self.refresh_calls += 1
raise AssertionError("a loser must not refresh")
src = _ExpiredThenWinnerPersisted()
store = RefreshingTokenStore(
src,
src,
expiry_skew_seconds=30,
coordinator=_LoserCoordinator(),
clock=_Clock(1000.0),
)
token = await store.fetch("u", "s")
assert token is not None and token.access_token == "winner-fresh"
assert src.refresh_calls == 0

View file

@ -0,0 +1,160 @@
import asyncio
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
OAuthTokenStore,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
ServerLookup,
)
class _RecordingStore:
def __init__(self, access_token: str) -> None:
self._access_token = access_token
self.calls: list[tuple[str, str]] = []
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
self.calls.append((user_id, server_id))
return OAuthToken(access_token=self._access_token)
class _BlockingStore:
def __init__(self, access_token: str) -> None:
self._access_token = access_token
self.started = asyncio.Event()
self.release = asyncio.Event()
self.calls: list[tuple[str, str]] = []
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
self.calls.append((user_id, server_id))
self.started.set()
await self.release.wait()
return OAuthToken(access_token=self._access_token)
class _RedisAvailability:
def __init__(self) -> None:
self.available = False
def __call__(self) -> bool:
return self.available
async def _wait_for_call_count(store: _BlockingStore, count: int) -> None:
for _ in range(100):
if len(store.calls) >= count:
return
await asyncio.sleep(0)
raise AssertionError(f"expected {count} calls, saw {len(store.calls)}")
@pytest.mark.asyncio
async def test_lazy_store_rebuilds_when_redis_becomes_available() -> None:
local_store = _RecordingStore("local")
redis_store = _RecordingStore("redis")
redis_available = _RedisAvailability()
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
if redis_available.available:
return redis_store, True
return local_store, False
def server_lookup(_server_id: str) -> None:
return None
store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=redis_available,
)
first = await store.fetch("u", "s")
redis_available.available = True
second = await store.fetch("u", "s")
third = await store.fetch("u", "s")
assert first is not None and first.access_token == "local"
assert second is not None and second.access_token == "redis"
assert third is not None and third.access_token == "redis"
assert build_calls == 2
assert local_store.calls == [("u", "s")]
assert redis_store.calls == [("u", "s"), ("u", "s")]
@pytest.mark.asyncio
async def test_lazy_store_allows_concurrent_local_fetches_without_redis() -> None:
local_store = _BlockingStore("local")
redis_available = _RedisAvailability()
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return local_store, False
def server_lookup(_server_id: str) -> None:
return None
store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=redis_available,
)
first_fetch = asyncio.create_task(store.fetch("u1", "s1"))
second_fetch = asyncio.create_task(store.fetch("u2", "s2"))
await asyncio.wait_for(_wait_for_call_count(local_store, 2), timeout=1)
local_store.release.set()
first, second = await asyncio.gather(first_fetch, second_fetch)
assert first is not None and first.access_token == "local"
assert second is not None and second.access_token == "local"
assert build_calls == 1
assert local_store.calls == [("u1", "s1"), ("u2", "s2")]
@pytest.mark.asyncio
async def test_lazy_store_waits_for_in_flight_local_fetch_before_redis_rebuild() -> None:
local_store = _BlockingStore("local")
redis_store = _RecordingStore("redis")
redis_available = _RedisAvailability()
def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
if redis_available.available:
return redis_store, True
return local_store, False
def server_lookup(_server_id: str) -> None:
return None
store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=redis_available,
)
first_fetch = asyncio.create_task(store.fetch("u", "s"))
await local_store.started.wait()
redis_available.available = True
second_fetch = asyncio.create_task(store.fetch("u", "s"))
await asyncio.sleep(0)
assert redis_store.calls == []
local_store.release.set()
first = await first_fetch
second = await second_fetch
assert first is not None and first.access_token == "local"
assert second is not None and second.access_token == "redis"
assert local_store.calls == [("u", "s")]
assert redis_store.calls == [("u", "s")]

View file

@ -0,0 +1,126 @@
"""Tests for the Redis lock: acquire NX/PX with a token, compare-and-delete release, namespacing."""
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_distributed_lock import (
RedisDistributedLock,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_coordinator import (
LockAcquisition,
)
class _FakeRedis:
"""Models just enough Redis to exercise SET NX (token store) and the compare-and-delete EVAL."""
def __init__(self, set_returns=True, exists_returns=1, raise_on=()):
self._set_returns = set_returns
self._exists_returns = exists_returns
self._raise_on = set(raise_on)
self.set_calls = []
self.eval_calls = []
self.expire_calls = []
self.deleted = []
self.store: dict = {}
async def set(self, name, value, *, nx=False, px=None):
if "set" in self._raise_on:
raise RuntimeError("redis down")
self.set_calls.append((name, value, nx, px))
if self._set_returns:
self.store[name] = value
return self._set_returns
async def eval(self, script, numkeys, *keys_and_args):
if "eval" in self._raise_on:
raise RuntimeError("redis down")
self.eval_calls.append((numkeys, keys_and_args))
key, token = keys_and_args[0], keys_and_args[1]
if len(keys_and_args) == 3:
ttl_ms = keys_and_args[2]
if self.store.get(key) == token:
self.expire_calls.append((key, ttl_ms))
return 1
return 0
if self.store.get(key) == token: # compare-and-delete: only the owner deletes
del self.store[key]
self.deleted.append(key)
return 1
return 0
async def exists(self, *names):
if "exists" in self._raise_on:
raise RuntimeError("redis down")
return self._exists_returns
@pytest.mark.asyncio
async def test_acquire_sets_token_with_nx_px_and_reports_acquired():
redis = _FakeRedis(set_returns=True)
assert await RedisDistributedLock(redis).acquire("k", "tok-1", 10.0) is (LockAcquisition.ACQUIRED)
assert redis.set_calls == [("k", "tok-1", True, 10000)] # token value, NX, px in ms
@pytest.mark.asyncio
async def test_acquire_reports_held_when_key_already_held():
# redis SET NX returns None when the key exists -> another worker holds it.
assert await RedisDistributedLock(_FakeRedis(set_returns=None)).acquire("k", "tok", 10.0) is LockAcquisition.HELD
@pytest.mark.asyncio
async def test_acquire_reports_error_on_redis_error_distinct_from_held():
# A dead backend must be distinguishable from a busy holder so the coordinator refreshes anyway.
assert await RedisDistributedLock(_FakeRedis(raise_on=["set"])).acquire("k", "tok", 10.0) is LockAcquisition.ERROR
@pytest.mark.asyncio
async def test_release_deletes_only_when_the_token_matches():
# Regression: a holder whose lock PX-expired and was re-acquired by another worker must not be
# able to delete the new holder's lock. release with a stale token is a no-op.
redis = _FakeRedis()
lock = RedisDistributedLock(redis)
await lock.acquire("k", "owner-B", 10.0) # B currently holds the lock
await lock.release("k", "owner-A") # A's stale token
assert redis.deleted == [] and redis.store.get("k") == "owner-B" # B's lock survives
await lock.release("k", "owner-B") # the real owner releases
assert redis.deleted == ["k"] and "k" not in redis.store
@pytest.mark.asyncio
async def test_keys_are_namespaced_before_reaching_redis():
redis = _FakeRedis()
lock = RedisDistributedLock(redis, namespace_key=lambda key: f"ns:{key}")
await lock.acquire("k", "tok", 10.0)
await lock.extend("k", "tok", 10.0)
await lock.release("k", "tok")
await lock.is_held("k")
assert redis.set_calls[0][0] == "ns:k" # acquire namespaced
assert redis.eval_calls[0][1][0] == "ns:k" # extend (EVAL KEYS[1]) namespaced
assert redis.eval_calls[1][1][0] == "ns:k" # release (EVAL KEYS[1]) namespaced
assert redis.deleted == ["ns:k"]
@pytest.mark.asyncio
async def test_extend_refreshes_ttl_only_when_the_token_matches():
redis = _FakeRedis()
lock = RedisDistributedLock(redis)
await lock.acquire("k", "owner-B", 10.0)
assert await lock.extend("k", "owner-A", 10.0) is False
assert await lock.extend("k", "owner-B", 10.0) is True
assert redis.expire_calls == [("k", "10000")]
@pytest.mark.asyncio
async def test_extend_degrades_to_false_on_redis_error():
assert await RedisDistributedLock(_FakeRedis(raise_on=["eval"])).extend("k", "tok", 10.0) is False
@pytest.mark.asyncio
async def test_is_held_reflects_exists():
assert await RedisDistributedLock(_FakeRedis(exists_returns=1)).is_held("k") is True
assert await RedisDistributedLock(_FakeRedis(exists_returns=0)).is_held("k") is False
@pytest.mark.asyncio
async def test_is_held_degrades_to_false_on_redis_error():
assert await RedisDistributedLock(_FakeRedis(raise_on=["exists"])).is_held("k") is False

View file

@ -0,0 +1,270 @@
"""Tests for the cross-replica refresh coordinator: winner refreshes, losers wait then re-read."""
import asyncio
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 (
LockAcquisition,
RedisRefreshCoordinator,
)
_KEY = "mcp:refresh_lock:u:s"
class _FakeLock:
def __init__(self, acquired, held_sequence=()):
self._acquired = acquired
self._held = list(held_sequence)
self.acquire_calls = []
self.extend_calls = []
self.released = []
self.extended = asyncio.Event()
async def acquire(self, key, token, ttl_seconds):
self.acquire_calls.append((key, token, ttl_seconds))
return self._acquired
async def release(self, key, token):
self.released.append((key, token))
async def extend(self, key, token, ttl_seconds):
self.extend_calls.append((key, token, ttl_seconds))
self.extended.set()
return True
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_acquires_and_releases_with_the_same_token_and_never_rereads():
lock = _FakeLock(acquired=LockAcquisition.ACQUIRED)
refreshed = OAuthToken(access_token="new")
reread_calls = []
async def refresh():
return refreshed
async def reread():
reread_calls.append(1)
return None
coord = RedisRefreshCoordinator(lock, new_token=lambda: "tok")
result = await coord.run("u", "s", refresh, reread)
assert result is refreshed
assert lock.acquire_calls == [(_KEY, "tok", 10.0)]
# released with the SAME token it acquired with, so it can only delete its own lock
assert lock.released == [(_KEY, "tok")]
assert reread_calls == []
@pytest.mark.asyncio
async def test_winner_releases_even_when_refresh_raises():
lock = _FakeLock(acquired=LockAcquisition.ACQUIRED)
async def refresh():
raise RuntimeError("boom")
async def reread():
return None
coord = RedisRefreshCoordinator(lock, new_token=lambda: "tok")
with pytest.raises(RuntimeError):
await coord.run("u", "s", refresh, reread)
assert lock.released == [(_KEY, "tok")] # the lock is freed even on failure
@pytest.mark.asyncio
async def test_winner_renews_the_lock_while_refresh_runs():
lock = _FakeLock(acquired=LockAcquisition.ACQUIRED)
refresh_finished = asyncio.Event()
sleep_started = asyncio.Event()
sleep_can_finish = asyncio.Event()
sleep_count = 0
async def sleep(seconds):
nonlocal sleep_count
sleep_count += 1
assert seconds == 5.0
sleep_started.set()
if sleep_count > 1:
await asyncio.Event().wait()
return
await sleep_can_finish.wait()
async def refresh():
await refresh_finished.wait()
return OAuthToken(access_token="new")
async def reread():
return None
coord = RedisRefreshCoordinator(lock, new_token=lambda: "tok", sleep=sleep)
task = asyncio.create_task(coord.run("u", "s", refresh, reread))
await sleep_started.wait()
sleep_can_finish.set()
await lock.extended.wait()
refresh_finished.set()
result = await task
assert result is not None and result.access_token == "new"
assert lock.extend_calls == [(_KEY, "tok", 10.0)]
assert lock.released == [(_KEY, "tok")]
@pytest.mark.asyncio
async def test_loser_waits_for_the_holder_then_rereads_persisted_token():
clock = _Clock()
lock = _FakeLock(acquired=LockAcquisition.HELD, 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 releases a lock it does not own
@pytest.mark.asyncio
async def test_loser_rereads_after_timeout_if_holder_never_releases():
clock = _Clock()
lock = _FakeLock(acquired=LockAcquisition.HELD, held_sequence=[True] * 100)
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"
@pytest.mark.asyncio
async def test_lock_backend_error_refreshes_anyway_instead_of_serving_stale():
# Regression: on a total lock-backend outage every worker gets ERROR (not HELD). If ERROR were
# treated as "someone else holds it", no worker would refresh and all would re-read the still-
# expired token and serve a stale bearer upstream. ERROR must instead refresh anyway.
lock = _FakeLock(acquired=LockAcquisition.ERROR)
refreshed = OAuthToken(access_token="refreshed-despite-redis-down")
refresh_calls = []
reread_calls = []
async def refresh():
refresh_calls.append(1)
return refreshed
async def reread():
reread_calls.append(1)
return OAuthToken(access_token="stale-expired-token")
result = await RedisRefreshCoordinator(lock).run("u", "s", refresh, reread)
assert result is refreshed # served the fresh token, not the stale re-read
assert refresh_calls == [1]
assert reread_calls == [] # never fell back to re-reading the expired token
assert lock.released == [] # nothing was acquired, so nothing is released
def test_wait_timeout_outlasts_the_holders_max_lock_hold():
# Regression: a loser's wait must exceed the holder's max lock-hold - the renewal budget plus one
# lease tail. When wait_timeout_seconds was merely == lock_ttl_seconds, a loser bailed at the lease
# TTL while the holder was still renewing mid-refresh and re-read the still-expired token.
coord = RedisRefreshCoordinator(_FakeLock(acquired=LockAcquisition.HELD))
assert coord.wait_timeout_seconds > coord.refresh_budget_seconds + coord.lock_ttl_seconds
@pytest.mark.asyncio
async def test_loser_with_default_timeout_waits_past_the_lease_ttl_for_the_holder():
# With the DEFAULT wait_timeout (not an inflated test value), a loser keeps waiting while the holder
# still holds the lock past lock_ttl_seconds, then re-reads what the holder persisted - rather than
# giving up at the lease TTL mid-refresh. At 0.5s/poll, 20 polls reach the 10s lease TTL.
clock = _Clock()
lock = _FakeLock(acquired=LockAcquisition.HELD, held_sequence=[True] * 30 + [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, step=0.5))
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 waits for the holder
assert lock._held == [] # waited until the holder released (whole sequence drained)
assert clock.t > coord.lock_ttl_seconds # ...past the lease TTL, not bailing at it
@pytest.mark.asyncio
async def test_winner_stops_renewing_after_the_refresh_budget():
# A holder whose refresh outruns the budget stops renewing - the lock then lapses for losers -
# instead of extending the lease forever and wedging the single-flight on one slow request.
clock = _Clock()
lock = _FakeLock(acquired=LockAcquisition.ACQUIRED)
release_refresh = asyncio.Event()
async def sleep(seconds):
clock.t += seconds # renewal sleeps lock_ttl/2 each cycle
async def refresh():
await release_refresh.wait() # outlives the budget
return OAuthToken(access_token="new")
async def reread():
return None
coord = RedisRefreshCoordinator(
lock,
new_token=lambda: "tok",
clock=clock,
sleep=sleep,
lock_ttl_seconds=10.0,
refresh_budget_seconds=20.0,
)
task = asyncio.create_task(coord.run("u", "s", refresh, reread))
for _ in range(200): # let the bounded renewal drain to its budget (clock 5..20)
if clock.t >= 20.0:
break
await asyncio.sleep(0)
extends_during_budget = len(lock.extend_calls)
release_refresh.set()
result = await task
assert result is not None and result.access_token == "new"
assert extends_during_budget == 4 # renewed at clock 5, 10, 15, 20 then stopped
assert len(lock.extend_calls) == 4 # no further renewals after the budget

View file

@ -0,0 +1,54 @@
"""Tests for the cache codec: encrypt on encode, drop the refresh_token, round-trip the bearer."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_codec import (
OAuthTokenCacheCodec,
)
def _wrapping_codec():
# A reversible stand-in for NaCl: proves encode() encrypts (output is wrapped) and decode()
# decrypts, without needing a salt key.
return OAuthTokenCacheCodec(
encrypt=lambda s: f"enc:{s}",
decrypt=lambda b: b[4:] if b.startswith("enc:") else None,
)
def test_round_trips_the_access_token():
codec = _wrapping_codec()
token = codec.decode(
codec.encode(OAuthToken(access_token="at-123", expires_at=1234.5))
)
assert token is not None
assert token.access_token == "at-123"
def test_encode_encrypts_and_omits_the_refresh_token():
codec = _wrapping_codec()
blob = codec.encode(OAuthToken(access_token="at", refresh_token="super-secret-rt"))
assert blob.startswith("enc:") # encryption was applied
assert (
"super-secret-rt" not in blob
) # the long-lived secret never reaches the cache
decoded = codec.decode(blob)
assert decoded is not None and decoded.refresh_token is None
def test_decoded_token_defers_expiry_to_the_cache_ttl():
codec = _wrapping_codec()
token = codec.decode(codec.encode(OAuthToken(access_token="at", expires_at=999.0)))
# The value carries no expiry; the cache entry's TTL bounds its life instead.
assert token is not None and token.expires_at is None
def test_undecryptable_blob_is_a_miss():
# e.g. master-key rotation makes an old entry unreadable -> treat as a miss, not a crash.
assert _wrapping_codec().decode("not-our-prefix") is None
def test_empty_plaintext_is_a_miss():
codec = OAuthTokenCacheCodec(encrypt=lambda s: s, decrypt=lambda b: b)
assert codec.decode("") is None