Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_endpoint.py: 36%

120 statements  

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

1"""An authenticated OAuth token-endpoint call plus a short-lived-token cache. 

2 

3`TokenEndpointClient.fetch` POSTs one grant to a token endpoint, authenticating the gateway as 

4an OAuth client via `client_auth` (RFC 7523 private-key JWT, or `client_secret_post`), and returns 

5the minted token or a typed `CredError`. `ExchangedTokenCache` memoizes the final token string per 

6opaque cache key with per-key single-flight, so concurrent callers share one round-trip and a hit 

7skips the endpoint entirely. 

8 

9Pure v2: no imports from the v1 MCP auth handlers. The multi-leg flows that compose these (ID-JAG, 

10and later token_exchange / client_credentials) live in the resolver arms; this collaborator owns 

11only the single authenticated call and the cache. 

12""" 

13 

14from __future__ import annotations 

15 

16import asyncio 

17import json 

18import time 

19import uuid 

20import weakref 

21from collections.abc import Awaitable, Callable, Mapping 

22from dataclasses import dataclass 

23from typing import Final 

24 

25import httpx 

26import jwt 

27from pydantic import BaseModel, TypeAdapter, ValidationError 

28from typing_extensions import assert_never 

29 

30from litellm._logging import verbose_proxy_logger 

31from litellm.caching.in_memory_cache import InMemoryCache 

32from litellm.constants import ( 

33 MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, 

34 MCP_OAUTH2_TOKEN_CACHE_MIN_TTL, 

35 MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS, 

36 MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE, 

37) 

38from litellm.exceptions import Timeout 

39from litellm.llms.custom_httpx.http_handler import ( 

40 get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # litellm http handler is untyped 

41) 

42from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( 

43 Error, 

44 Ok, 

45 Result, 

46) 

47from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( 

48 ClientAuth, 

49 ClientSecretAuth, 

50 CredError, 

51 PrivateKeyJwtAuth, 

52) 

53from litellm.types.llms.custom_http import httpxSpecialProvider 

54 

55# The cache stores (fingerprint, token); anything else in the slot is treated as absent. 

56_CACHED_ENTRY_ADAPTER: Final[TypeAdapter[tuple[str, str]]] = TypeAdapter(tuple[str, str]) 

57 

58CLIENT_ASSERTION_TYPE: Final = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" 

59CLIENT_ASSERTION_LIFETIME_SECONDS: Final = 60 

60 

61 

62@dataclass(frozen=True, slots=True) 

63class ExchangedToken: 

64 access_token: str 

65 expires_in: int | None 

66 

67 

68class _TokenEndpointResponse(BaseModel): 

69 access_token: str 

70 expires_in: int | None = None 

71 

72 

73class TokenEndpointClient: 

74 """One authenticated POST to an OAuth token endpoint, returning the minted token as a value.""" 

75 

76 async def fetch( 

77 self, 

78 endpoint: str, 

79 client_id: str, 

80 grant_params: Mapping[str, str], 

81 client_auth: ClientAuth, 

82 ) -> Result[ExchangedToken, CredError]: 

83 try: 

84 data: Final = {**grant_params, **_client_auth_params(endpoint, client_id, client_auth)} 

85 except (ValueError, TypeError, NotImplementedError, jwt.PyJWTError): 

86 verbose_proxy_logger.warning("MCP token endpoint %s: could not sign the client assertion", endpoint) 

87 return Error( 

88 CredError.of_misconfigured( 

89 "token exchange failed: could not sign the client assertion; " 

90 "check client_private_key and client_assertion_signing_alg" 

91 ) 

92 ) 

93 try: 

94 raw: Final = await _post_form(endpoint, data) 

95 except httpx.HTTPStatusError as exc: 

96 verbose_proxy_logger.warning( 

97 "MCP token endpoint %s failed with status %s", endpoint, exc.response.status_code 

98 ) 

99 return Error( 

100 CredError.of_upstream_unavailable(f"token exchange failed with status {exc.response.status_code}") 

101 ) 

102 except (httpx.RequestError, Timeout) as exc: 

