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
« 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.
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.
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"""
14from __future__ import annotations
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
25import httpx
26import jwt
27from pydantic import BaseModel, TypeAdapter, ValidationError
28from typing_extensions import assert_never
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
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])
58CLIENT_ASSERTION_TYPE: Final = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer"
59CLIENT_ASSERTION_LIFETIME_SECONDS: Final = 60
62@dataclass(frozen=True, slots=True)
63class ExchangedToken:
64 access_token: str
65 expires_in: int | None
68class _TokenEndpointResponse(BaseModel):
69 access_token: str
70 expires_in: int | None = None
73class TokenEndpointClient:
74 """One authenticated POST to an OAuth token endpoint, returning the minted token as a value."""
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))
124class _KeyGuard:
125 """The per-key single-flight lock plus the invalidation generation that lock protects.
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 """
134 __slots__ = ("__weakref__", "generation", "lock")
136 def __init__(self) -> None:
137 self.lock = asyncio.Lock()
138 self.generation = 0
141class ExchangedTokenCache:
142 """Memoizes the final token string per key, single-flighting concurrent misses on one lock."""
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()
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.
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.
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)
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).
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.
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
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 )
209 def _get(self, cache_key: str, fingerprint: str) -> str | None:
210 """The stored token, or None when absent or minted for different inputs.
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
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
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 )
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
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)
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 )