mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(mcp): cross-replica single-flight refresh for the v2 per-user OAuth store [2/2] (#31474)
* feat(mcp): encrypt+serialize codec for caching OAuth tokens in Redis (step 1b §1.5) The serialize+encrypt boundary a cross-replica cache needs: a plaintext bearer in Redis is a leak, so encode() encrypts (NaCl in prod via the injected encrypt, identity in tests). Caches only access_token and expires_at, never the refresh_token - the hot path needs just the bearer, and the long-lived refresh_token stays in the DB (the refresh path is always a cache miss), matching v1. A decoded token always has refresh_token=None. Undecryptable (key rotation) or corrupt entries read as a miss. * feat(mcp): DualCache-backed token cache backend (step 1b §1.5) The cross-replica TokenCacheBackend implementation that plugs into the foundation's CachedOAuthTokenStore seam: encrypts+serializes the token via the codec and stores it in LiteLLM's shared DualCache under the same per-(user,server) key v1 used, so workers share one refresh and a token cached by v1 or v2 is readable by the other across the cutover. Cache and codec are injected; a non-positive TTL (already-expired token) is not cached, and a missing/corrupt entry reads as a miss. * feat(mcp): Redis SET NX PX refresh coordinator (step 1b §1.5) The cross-replica RefreshCoordinator that plugs into the foundation's RefreshingTokenStore seam: a SET NX PX lock elects one worker to refresh per (user, server) while the rest wait for it and re-read the token it persisted, so a rotating refresh_token is used once across the fleet, not once per worker. The lock self-expires (PX) so a crashed holder can't wedge refresh; a loser falls back to a bounded re-read and the surrounding store re-checks expiry next fetch, so a crash self-heals. The lock (a thin Redis SET NX/DEL/EXISTS wrapper in prod) is injected, so the single-flight logic is testable without Redis. * feat(mcp): Redis SET NX PX distributed lock (step 1b §1.5) The concrete DistributedLock the RedisRefreshCoordinator elects refreshers with: acquire is an atomic SET key NX PX ttl (first caller wins, entry self-expires so a crashed holder can't wedge refresh), release is DEL, is_held is EXISTS. The async Redis client is injected (the client from LiteLLM's RedisCache in prod), so it is unit-testable with a fake. Any Redis error degrades to not-acquired / not-held so a cache blip causes an extra refresh, never a crash on the resolve path. * feat(mcp): wire the cross-replica cache + coordinator into the per-user store (step 1b §1.5) Upgrade the composition root to use the DualCache-backed cache and SET NX PX refresh coordinator when Redis is wired, falling back to the foundation's in-process defaults on a single replica. Layers the cross-replica path on top of the single-replica dispatch store. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> * fix(mcp): refresh on lock-backend error instead of serving a stale token The cross-replica refresh coordinator elected refreshers with a boolean acquire: a Redis transport error was caught and returned as False, which is indistinguishable from "another worker holds the lock". On a total Redis outage every worker therefore took the wait-then-reread branch and served the still-expired token upstream (the upstream then 401s), even though the lock and coordinator docstrings claimed a Redis blip "degrades to an extra refresh". Make acquire tristate (LockAcquisition: ACQUIRED / HELD / ERROR) so the coordinator can tell a busy holder from a dead backend, and refresh anyway on ERROR. This single-flight lock is a load optimization, not a correctness mutex, so failing open is correct: it degrades a lock-backend outage to the no-coordinator behavior (an extra refresh), never a stale bearer. Add a regression test asserting an acquire error refreshes rather than re-reading the expired token, and update the docstrings to match. * style(mcp): wrap redis lock signatures at line-length 88 for CI ruff format * fix(mcp): a refresh loser surfaces None, not a stale token, when the winner failed The cross-replica coordinator's losers re-read the token the winner persisted. If the winner's refresh failed, the store still holds the expired token, so the loser re-read it and RefreshingTokenStore handed that expired bearer to the caller (the upstream then 401s) instead of the re-auth challenge the winner returned via None. Make the loser's re-read expiry-aware, mirroring refresh_latest_token: a re-read that is still expired surfaces None so the arm challenges. This only affects the loser path; the winner's freshly refreshed token is returned directly by the coordinator and is unaffected. --------- Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
84b77fc72c
commit
cd2fb6b0f2
11 changed files with 782 additions and 67 deletions
|
|
@ -0,0 +1,65 @@
|
|||
"""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 typing import Protocol
|
||||
|
||||
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: ...
|
||||
|
||||
|
||||
class DualCacheTokenCacheBackend:
|
||||
def __init__(
|
||||
self,
|
||||
cache: AsyncCache,
|
||||
codec: OAuthTokenCacheCodec,
|
||||
*,
|
||||
key_prefix: str = "mcp:per_user_token:",
|
||||
) -> None:
|
||||
self._cache = cache
|
||||
self._codec = codec
|
||||
self._key_prefix = key_prefix
|
||||
|
||||
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:
|
||||
blob = await self._cache.async_get_cache(self._key(user_id, server_id))
|
||||
return self._codec.decode(blob) if isinstance(blob, str) else None
|
||||
|
||||
async def set(
|
||||
self, user_id: str, server_id: str, token: OAuthToken, ttl_seconds: float
|
||||
) -> None:
|
||||
if ttl_seconds <= 0:
|
||||
return
|
||||
await self._cache.async_set_cache(
|
||||
self._key(user_id, server_id),
|
||||
self._codec.encode(token),
|
||||
ttl=ttl_seconds,
|
||||
)
|
||||
|
||||
async def delete(self, user_id: str, server_id: str) -> None:
|
||||
await self._cache.async_delete_cache(self._key(user_id, server_id))
|
||||
|
|
@ -289,9 +289,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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,11 +1,11 @@
|
|||
"""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
|
||||
|
|
@ -17,12 +17,27 @@ 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,
|
||||
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,
|
||||
)
|
||||
|
|
@ -99,16 +114,48 @@ async def _post_token_endpoint(
|
|||
return body # pyright: ignore
|
||||
|
||||
|
||||
def _runtime_backend_and_coordinator() -> tuple[
|
||||
TokenCacheBackend | None, RefreshCoordinator | None
|
||||
]:
|
||||
"""The cross-replica cache + coordinator when Redis is wired, else ``(None, None)`` 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
|
||||
codec = OAuthTokenCacheCodec(
|
||||
encrypt_value_helper,
|
||||
lambda blob: decrypt_value_helper(blob, "mcp_per_user_token"),
|
||||
)
|
||||
# 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
|
||||
backend = DualCacheTokenCacheBackend(cache, codec)
|
||||
coordinator = RedisRefreshCoordinator(lock)
|
||||
return backend, coordinator
|
||||
|
||||
|
||||
def build_per_user_oauth_token_store(
|
||||
server_lookup: ServerLookup,
|
||||
) -> CachedOAuthTokenStore:
|
||||
backend, coordinator = _runtime_backend_and_coordinator()
|
||||
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)
|
||||
refreshing = RefreshingTokenStore(
|
||||
V2PerUserTokenStore(_read_credential), refresher, coordinator=coordinator
|
||||
)
|
||||
return CachedOAuthTokenStore(
|
||||
refreshing, default_ttl_seconds=_DEFAULT_TTL_SECONDS, backend=backend
|
||||
)
|
||||
|
||||
|
||||
class LazyPerUserOAuthTokenStore:
|
||||
|
|
|
|||
|
|
@ -0,0 +1,66 @@
|
|||
"""Concrete ``DistributedLock`` over a Redis client: ``SET NX PX`` / ``DEL`` / ``EXISTS``.
|
||||
|
||||
The cross-replica lock the ``RedisRefreshCoordinator`` elects refreshers with. ``acquire`` is an
|
||||
atomic ``SET key NX PX ttl`` (only the first caller wins; the entry self-expires so a crashed holder
|
||||
can't wedge refresh), ``release`` is ``DEL``, ``is_held`` is ``EXISTS``. 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 typing import Protocol
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_coordinator import (
|
||||
LockAcquisition,
|
||||
)
|
||||
|
||||
|
||||
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 delete(self, *names: str) -> int: ...
|
||||
|
||||
async def exists(self, *names: str) -> int: ...
|
||||
|
||||
|
||||
class RedisDistributedLock:
|
||||
def __init__(self, client: RedisCommands) -> None:
|
||||
self._client = client
|
||||
|
||||
async def acquire(self, key: str, ttl_seconds: float) -> LockAcquisition:
|
||||
try:
|
||||
result = await self._client.set(
|
||||
key, "1", 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 release(self, key: str) -> None:
|
||||
try:
|
||||
await self._client.delete(key)
|
||||
# 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(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
|
||||
|
|
@ -0,0 +1,102 @@
|
|||
"""Cross-replica ``RefreshCoordinator``: one refresh per ``(user, server)`` across all workers.
|
||||
|
||||
Plugs into the foundation's ``RefreshingTokenStore`` via the ``RefreshCoordinator`` seam. A ``SET NX
|
||||
PX`` lock elects one worker to run the refresh while the rest wait for it and re-read the token it
|
||||
persisted - so a rotating refresh_token is used once across the fleet, not once per worker. The lock
|
||||
auto-expires (``PX``), so a crashed holder can't wedge refresh; a loser that times out (or whose
|
||||
holder crashed mid-refresh) falls back to a re-read, and the surrounding store re-checks expiry on the
|
||||
next fetch, so a crash self-heals rather than serving stale forever. Reading needs no lock, so losers
|
||||
don't serialize behind each other. The lock is injected (a thin Redis ``SET NX``/``DEL``/``EXISTS``
|
||||
wrapper in production, a fake in tests).
|
||||
|
||||
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
|
||||
from collections.abc import Awaitable, Callable
|
||||
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 NX PX ttl`` reported as a
|
||||
``LockAcquisition`` (won / held by another / backend error); ``release`` is ``DEL``; ``is_held``
|
||||
is ``EXISTS`` (so a waiter can poll without taking the lock)."""
|
||||
|
||||
async def acquire(self, key: str, ttl_seconds: float) -> LockAcquisition: ...
|
||||
|
||||
async def release(self, key: str) -> None: ...
|
||||
|
||||
async def is_held(self, key: str) -> bool: ...
|
||||
|
||||
|
||||
class RedisRefreshCoordinator:
|
||||
def __init__(
|
||||
self,
|
||||
lock: DistributedLock,
|
||||
*,
|
||||
key_prefix: str = "mcp:refresh_lock:",
|
||||
lock_ttl_seconds: float = 10.0,
|
||||
wait_timeout_seconds: float = 10.0,
|
||||
poll_interval_seconds: float = 0.05,
|
||||
sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
|
||||
clock: Callable[[], float] = time.monotonic,
|
||||
) -> None:
|
||||
self._lock = lock
|
||||
self._key_prefix = key_prefix
|
||||
self._lock_ttl_seconds = lock_ttl_seconds
|
||||
self._wait_timeout_seconds = wait_timeout_seconds
|
||||
self._poll_interval_seconds = poll_interval_seconds
|
||||
self._sleep = sleep
|
||||
self._clock = clock
|
||||
|
||||
def _key(self, user_id: str, server_id: str) -> str:
|
||||
return f"{self._key_prefix}{user_id}:{server_id}"
|
||||
|
||||
async def run(
|
||||
self,
|
||||
user_id: str,
|
||||
server_id: str,
|
||||
refresh: Callable[[], Awaitable[OAuthToken | None]],
|
||||
reread: Callable[[], Awaitable[OAuthToken | None]],
|
||||
) -> OAuthToken | None:
|
||||
key = self._key(user_id, server_id)
|
||||
match await self._lock.acquire(key, self._lock_ttl_seconds):
|
||||
case LockAcquisition.ACQUIRED:
|
||||
try:
|
||||
return await refresh()
|
||||
finally:
|
||||
await self._lock.release(key)
|
||||
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()
|
||||
|
|
@ -0,0 +1,37 @@
|
|||
"""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 litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
OAuthToken,
|
||||
)
|
||||
|
||||
|
||||
class OAuthTokenCacheCodec:
|
||||
def __init__(
|
||||
self,
|
||||
encrypt: Callable[[str], str],
|
||||
decrypt: Callable[[str], str | None],
|
||||
) -> None:
|
||||
self._encrypt = encrypt
|
||||
self._decrypt = decrypt
|
||||
|
||||
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)
|
||||
|
|
@ -0,0 +1,93 @@
|
|||
"""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_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"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -0,0 +1,73 @@
|
|||
"""Tests for the Redis SET NX PX lock: acquire semantics, release, is_held, error degradation."""
|
||||
|
||||
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:
|
||||
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.deleted = []
|
||||
|
||||
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))
|
||||
return self._set_returns
|
||||
|
||||
async def delete(self, *names):
|
||||
self.deleted.extend(names)
|
||||
return len(names)
|
||||
|
||||
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_uses_set_nx_px_and_reports_acquired():
|
||||
redis = _FakeRedis(set_returns=True)
|
||||
lock = RedisDistributedLock(redis)
|
||||
assert await lock.acquire("k", 10.0) is LockAcquisition.ACQUIRED
|
||||
assert redis.set_calls == [("k", "1", True, 10000)] # NX + px in milliseconds
|
||||
|
||||
|
||||
@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", 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
|
||||
# instead of waiting and serving a stale token.
|
||||
assert await RedisDistributedLock(_FakeRedis(raise_on=["set"])).acquire("k", 10.0) is LockAcquisition.ERROR
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_deletes_the_key():
|
||||
redis = _FakeRedis()
|
||||
await RedisDistributedLock(redis).release("k")
|
||||
assert redis.deleted == ["k"]
|
||||
|
||||
|
||||
@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
|
||||
|
|
@ -0,0 +1,149 @@
|
|||
"""Tests for the cross-replica refresh coordinator: winner refreshes, losers wait then re-read."""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
OAuthToken,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_coordinator import (
|
||||
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.acquired_keys = []
|
||||
self.released = []
|
||||
|
||||
async def acquire(self, key, ttl_seconds):
|
||||
self.acquired_keys.append((key, ttl_seconds))
|
||||
return self._acquired
|
||||
|
||||
async def release(self, key):
|
||||
self.released.append(key)
|
||||
|
||||
async def is_held(self, key):
|
||||
return self._held.pop(0) if self._held else False
|
||||
|
||||
|
||||
class _Clock:
|
||||
def __init__(self):
|
||||
self.t = 0.0
|
||||
|
||||
def __call__(self):
|
||||
return self.t
|
||||
|
||||
|
||||
def _advancing_sleep(clock, step=0.5):
|
||||
async def sleep(_seconds):
|
||||
clock.t += step
|
||||
|
||||
return sleep
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_winner_refreshes_then_releases_and_never_rereads():
|
||||
lock = _FakeLock(acquired=LockAcquisition.ACQUIRED)
|
||||
refreshed = OAuthToken(access_token="new")
|
||||
reread_calls = []
|
||||
|
||||
async def refresh():
|
||||
return refreshed
|
||||
|
||||
async def reread():
|
||||
reread_calls.append(1)
|
||||
return None
|
||||
|
||||
result = await RedisRefreshCoordinator(lock).run("u", "s", refresh, reread)
|
||||
assert result is refreshed
|
||||
assert lock.acquired_keys == [(_KEY, 10.0)]
|
||||
assert lock.released == [_KEY]
|
||||
assert reread_calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_winner_releases_even_when_refresh_raises():
|
||||
lock = _FakeLock(acquired=LockAcquisition.ACQUIRED)
|
||||
|
||||
async def refresh():
|
||||
raise RuntimeError("boom")
|
||||
|
||||
async def reread():
|
||||
return None
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
await RedisRefreshCoordinator(lock).run("u", "s", refresh, reread)
|
||||
assert lock.released == [_KEY] # the lock is freed even on failure
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_loser_waits_for_the_holder_then_rereads_persisted_token():
|
||||
clock = _Clock()
|
||||
lock = _FakeLock(acquired=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 holds the lock
|
||||
|
||||
|
||||
@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) # holder never releases
|
||||
|
||||
async def refresh():
|
||||
return None
|
||||
|
||||
async def reread():
|
||||
return OAuthToken(access_token="whatever-is-there")
|
||||
|
||||
coord = RedisRefreshCoordinator(
|
||||
lock,
|
||||
clock=clock,
|
||||
sleep=_advancing_sleep(clock, step=0.5),
|
||||
wait_timeout_seconds=1.0,
|
||||
poll_interval_seconds=0.1,
|
||||
)
|
||||
result = await coord.run("u", "s", refresh, reread)
|
||||
# Gave up waiting (bounded) and returned what's persisted rather than blocking forever.
|
||||
assert result is not None and result.access_token == "whatever-is-there"
|
||||
|
||||
|
||||
@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
|
||||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue