Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py: 32%

133 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1""" 

2OAuth2 client_credentials token cache for MCP servers. 

3 

4Automatically fetches and refreshes access tokens for MCP servers configured 

5with ``client_id``, ``client_secret``, and ``token_url``. 

6""" 

7 

8import asyncio 

9import hashlib 

10from collections.abc import Mapping 

11from typing import TYPE_CHECKING, Final 

12 

13import httpx 

14 

15from litellm._logging import verbose_logger 

16from litellm.caching.in_memory_cache import InMemoryCache 

17from litellm.constants import ( 

18 MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, 

19 MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE, 

20 MCP_OAUTH2_TOKEN_CACHE_MIN_TTL, 

21 MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS, 

22 MCP_PER_USER_TOKEN_DEFAULT_TTL, 

23 MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS, 

24 MCP_PER_USER_TOKEN_REDIS_KEY_PREFIX, 

25) 

26from litellm.llms.custom_httpx.http_handler import get_async_httpx_client 

27from litellm.proxy._experimental.mcp_server.oauth_utils import ( 

28 build_upstream_oauth2_token_request, 

29 resolve_upstream_resource, 

30) 

31from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken 

32from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_codec import OAuthTokenCacheCodec 

33from litellm.proxy.common_utils.encrypt_decrypt_utils import ( 

34 decrypt_value_helper, 

35 encrypt_value_helper, 

36) 

37from litellm.types.llms.custom_http import httpxSpecialProvider 

38 

39if TYPE_CHECKING: 39 ↛ 40line 39 didn't jump to line 40 because the condition on line 39 was never true

40 from litellm.types.mcp_server.mcp_server_manager import MCPServer 

41 

42 

43class MCPOAuth2TokenCache(InMemoryCache): 

44 """ 

45 In-memory cache for OAuth2 client_credentials tokens, keyed by the identity of the token 

46 request rather than by server_id alone. 

47 

48 A minted token is only reusable for the exact request that produced it. Keying on server_id 

49 alone served a token minted under the previous configuration whenever any of those inputs 

50 changed, so editing scopes, rotating the client secret, or setting ``upstream_resource`` 

51 silently kept handing out a token carrying the old scopes or audience until it expired. The 

52 identity below covers every input ``_fetch_token`` puts on the wire, so a change to any of 

53 them misses the cache and mints afresh. 

54 

55 Inherits from ``InMemoryCache`` for TTL-based storage and eviction. 

56 Adds a per-identity ``asyncio.Lock`` to prevent duplicate concurrent fetches. 

57 """ 

58 

59 def __init__(self) -> None: 

60 super().__init__( 

61 max_size_in_memory=MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE, 

62 default_ttl=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, 

63 ) 

64 self._locks: dict[str, asyncio.Lock] = {} 

65 

66 @staticmethod 

67 def _token_identity(server: "MCPServer") -> str: 

68 """Cache key for the token this server's config would mint, prefixed by server_id so a 

69 single server's entries stay greppable and invalidatable. The secret is hashed with the 

70 rest of the identity rather than stored in a key.""" 

71 material: Final = "\x00".join( 

72 ( 

73 server.effective_token_url or "", 

74 server.client_id or "", 

75 server.client_secret or "", 

76 " ".join(server.scopes or ()), 

77 resolve_upstream_resource(server) or "", 

78 server.token_endpoint_auth_method or "", 

79 ) 

80 ) 

81 return f"{server.server_id}:{hashlib.sha256(material.encode()).hexdigest()}" 

82 

83 def _get_lock(self, identity: str) -> asyncio.Lock: 

84 return self._locks.setdefault(identity, asyncio.Lock()) 

85 

86 @staticmethod 

87 def _has_client_credentials_config(server: "MCPServer") -> bool: 

88 return bool(server.client_id and server.client_secret and server.effective_token_url) 

89 

90 async def async_get_token(self, server: "MCPServer") -> str | None: 

91 """Return a valid access token, fetching or refreshing as needed. 

92 

93 Returns ``None`` when the server lacks client credentials config. 

94 """ 

95 if not server.has_client_credentials: 

96 return None 

97 if not self._has_client_credentials_config(server): 

98 return None 

99 

100 identity: Final = self._token_identity(server) 

101 

102 # Fast path — cached token is still valid 

103 cached = self.get_cache(identity) 

104 if cached is not None: 

105 return cached 

106 

107 # Slow path — acquire per-identity lock then double-check 

108 async with self._get_lock(identity): 

109 cached = self.get_cache(identity) 

110 if cached is not None: 

111 return cached 

112 

113 token, ttl = await self._fetch_token(server) 

114 self.set_cache(identity, token, ttl=ttl) 

115 return token 

116 

117 async def _fetch_token(self, server: "MCPServer") -> tuple[str, int]: 

118 """POST to ``effective_token_url`` with ``grant_type=client_credentials``. 

119 

120 Returns ``(access_token, ttl_seconds)`` where ttl accounts for the 

121 expiry buffer so the cache entry expires before the real token does. 

122 """ 

123 client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) 

124 

125 token_url: Final = server.effective_token_url 

126 if not server.client_id or not server.client_secret or not token_url: 

127 raise ValueError( 

128 f"MCP server '{server.server_id}' missing required OAuth2 fields: " 

129 f"client_id={bool(server.client_id)}, " 

130 f"client_secret={bool(server.client_secret)}, " 

131 f"token_url={bool(token_url)}" 

132 ) 

133 

134 token_request: Final = build_upstream_oauth2_token_request( 

135 server, 

136 auth_method=server.token_endpoint_auth_method, 

137 client_id=server.client_id, 

138 client_secret=server.client_secret, 

139 ) 

140 data: Final[dict[str, str]] = { 

141 "grant_type": "client_credentials", 

142 **token_request.body, 

143 } 

144 if server.scopes: 

145 data["scope"] = " ".join(server.scopes) 

146 

147 verbose_logger.debug( 

148 "Fetching OAuth2 client_credentials token for MCP server %s", 

149 server.server_id, 

150 ) 

151 

152 try: 

153 response: Final = await client.post(token_url, data=data, headers=token_request.headers or None) 

154 response.raise_for_status() 

155 except httpx.HTTPStatusError as exc: 

156 raise ValueError( 

157 f"OAuth2 token request for MCP server '{server.server_id}' " 

158 f"failed with status {exc.response.status_code}" 

159 ) from exc 

160 

161 body: Final = response.json() 

162 

163 if not isinstance(body, dict): 

164 raise ValueError( 

165 f"OAuth2 token response for MCP server '{server.server_id}' " 

166 f"returned non-object JSON (got {type(body).__name__})" 

167 ) 

168 

169 access_token: Final = body.get("access_token") 

170 if not access_token: 

171 raise ValueError(f"OAuth2 token response for MCP server '{server.server_id}' missing 'access_token'") 

172 

173 # Safely parse expires_in — providers may return null or non-numeric values 

174 raw_expires_in: Final = body.get("expires_in") 

175 try: 

176 expires_in = int(raw_expires_in) if raw_expires_in is not None else MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL 

177 except (TypeError, ValueError): 

178 expires_in = MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL 

179 

180 ttl: Final = max( 

181 expires_in - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS, 

182 MCP_OAUTH2_TOKEN_CACHE_MIN_TTL, 

183 ) 

184 

185 verbose_logger.info( 

186 "Fetched OAuth2 token for MCP server %s (expires in %ds)", 

187 server.server_id, 

188 expires_in, 

189 ) 

190 return access_token, ttl 

191 

192 def invalidate(self, server_id: str) -> None: 

193 """Remove every cached token for a server (e.g. after a 401). 

194 

195 Entries are keyed by token identity, so one server can hold more than one entry across a 

196 config change; a 401 invalidates all of them rather than only the current configuration's. 

197 """ 

198 prefix: Final = f"{server_id}:" 

199 for key in [k for k in self.cache_dict if isinstance(k, str) and k.startswith(prefix)]: 

200 self.delete_cache(key) 

201 

202 

203mcp_oauth2_token_cache: Final = MCPOAuth2TokenCache() 

204 

205 

206def _compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) -> int: 

207 """Compute Redis TTL for a per-user token. 

208 

209 Uses server.token_storage_ttl_seconds when configured, capped at the token's 

210 remaining lifetime (expires_in minus the expiry buffer) so a cached entry never 

211 outlives the token itself; otherwise derives TTL from expires_in minus the 

212 expiry buffer; falls back to the default TTL. 

213 """ 

214 lifetime_bound: Final = expires_in - MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS if expires_in is not None else None 

215 if server.token_storage_ttl_seconds is not None: 

216 if lifetime_bound is None: 

217 return max(server.token_storage_ttl_seconds, 1) 

218 return max(min(server.token_storage_ttl_seconds, lifetime_bound), 1) 

219 if lifetime_bound is not None: 

220 return max(lifetime_bound, 1) 

221 return MCP_PER_USER_TOKEN_DEFAULT_TTL 

222 

223 

224class MCPPerUserTokenCache: 

225 """Redis-backed cache for per-user OAuth2 access tokens. 

226 

227 Uses LiteLLM's existing ``user_api_key_cache`` (DualCache with optional 

228 Redis backend). Tokens are NaCl-encrypted with ``encrypt_value_helper`` 

229 before storage so they are safe at rest in Redis. 

230 

231 Redis key format: ``mcp:per_user_token:{user_id}:{server_id}`` 

232 Redis value: ``encrypt_value_helper(access_token)`` — URL-safe base64 

233 """ 

234 

235 def _cache_key(self, user_id: str, server_id: str) -> str: 

236 return f"{MCP_PER_USER_TOKEN_REDIS_KEY_PREFIX}:{user_id}:{server_id}" 

237 

238 def _codec(self) -> OAuthTokenCacheCodec: 

239 return OAuthTokenCacheCodec( 

240 encrypt_value_helper, 

241 lambda blob: decrypt_value_helper(blob, key="mcp_per_user_token", exception_type="debug"), 

242 ) 

243 

244 async def get(self, user_id: str, server_id: str) -> str | None: 

245 token: Final = await self.get_token(user_id, server_id) 

246 return token.access_token if token is not None else None 

247 

248 async def get_token(self, user_id: str, server_id: str) -> OAuthToken | None: 

249 try: 

250 from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 

251 

252 key: Final = self._cache_key(user_id, server_id) 

253 encrypted: Final = await user_api_key_cache.async_get_cache(key) 

254 if encrypted is None: 

255 return None 

256 return self._codec().decode(encrypted) 

257 except Exception as exc: 

258 verbose_logger.debug( 

259 "MCPPerUserTokenCache.get failed for user=%s server=%s: %s", 

260 user_id, 

261 server_id, 

262 exc, 

263 ) 

264 return None 

265 

266 async def set( 

267 self, 

268 user_id: str, 

269 server_id: str, 

270 access_token: str, 

271 ttl: int, 

272 identity_binding_proof: str | None = None, 

273 ) -> None: 

274 """Store NaCl-encrypted access_token in Redis with the given TTL.""" 

275 try: 

276 from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 

277 

278 key: Final = self._cache_key(user_id, server_id) 

279 encrypted: Final = self._codec().encode( 

280 OAuthToken(access_token=access_token, identity_binding_proof=identity_binding_proof) 

281 ) 

282 await user_api_key_cache.async_set_cache(key, encrypted, ttl=ttl) 

283 verbose_logger.debug( 

284 "MCPPerUserTokenCache.set: cached token for user=%s server=%s ttl=%ds", 

285 user_id, 

286 server_id, 

287 ttl, 

288 ) 

289 except Exception as exc: 

290 verbose_logger.debug( 

291 "MCPPerUserTokenCache.set failed for user=%s server=%s: %s", 

292 user_id, 

293 server_id, 

294 exc, 

295 ) 

296 

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

298 """Invalidate the cached token in Redis, here, and in every peer worker's in-memory layer.""" 

299 try: 

300 from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( # noqa: PLC0415 # proxy import cycle 

301 evict_and_broadcast, 

302 ) 

303 from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 

304 

305 key: Final = self._cache_key(user_id, server_id) 

306 await evict_and_broadcast((key,), user_api_key_cache) 

307 except Exception as exc: 

308 verbose_logger.debug( 

309 "MCPPerUserTokenCache.delete failed for user=%s server=%s: %s", 

310 user_id, 

311 server_id, 

312 exc, 

313 ) 

314 

315 

316mcp_per_user_token_cache: Final = MCPPerUserTokenCache() 

317 

318 

319async def resolve_mcp_auth( 

320 server: "MCPServer", 

321 mcp_auth_header: str | dict[str, str] | None = None, 

322) -> str | dict[str, str] | None: 

323 """Resolve the auth value for an MCP server. 

324 

325 Priority: 

326 1. ``mcp_auth_header`` — per-request/per-user override 

327 2. OAuth2 client_credentials token — auto-fetched and cached 

328 3. ``server.authentication_token`` — static token from config/DB 

329 

330 ``resolved_token_header`` answers, for the same two inputs, which header the value belongs in. 

331 """ 

332 if mcp_auth_header: 332 ↛ 333line 332 didn't jump to line 333 because the condition on line 332 was never true

333 return mcp_auth_header 

334 if server.has_client_credentials: 334 ↛ 335line 334 didn't jump to line 335 because the condition on line 334 was never true

335 return await mcp_oauth2_token_cache.async_get_token(server) 

336 return server.authentication_token 

337 

338 

339def resolved_token_header( 

340 server: "MCPServer", 

341 mcp_auth_header: str | Mapping[str, str] | None = None, 

342) -> str | None: 

343 """Which upstream header the value ``resolve_mcp_auth`` just returned belongs in. 

344 

345 ``None`` means keep the auth_type default. A caller-supplied ``mcp_auth_header`` is the caller's 

346 own credential aimed at the slot the upstream normally uses, so it never moves; only the values 

347 the gateway resolved from its own config (the minted M2M token, the static token) follow 

348 ``upstream_token_header``. Same inputs and same branch order as ``resolve_mcp_auth``, so the two 

349 cannot disagree about which case they are in. 

350 """ 

351 return None if mcp_auth_header else server.upstream_token_header