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
« 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.
4Automatically fetches and refreshes access tokens for MCP servers configured
5with ``client_id``, ``client_secret``, and ``token_url``.
6"""
8import asyncio
9import hashlib
10from collections.abc import Mapping
11from typing import TYPE_CHECKING, Final
13import httpx
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
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
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.
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.
55 Inherits from ``InMemoryCache`` for TTL-based storage and eviction.
56 Adds a per-identity ``asyncio.Lock`` to prevent duplicate concurrent fetches.
57 """
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] = {}
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()}"
83 def _get_lock(self, identity: str) -> asyncio.Lock:
84 return self._locks.setdefault(identity, asyncio.Lock())
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)
90 async def async_get_token(self, server: "MCPServer") -> str | None:
91 """Return a valid access token, fetching or refreshing as needed.
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
100 identity: Final = self._token_identity(server)
102 # Fast path — cached token is still valid
103 cached = self.get_cache(identity)
104 if cached is not None:
105 return cached
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
113 token, ttl = await self._fetch_token(server)
114 self.set_cache(identity, token, ttl=ttl)
115 return token
117 async def _fetch_token(self, server: "MCPServer") -> tuple[str, int]:
118 """POST to ``effective_token_url`` with ``grant_type=client_credentials``.
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)
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 )
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)
147 verbose_logger.debug(
148 "Fetching OAuth2 client_credentials token for MCP server %s",
149 server.server_id,
150 )
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
161 body: Final = response.json()
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 )
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'")
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
180 ttl: Final = max(
181 expires_in - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS,
182 MCP_OAUTH2_TOKEN_CACHE_MIN_TTL,
183 )
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
192 def invalidate(self, server_id: str) -> None:
193 """Remove every cached token for a server (e.g. after a 401).
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)
203mcp_oauth2_token_cache: Final = MCPOAuth2TokenCache()
206def _compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) -> int:
207 """Compute Redis TTL for a per-user token.
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
224class MCPPerUserTokenCache:
225 """Redis-backed cache for per-user OAuth2 access tokens.
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.
231 Redis key format: ``mcp:per_user_token:{user_id}:{server_id}``
232 Redis value: ``encrypt_value_helper(access_token)`` — URL-safe base64
233 """
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}"
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 )
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
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
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
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
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 )
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
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 )
316mcp_per_user_token_cache: Final = MCPPerUserTokenCache()
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.
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
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
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.
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