103 verbose_proxy_logger.warning("MCP token endpoint %s unreachable: %s", endpoint, type(exc).__name__) 

104 return Error( 

105 CredError.of_upstream_unavailable( 

106 f"token exchange failed: token endpoint unreachable ({type(exc).__name__})" 

107 ) 

108 ) 

109 except json.JSONDecodeError: 

110 verbose_proxy_logger.warning("MCP token endpoint %s returned a non-JSON response", endpoint) 

111 return Error( 

112 CredError.of_upstream_unavailable("token exchange failed: token endpoint returned a non-JSON response") 

113 ) 

114 try: 

115 parsed: Final = _TokenEndpointResponse.model_validate(raw) 

116 except ValidationError: 

117 verbose_proxy_logger.warning("MCP token endpoint %s response missing access_token", endpoint) 

118 return Error( 

119 CredError.of_upstream_unavailable("token exchange failed: token endpoint response missing access_token") 

120 ) 

121 return Ok(ExchangedToken(access_token=parsed.access_token, expires_in=parsed.expires_in)) 

122 

123 

124class _KeyGuard: 

125 """The per-key single-flight lock plus the invalidation generation that lock protects. 

126 

127 Both live on one object so their lifetimes cannot diverge. `get_or_compute` binds the guard to 

128 a local for its whole critical section, which keeps the weak map's entry alive for as long as 

129 that compute could still write; an `invalidate` overlapping the compute therefore reaches the 

130 very same object and its bump is guaranteed to be observed. Conversely a guard nobody holds is 

131 collectible precisely because no write is outstanding for it to fence. 

132 """ 

133 

134 __slots__ = ("__weakref__", "generation", "lock") 

135 

136 def __init__(self) -> None: 

137 self.lock = asyncio.Lock() 

138 self.generation = 0 

139 

140 

141class ExchangedTokenCache: 

142 """Memoizes the final token string per key, single-flighting concurrent misses on one lock.""" 

143 

144 def __init__(self) -> None: 

145 self._cache = InMemoryCache( 

146 max_size_in_memory=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE, 

147 default_ttl=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, 

148 ) 

149 self._guards: weakref.WeakValueDictionary[str, _KeyGuard] = weakref.WeakValueDictionary() 

150 

151 async def get_or_compute( 

152 self, 

153 cache_key: str, 

154 compute: Callable[[], Awaitable[Result[ExchangedToken, CredError]]], 

155 *, 

156 fingerprint: str = "", 

157 ) -> Result[str, CredError]: 

158 """The cached token for `cache_key`, minting one when absent. 

159 

160 `fingerprint` lets a caller address a slot by something stable (a principal) while still 

161 guaranteeing the token it gets back was minted for the *current* inputs: a stored entry 

162 whose fingerprint differs reads as a miss and is re-minted over. That keeps eviction 

163 addressable without the key having to encode the credential material it protects. 

164 

165 An `invalidate` landing while `compute` is in flight wins over that compute's write. The 

166 token is still returned to the caller it was minted for, but it is not stored, so the next 

167 resolution re-mints rather than serving a bearer that predates the invalidation for the 

168 rest of its TTL. 

169 """ 

170 cached = self._get(cache_key, fingerprint) 

171 if cached is not None: 

172 return Ok(cached) 

173 guard = self._guard(cache_key) 

174 async with guard.lock: 

175 cached = self._get(cache_key, fingerprint) 

176 if cached is not None: 

177 return Ok(cached) 

178 generation = guard.generation 

179 match await compute(): 

180 case Ok(token): 

181 if guard.generation == generation: 

182 self._store(cache_key, fingerprint, token) 

183 return Ok(token.access_token) 

184 case Error(err): 

185 return Error(err) 

186 

187 def invalidate(self, cache_key: str) -> None: 

188 """Evict one cached token so the next `get_or_compute` re-mints (e.g. after an upstream 401). 

189 

190 Bumping the guard's generation is what makes the eviction stick against a compute already 

191 awaiting the token endpoint: that compute snapshotted the old generation and so skips its 

192 write. No guard means no compute is in flight, since an in-flight one pins its own. 

193 

194 Stays synchronous: callers invalidate from plain `def`s. 

195 """ 

