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