refactor(mcp): cache positive tokens only, matching v1 (no negative caching)

CachedOAuthTokenStore no longer caches the "not authorized" None result; every miss re-reads the
inner store. v1's per-user token cache never caches misses, so a token written by the OAuth flow
is visible on the next request without an invalidation hook, and uniformly across replicas since
the in-process cache holds no stale None to clear. invalidate() now only covers rotation or
revocation of a cached token. Negative caching (with distributed invalidation) can return later
if a slow DB-backed v2-native source makes per-miss reads expensive.
This commit is contained in:
Tin Chi Lo 2026-06-25 10:00:15 -07:00
parent 6ed9ecfaa7
commit 92b2b8c26e
2 changed files with 43 additions and 17 deletions

View file

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

View file

@ -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():