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

1"""Per-user OAuth token store for the ``authorization_code`` mode. 

2 

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

6 

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""" 

14 

15from __future__ import annotations 

16 

17import asyncio 

18import time 

19from collections.abc import Awaitable, Callable 

20from dataclasses import dataclass 

21from typing import Final, Protocol 

22 

23 

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. 

27 

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

34 

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 """ 

39 

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 

45 

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})" 

49 

50 

51class TokenStoreUnavailable(Exception): 

52 """Raised by ``fetch`` when the backing token store is unreachable (e.g. the DB is down). 

53 

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 """ 

58 

59 

60class OAuthTokenStore(Protocol): 

61 """Per-user OAuth token lookup for the ``authorization_code`` mode. 

62 

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 """ 

69 

70 async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: ... 70 ↛ exitline 70 didn't return from function 'fetch' because

71 

72 

73class InvalidatableOAuthTokenStore(OAuthTokenStore, Protocol): 

74 """An ``OAuthTokenStore`` whose cached entry for a ``(user, server)`` pair can be dropped. 

75 

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 """ 

80 

81 async def invalidate(self, user_id: str, server_id: str) -> None: ... 81 ↛ exitline 81 didn't return from function 'invalidate' because

82 

83 

84class TokenRefresher(Protocol): 

85 """Mints a fresh token from an expired one and persists it, returning the new token. 

86 

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. 

91 

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 """ 

96 

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

98 

99 

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 """ 

106 

107 async def get(self, user_id: str, server_id: str) -> OAuthToken | None: ... 107 ↛ exitline 107 didn't return from function 'get' because

108 

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

110 

111 async def delete(self, user_id: str, server_id: str) -> None: ... 111 ↛ exitline 111 didn't return from function 'delete' because

112 

113 

114class InMemoryTokenCacheBackend: 

115 """Per-process token cache: a bounded dict with wall-clock TTLs (the default backend).""" 

116 

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]] = {} 

121 

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 

132 

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) 

140 

141 async def delete(self, user_id: str, server_id: str) -> None: 

142 self._cache.pop((user_id, server_id), None) 

143 

144 

145class CachedOAuthTokenStore: 

146 """Expiry-aware cache over an ``OAuthTokenStore``. Caches positive tokens only. 

147 

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 """ 

156 

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) 

172 

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 

177 

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 

182 

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 

191 

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) 

196 

197 

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 """ 

204 

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

212 

213 

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 """ 

220 

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]] = {} 

225 

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 

241 

242 

243class RefreshingTokenStore: 

244 """An ``OAuthTokenStore`` that proactively refreshes a near-expiry token. 

245 

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. 

252 

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 """ 

257 

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

272 

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 

275 

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 

280 

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) 

286 

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 

295 

296 return await self._coordinator.run( 

297 user_id, 

298 server_id, 

299 refresh=refresh_latest_token, 

300 reread=reread_fresh_token, 

301 )