Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py: 47%
109 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""Per-user OAuth token store for the ``authorization_code`` mode.
3The resolver reads a user's token through the injected ``OAuthTokenStore`` seam;
4``CachedOAuthTokenStore`` is an expiry-aware cache in front of it. ``TokenStoreUnavailable``
5signals an unreachable backing store, so an outage is never cached or read as "not authorized".
7``RefreshingTokenStore`` mints a fresh token through an injected ``TokenRefresher`` when the stored
8one is near expiry, under in-process per-(user, server) single-flight so concurrent callers share
9one refresh. Distributed (cross-replica) single-flight and reactive-401 refresh are the later
10hardening. The mode plugs in its own source and refresher; the cache, store seam, and refresh
11machinery are shared across the oauth2 modes (authorization_code / client_credentials /
12token_exchange).
13"""
15from __future__ import annotations
17import asyncio
18import time
19from collections.abc import Awaitable, Callable
20from dataclasses import dataclass
21from typing import Final, Protocol
24@dataclass(frozen=True, slots=True, repr=False)
25class OAuthToken:
26 """A user's OAuth credential: the bearer value, when it expires, and how to refresh it.
28 ``expires_at`` is epoch seconds (``None`` means no known expiry). ``refresh_token`` is what a
29 ``TokenRefresher`` uses to mint a new access token when this one nears expiry (the refresh
30 mechanism, ``RefreshingTokenStore``, is in this module; the concrete per-mode refresher lands
31 with each mode); it is never minted into a header directly. ``repr`` masks both secrets so a
32 stray log line cannot leak them (the values are still plain ``str`` for the header path, since
33 ``SecretStr`` resolves as unknown under this repo's basedpyright).
35 ``scopes`` is the recorded grant. A refresh response that omits ``scope`` (RFC 6749 §5.1: an
36 omitted ``scope`` means unchanged) carries the prior value forward, so a refresh never silently
37 drops it; the resolver itself does not read it.
38 """
40 access_token: str
41 expires_at: float | None = None
42 refresh_token: str | None = None
43 scopes: tuple[str, ...] = ()
44 identity_binding_proof: str | None = None
46 def __repr__(self) -> str:
47 has_refresh: Final = self.refresh_token is not None
48 return f"OAuthToken(access_token=***, expires_at={self.expires_at!r}, has_refresh_token={has_refresh}, scopes={self.scopes!r})"
51class TokenStoreUnavailable(Exception):
52 """Raised by ``fetch`` when the backing token store is unreachable (e.g. the DB is down).
54 Distinct from returning ``None`` for "the user has not authorized this server": a read-through
55 cache skips caching the failure, and the resolver maps it to its fail-closed status rather than
56 treating an outage as a definite absence.
57 """
60class OAuthTokenStore(Protocol):
61 """Per-user OAuth token lookup for the ``authorization_code`` mode.
63 Returns the user's token for an upstream, or ``None`` when they have not completed the OAuth
64 flow (the arm turns that into a 401 challenge). The ``(user_id, server_id)`` pair fully scopes
65 the lookup, so an implementation must never return one subject's token to another. Raises
66 ``TokenStoreUnavailable`` when the backing store is unreachable, so an outage is never cached or
67 read as a definite absence.
68 """
70 async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: ... 70 ↛ exitline 70 didn't return from function 'fetch' because
73class InvalidatableOAuthTokenStore(OAuthTokenStore, Protocol):
74 """An ``OAuthTokenStore`` whose cached entry for a ``(user, server)`` pair can be dropped.
76 The write side calls ``invalidate`` after a (re)authorization or revocation changes the
77 credential row, so reads stop serving the replaced token immediately instead of until its
78 cache TTL. ``CachedOAuthTokenStore`` (the top of the per-user chain) satisfies this.
79 """
81 async def invalidate(self, user_id: str, server_id: str) -> None: ... 81 ↛ exitline 81 didn't return from function 'invalidate' because
84class TokenRefresher(Protocol):
85 """Mints a fresh token from an expired one and persists it, returning the new token.
87 The action is mode-specific: the ``authorization_code`` refresh_token grant, the
88 ``client_credentials`` grant, or an RFC 8693 re-exchange. Returns ``None`` when it cannot
89 refresh (e.g. no ``refresh_token``), which the caller turns into a 401 challenge. It must
90 persist the new token so later requests (and the surrounding cache) read it without refreshing.
92 ``server_id`` selects the upstream's config (token endpoint, client credentials, scopes) the
93 grant runs against; ``(user_id, server_id)`` is the key the new token is persisted under. They
94 are not derivable from ``token``, so the seam threads them alongside it.
95 """
97 async def refresh(self, user_id: str, server_id: str, token: OAuthToken) -> OAuthToken | None: ... 97 ↛ exitline 97 didn't return from function 'refresh' because
100class TokenCacheBackend(Protocol):
101 """Storage behind ``CachedOAuthTokenStore``: hold a token under ``(user_id, server_id)`` for
102 ``ttl_seconds``, then forget it. The default ``InMemoryTokenCacheBackend`` is per-process; a
103 cross-replica deployment injects a shared (Redis) backend so every worker reads one refresh,
104 matching v1. ``get`` returns ``None`` once the entry's TTL has elapsed.
105 """
107 async def get(self, user_id: str, server_id: str) -> OAuthToken | None: ... 107 ↛ exitline 107 didn't return from function 'get' because
109 async def set(self, user_id: str, server_id: str, token: OAuthToken, ttl_seconds: float) -> None: ... 109 ↛ exitline 109 didn't return from function 'set' because
111 async def delete(self, user_id: str, server_id: str) -> None: ... 111 ↛ exitline 111 didn't return from function 'delete' because
114class InMemoryTokenCacheBackend:
115 """Per-process token cache: a bounded dict with wall-clock TTLs (the default backend)."""
117 def __init__(self, *, max_size: int = 4096, clock: Callable[[], float] = time.time) -> None:
118 self._max_size = max_size
119 self._clock = clock
120 self._cache: dict[tuple[str, str], tuple[OAuthToken, float]] = {}
122 async def get(self, user_id: str, server_id: str) -> OAuthToken | None:
123 key: Final = (user_id, server_id)
124 hit: Final = self._cache.get(key)
125 if hit is None:
126 return None
127 token, valid_until = hit
128 if self._clock() < valid_until:
129 return token
130 self._cache.pop(key, None)
131 return None
133 async def set(self, user_id: str, server_id: str, token: OAuthToken, ttl_seconds: float) -> None:
134 key: Final = (user_id, server_id)
135 if key not in self._cache and len(self._cache) >= self._max_size:
136 # Evict the oldest entry (insertion order), rather than clearing the whole cache and
137 # forcing every key to re-read the store at once.
138 self._cache.pop(next(iter(self._cache)), None)
139 self._cache[key] = (token, self._clock() + ttl_seconds)
141 async def delete(self, user_id: str, server_id: str) -> None:
142 self._cache.pop((user_id, server_id), None)
145class CachedOAuthTokenStore:
146 """Expiry-aware cache over an ``OAuthTokenStore``. Caches positive tokens only.
148 A cached token is served only while it is unexpired (minus ``expiry_skew_seconds``), or for
149 ``default_ttl_seconds`` if it carries no expiry; past that the inner store is read again. A
150 "not authorized" (``None``) result is never cached: every miss re-reads the inner store, so a
151 token written after the OAuth flow is visible immediately on every replica, matching v1 (which
152 never caches misses). The clock is injected (wall-clock, since ``expires_at`` is epoch) so
153 expiry is deterministic in tests, and a store outage (``TokenStoreUnavailable``) propagates
154 without being cached.
155 """
157 def __init__(
158 self,
159 inner: OAuthTokenStore,
160 *,
161 default_ttl_seconds: float,
162 expiry_skew_seconds: float = 60.0,
163 max_size: int = 4096,
164 backend: TokenCacheBackend | None = None,
165 clock: Callable[[], float] = time.time,
166 ) -> None:
167 self._inner = inner
168 self._default_ttl_seconds = default_ttl_seconds
169 self._expiry_skew_seconds = expiry_skew_seconds
170 self._clock = clock
171 self._backend: TokenCacheBackend = backend or InMemoryTokenCacheBackend(max_size=max_size, clock=clock)
173 def _ttl(self, token: OAuthToken) -> float:
174 if token.expires_at is not None:
175 return max(0.0, token.expires_at - self._expiry_skew_seconds - self._clock())
176 return self._default_ttl_seconds
178 async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
179 hit: Final = await self._backend.get(user_id, server_id)
180 if hit is not None:
181 return hit
183 token: Final = await self._inner.fetch(user_id, server_id)
184 if token is None:
185 # Never cache "not authorized": drop any stale entry and re-read on the next call, so
186 # a token stored after the OAuth flow is seen immediately rather than after a TTL.
187 await self._backend.delete(user_id, server_id)
188 return token
189 await self._backend.set(user_id, server_id, token, self._ttl(token))
190 return token
192 async def invalidate(self, user_id: str, server_id: str) -> None:
193 """Drop a cached entry after the user (re)authorizes or revokes, so a stale token or a
194 stale "not authorized" None cannot mask the change."""
195 await self._backend.delete(user_id, server_id)
198class RefreshCoordinator(Protocol):
199 """Ensures one refresh runs per ``(user_id, server_id)`` at a time. Concurrent callers either
200 share the winner's result (the default ``InProcessRefreshCoordinator``) or, in a cross-replica
201 coordinator, wait for the holder and ``reread`` the token it persisted - so the IdP sees one
202 refresh per key across all workers, not one per worker.
203 """
205 async def run( 205 ↛ exitline 205 didn't return from function 'run' because
206 self,
207 user_id: str,
208 server_id: str,
209 refresh: Callable[[], Awaitable[OAuthToken | None]],
210 reread: Callable[[], Awaitable[OAuthToken | None]],
211 ) -> OAuthToken | None: ...
214class InProcessRefreshCoordinator:
215 """Single-flight within one event loop (the default): the first caller per key refreshes while
216 concurrent callers await the same in-flight task and share its result. ``reread`` is unused here -
217 the shared task already yields the new token - and exists for the cross-replica coordinator, where
218 losers re-read the persisted token instead of sharing an in-process future.
219 """
221 def __init__(self) -> None:
222 # In-flight refreshes, one task per (user, server); each entry is removed by the task's
223 # done-callback, so the map is bounded by concurrent refreshes, not by distinct keys seen.
224 self._inflight: dict[tuple[str, str], asyncio.Future[OAuthToken | None]] = {}
226 async def run(
227 self,
228 user_id: str,
229 server_id: str,
230 refresh: Callable[[], Awaitable[OAuthToken | None]],
231 reread: Callable[[], Awaitable[OAuthToken | None]],
232 ) -> OAuthToken | None:
233 key: Final = (user_id, server_id)
234 task = self._inflight.get(key)
235 if task is None:
236 # The task is detached from the caller, so a cancelled caller does not abort the refresh.
237 task = asyncio.ensure_future(refresh())
238 self._inflight[key] = task
239 task.add_done_callback(lambda _t, k=key: self._inflight.pop(k, None))
240 return await task
243class RefreshingTokenStore:
244 """An ``OAuthTokenStore`` that proactively refreshes a near-expiry token.
246 Reads from an inner store; if the token is within ``expiry_skew_seconds`` of expiry, it mints a
247 fresh one via the injected ``TokenRefresher``, serialized per ``(user, server)`` by the injected
248 ``RefreshCoordinator`` so callers don't stampede the IdP. The refresher persists the new token so
249 later requests (and the surrounding cache) read it without refreshing again. An expired token the
250 refresher cannot renew (``None``) is surfaced as ``None`` so the arm challenges, never a stale
251 bearer.
253 The default coordinator is in-process; a cross-replica deployment injects a distributed one (Redis
254 SET NX). Reactive-401 refresh is later hardening (it lives in the egress transport, which sees the
255 upstream's 401). Composes under ``CachedOAuthTokenStore`` so the refreshed token is cached.
256 """
258 def __init__(
259 self,
260 inner: OAuthTokenStore,
261 refresher: TokenRefresher,
262 *,
263 expiry_skew_seconds: float = 60.0,
264 coordinator: RefreshCoordinator | None = None,
265 clock: Callable[[], float] = time.time,
266 ) -> None:
267 self._inner = inner
268 self._refresher = refresher
269 self._expiry_skew_seconds = expiry_skew_seconds
270 self._clock = clock
271 self._coordinator: RefreshCoordinator = coordinator or InProcessRefreshCoordinator()
273 def _is_expired(self, token: OAuthToken) -> bool:
274 return token.expires_at is not None and self._clock() >= token.expires_at - self._expiry_skew_seconds
276 async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
277 token: Final = await self._inner.fetch(user_id, server_id)
278 if token is None or not self._is_expired(token):
279 return token
281 async def refresh_latest_token() -> OAuthToken | None:
282 latest_token: Final = await self._inner.fetch(user_id, server_id)
283 if latest_token is None or not self._is_expired(latest_token):
284 return latest_token
285 return await self._refresher.refresh(user_id, server_id, latest_token)
287 async def reread_fresh_token() -> OAuthToken | None:
288 # A loser re-reads what the winner persisted. If the winner's refresh failed, the store
289 # still holds the expired token; surface None (-> challenge) like the winner did rather
290 # than the stale bearer the upstream would 401.
291 latest_token: Final = await self._inner.fetch(user_id, server_id)
292 if latest_token is None or self._is_expired(latest_token):
293 return None
294 return latest_token
296 return await self._coordinator.run(
297 user_id,
298 server_id,
299 refresh=refresh_latest_token,
300 reread=reread_fresh_token,
301 )