196 self._cache.delete_cache(cache_key) # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped 

197 guard = self._guards.get(cache_key) 

198 if guard is None: 

199 return 

200 guard.generation += 1 

201 

202 def _store(self, cache_key: str, fingerprint: str, token: ExchangedToken) -> None: 

203 self._cache.set_cache( # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped 

204 cache_key, 

205 (fingerprint, token.access_token), 

206 ttl=_cache_ttl_seconds(token.expires_in), 

207 ) 

208 

209 def _get(self, cache_key: str, fingerprint: str) -> str | None: 

210 """The stored token, or None when absent or minted for different inputs. 

211 

212 The fingerprint comparison is what makes a shared slot safe: a mismatch never returns the 

213 other party's token, it just reads as a miss. 

214 """ 

215 value = self._cache.get_cache(cache_key) # pyright: ignore[reportUnknownMemberType,reportUnknownVariableType] # InMemoryCache is untyped; the adapter below is the type gate 

216 try: 

217 stored_fingerprint, token = _CACHED_ENTRY_ADAPTER.validate_python(value) 

218 except ValidationError: 

219 return None 

220 return token if stored_fingerprint == fingerprint else None 

221 

222 def _guard(self, cache_key: str) -> _KeyGuard: 

223 guard = self._guards.get(cache_key) 

224 if guard is None: 

225 guard = _KeyGuard() 

226 self._guards[cache_key] = guard 

227 return guard 

228 

229 

230def _cache_ttl_seconds(expires_in: int | None) -> int: 

231 lifetime: Final = expires_in if expires_in is not None else MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL 

232 return max( 

233 lifetime - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS, 

234 MCP_OAUTH2_TOKEN_CACHE_MIN_TTL, 

235 ) 

236 

237 

238async def _post_form(endpoint: str, data: dict[str, str]) -> object: 

239 # litellm's httpx handler and httpx.Response are only partially typed; the token endpoint 

240 # returns a JSON object that `_TokenEndpointResponse` validates, so the untyped boundary is 

241 # contained here. A non-2xx raises `httpx.HTTPStatusError`, an unreachable endpoint raises 

242 # `httpx.RequestError` (or litellm's `Timeout`, which the handler substitutes for 

243 # `httpx.TimeoutException`), and a non-JSON body raises `json.JSONDecodeError`; `fetch` maps 

244 # each to a CredError. 

245 client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore[reportUnknownVariableType] # litellm http handler is untyped 

246 response = await client.post(endpoint, data=data) # pyright: ignore[reportUnknownMemberType,reportUnknownVariableType] # litellm http handler is untyped 

247 response.raise_for_status() 

248 return response.json() # pyright: ignore[reportAny] # untyped JSON; validated by _TokenEndpointResponse in fetch 

249 

250 

251def _client_auth_params(endpoint: str, client_id: str, client_auth: ClientAuth) -> dict[str, str]: 

252 match client_auth: 

253 case PrivateKeyJwtAuth() as auth: 

254 return { 

255 "client_id": client_id, 

256 "client_assertion_type": CLIENT_ASSERTION_TYPE, 

257 "client_assertion": _client_assertion(endpoint, client_id, auth), 

258 } 

259 case ClientSecretAuth() as auth: 

260 return { 

261 "client_id": client_id, 

262 "client_secret": auth.client_secret.get_secret_value(), 

263 } 

264 assert_never(client_auth) 

265 

266 

267def _client_assertion(endpoint: str, client_id: str, auth: PrivateKeyJwtAuth) -> str: 

268 now: Final = int(time.time()) 

269 return jwt.encode( 

270 { 

271 "iss": client_id, 

272 "sub": client_id, 

273 "aud": endpoint, 

274 "jti": uuid.uuid4().hex, 

275 "iat": now, 

276 "exp": now + CLIENT_ASSERTION_LIFETIME_SECONDS, 

277 }, 

278 auth.private_key.get_secret_value(), 

279 algorithm=auth.signing_alg, 

280 headers={"kid": auth.key_id} if auth.key_id else None, 

281 )