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 2e9044a3db6..8a379cc06ec 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 @@ -77,13 +77,15 @@ class TokenRefresher(Protocol): class CachedOAuthTokenStore: - """Expiry-aware cache over an ``OAuthTokenStore``. + """Expiry-aware cache over an ``OAuthTokenStore``. Caches positive tokens only. - A cached token is served only while it is unexpired (minus ``expiry_skew_seconds``); past that - the inner store is read again. Tokens with no known expiry, and the ``None`` "not authorized" - result, are held for ``default_ttl_seconds`` so the store is not hit on every call. The clock is - injected (wall-clock, since ``expires_at`` is epoch) so expiry is deterministic in tests, and a - store outage (``TokenStoreUnavailable``) propagates without being cached. + A cached token is served only while it is unexpired (minus ``expiry_skew_seconds``), or for + ``default_ttl_seconds`` if it carries no expiry; past that the inner store is read again. A + "not authorized" (``None``) result is never cached: every miss re-reads the inner store, so a + token written after the OAuth flow is visible immediately on every replica, matching v1 (which + never caches misses). The clock is injected (wall-clock, since ``expires_at`` is epoch) so + expiry is deterministic in tests, and a store outage (``TokenStoreUnavailable``) propagates + without being cached. """ def __init__( @@ -100,10 +102,10 @@ class CachedOAuthTokenStore: self._expiry_skew_seconds = expiry_skew_seconds self._max_size = max_size self._clock = clock - self._cache: dict[tuple[str, str], tuple[OAuthToken | None, float]] = {} + self._cache: dict[tuple[str, str], tuple[OAuthToken, float]] = {} - def _valid_until(self, token: OAuthToken | None) -> float: - if token is not None and token.expires_at is not None: + def _valid_until(self, token: OAuthToken) -> float: + if token.expires_at is not None: return token.expires_at - self._expiry_skew_seconds return self._clock() + self._default_ttl_seconds @@ -116,6 +118,11 @@ class CachedOAuthTokenStore: return token token = await self._inner.fetch(user_id, server_id) + if token is None: + # Never cache "not authorized": drop any stale entry and re-read on the next call, so + # a token stored after the OAuth flow is seen immediately rather than after a TTL. + self._cache.pop(key, None) + return token if key not in self._cache and len(self._cache) >= self._max_size: # Evict the oldest entry (insertion order) to make room, rather than clearing the # whole cache and forcing every key to re-read the store at once. 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 e8baae782b6..4fd4c27ac67 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 @@ -61,7 +61,7 @@ async def test_refetches_once_token_has_expired(): assert len(inner.calls) == 2 # re-read after the cached token expired -async def test_caches_not_authorized_none_for_default_ttl(): +async def test_does_not_cache_the_not_authorized_miss(): inner = _FakeStore({}) # user has not authorized clock = _Clock(1000.0) store = CachedOAuthTokenStore(inner, default_ttl_seconds=60, clock=clock) @@ -69,7 +69,22 @@ async def test_caches_not_authorized_none_for_default_ttl(): assert await store.fetch("u", "s") is None clock.t = 1059.0 assert await store.fetch("u", "s") is None - assert inner.calls == [("u", "s")] # None cached for the TTL window + assert inner.calls == [ + ("u", "s"), + ("u", "s"), + ] # misses re-read the store, never cached + + +async def test_token_stored_after_a_miss_is_visible_immediately(): + # No invalidation needed: a miss is never cached, so a token written after the OAuth flow is + # served on the very next call (matching v1, where misses always re-read the source). + inner = _FakeStore({}) + store = CachedOAuthTokenStore(inner, default_ttl_seconds=60, clock=_Clock()) + + assert await store.fetch("u", "s") is None + inner._values[("u", "s")] = OAuthToken(access_token="fresh") + result = await store.fetch("u", "s") + assert result is not None and result.access_token == "fresh" async def test_default_ttl_applies_to_tokens_without_expiry(): @@ -84,15 +99,19 @@ async def test_default_ttl_applies_to_tokens_without_expiry(): assert len(inner.calls) == 2 # no-expiry token re-read after the default TTL -async def test_invalidate_forces_refetch(): - inner = _FakeStore({}) +async def test_invalidate_drops_a_cached_token(): + # invalidate covers rotation/revocation of a *cached* token (the miss path needs no invalidate). + inner = _FakeStore({("u", "s"): OAuthToken(access_token="t1")}) store = CachedOAuthTokenStore(inner, default_ttl_seconds=60, clock=_Clock()) - assert await store.fetch("u", "s") is None - inner._values[("u", "s")] = OAuthToken(access_token="fresh") + first = await store.fetch("u", "s") + assert first is not None and first.access_token == "t1" # cached + inner._values[("u", "s")] = OAuthToken(access_token="t2") # rotated store.invalidate("u", "s") - result = await store.fetch("u", "s") - assert result is not None and result.access_token == "fresh" + second = await store.fetch("u", "s") + assert ( + second is not None and second.access_token == "t2" + ) # re-read after invalidate async def test_store_unavailable_is_not_cached():