From 5963b9320f4b70fdd02cdaa4f6c11bde0da297e0 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Sat, 27 Jun 2026 16:27:34 -0700 Subject: [PATCH] feat(mcp): cross-replica single-flight refresh for the v2 per-user OAuth store [2/2] (#31493) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 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) * 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:: 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 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) Co-authored-by: Cursor Agent --- .../dual_cache_token_backend.py | 74 +++++ .../outbound_credentials/oauth_token_store.py | 11 +- .../per_user_oauth_store.py | 132 +++++++-- .../redis_distributed_lock.py | 94 ++++++ .../redis_refresh_coordinator.py | 136 +++++++++ .../outbound_credentials/token_cache_codec.py | 34 +++ .../test_dual_cache_token_backend.py | 156 ++++++++++ .../test_oauth_token_store.py | 134 +++++---- .../test_per_user_oauth_store.py | 160 +++++++++++ .../test_redis_distributed_lock.py | 126 ++++++++ .../test_redis_refresh_coordinator.py | 270 ++++++++++++++++++ .../test_token_cache_codec.py | 54 ++++ 12 files changed, 1307 insertions(+), 74 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/outbound_credentials/dual_cache_token_backend.py create mode 100644 litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_distributed_lock.py create mode 100644 litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_refresh_coordinator.py create mode 100644 litellm/proxy/_experimental/mcp_server/outbound_credentials/token_cache_codec.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_dual_cache_token_backend.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_refresh_coordinator.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_cache_codec.py diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/dual_cache_token_backend.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/dual_cache_token_backend.py new file mode 100644 index 00000000000..66fe2169a46 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/dual_cache_token_backend.py @@ -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) 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 9c7ab326f29..fd2cb2f3e06 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 @@ -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, ) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py index cedd13e7b2f..7453709f358 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py @@ -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() diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_distributed_lock.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_distributed_lock.py new file mode 100644 index 00000000000..e3153907353 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_distributed_lock.py @@ -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 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..317f7c703e7 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_refresh_coordinator.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_cache_codec.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_cache_codec.py new file mode 100644 index 00000000000..b0ed708f607 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_cache_codec.py @@ -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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_dual_cache_token_backend.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_dual_cache_token_backend.py new file mode 100644 index 00000000000..f27bc987434 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_dual_cache_token_backend.py @@ -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 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 63cb407de19..d72afaf46b6 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 @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py new file mode 100644 index 00000000000..f64ce594efa --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py @@ -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")] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py new file mode 100644 index 00000000000..e06f3bcd5cf --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py @@ -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 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..a6831562018 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_refresh_coordinator.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_cache_codec.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_cache_codec.py new file mode 100644 index 00000000000..17a7e13af03 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_cache_codec.py @@ -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