Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py: 28%
170 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"""The ``client_credentials`` (M2M) arm's token source and retrying bearer auth.
3Implements the client-credentials behavior contract for the v2 resolver:
5- **Acquisition**: POST ``grant_type=client_credentials`` to the configured token endpoint with
6 the configured scopes and (when set) the IdP's ``audience`` parameter, authenticating the
7 client per ``token_endpoint_auth_method`` (RFC 6749 section 2.3.1, shared helper).
8- **Caching**: tokens are cached per ``(client identity, server)`` where the identity key hashes
9 ``token_url`` / ``client_id`` / ``client_secret`` / auth method / scopes / audience — rotating
10 or re-scoping the credentials changes the key, so a stale token can never be served for the
11 new identity (the contract's rotation-invalidation clause).
12- **Expiry**: the cache TTL respects ``expires_in`` minus a skew so an entry lapses before the
13 real token does; a response with no ``expires_in`` is cached briefly
14 (``default_ttl_seconds``), not assumed long-lived. No refresh_token is ever expected.
15- **401 recovery**: ``ClientCredentialsBearerAuth`` retries an upstream request exactly once
16 after a 401 — discard the cached token, mint a fresh one, resend; a second failure surfaces
17 the upstream's own auth error unchanged.
18- **No user context**: nothing here reads a ``Subject``; every caller shares the one client
19 identity.
21The token-endpoint POST is injected (``M2MTokenEndpointPost``) so the grant orchestration is
22testable without a live IdP; ``post_client_credentials_grant`` is the httpx edge. Failures are
23values: the source returns ``Result[OAuthToken, CredError]``; only the httpx edge touches
24exceptions.
25"""
27from __future__ import annotations
29import asyncio
30import hashlib
31import time
32from collections.abc import AsyncGenerator, Awaitable, Callable, Generator
33from dataclasses import dataclass
34from typing import Annotated, Final, Literal
36import httpx
37import httpx2
38from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError
39from typing_extensions import assert_never
41from litellm._logging import verbose_logger
42from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
43 InMemoryTokenCacheBackend,
44 OAuthToken,
45 TokenCacheBackend,
46)
47from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
48 Error,
49 Ok,
50 Result,
51)
52from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
53 ClientCredentialsConfig,
54 CredError,
55 HeaderCarrier,
56)
59class TokenEndpointSuccess(BaseModel):
60 """The endpoint returned a JSON object; field validation is the caller's job."""
62 model_config = ConfigDict(frozen=True)
63 tag: Literal["success"] = "success"
64 body: dict[str, object]
67class TokenEndpointDenied(BaseModel):
68 """The endpoint answered but did not grant a token (an HTTP error or a non-JSON body)."""
70 model_config = ConfigDict(frozen=True)
71 tag: Literal["denied"] = "denied"
72 status_code: int
73 detail: str
76class TokenEndpointUnreachable(BaseModel):
77 """The endpoint could not be reached (DNS, TLS, connect/read failure)."""
79 model_config = ConfigDict(frozen=True)
80 tag: Literal["unreachable"] = "unreachable"
81 detail: str
84TokenEndpointOutcome = Annotated[
85 TokenEndpointSuccess | TokenEndpointDenied | TokenEndpointUnreachable,
86 Field(discriminator="tag"),
87]
89M2MTokenEndpointPost = Callable[[str, "dict[str, str]", "dict[str, str]"], Awaitable[TokenEndpointOutcome]]
92_TOKEN_BODY_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object])
95async def post_client_credentials_grant(
96 url: str, form: dict[str, str], headers: dict[str, str]
97) -> TokenEndpointOutcome:
98 """POST the grant to the token endpoint and classify the transport outcome.
100 The httpx edge: litellm's handler raises ``HTTPStatusError`` itself on a 4xx/5xx, and every
101 field the caller reads comes out of a validated ``TokenEndpointOutcome``.
102 """
103 from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415 # defer heavy handler import to call time
104 get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # handler factory params are coarsely typed
105 )
106 from litellm.proxy._experimental.mcp_server.mcp_debug import ( # noqa: PLC0415 # diagnostics import credential enums through this package
107 describe_upstream_http_failure,
108 describe_upstream_response,
109 safe_upstream_url,
110 )
111 from litellm.types.llms.custom_http import httpxSpecialProvider # noqa: PLC0415 # deferred with the handler import
113 try:
114 client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
115 response: Final = await client.post( # pyright: ignore[reportUnknownMemberType] # handler params are coarsely typed
116 url, headers={"Accept": "application/json", **headers}, data=form
117 )
118 except httpx.HTTPStatusError as status_err:
119 status_code: Final = status_err.response.status_code
120 verbose_logger.warning(
121 "OAuth2 client_credentials token request denied:\n upstream exchange: %s",
122 describe_upstream_http_failure(status_err),
123 )
124 return TokenEndpointDenied(status_code=status_code, detail=f"token endpoint returned HTTP {status_code}")
125 except Exception as exc: # noqa: BLE001 # any transport failure is the same outcome: unreachable
126 verbose_logger.warning(
127 "OAuth2 client_credentials POST %s failed: %s", safe_upstream_url(httpx.URL(url)), type(exc).__name__
128 )
129 return TokenEndpointUnreachable(detail=type(exc).__name__)
130 try:
131 body: Final = _TOKEN_BODY_ADAPTER.validate_json(response.content)
132 except ValidationError:
133 verbose_logger.warning("OAuth2 client_credentials invalid response: %s", describe_upstream_response(response))
134 return TokenEndpointDenied(
135 status_code=response.status_code, detail="token endpoint returned a non-JSON-object body"
136 )
137 access_token: Final = body.get("access_token")
138 if not isinstance(access_token, str) or not access_token:
139 verbose_logger.warning(
140 "OAuth2 client_credentials response has no access token | %s", describe_upstream_response(response)
141 )
142 return TokenEndpointSuccess(body=body)
145def _parse_expires_in(raw: object) -> int | None:
146 if isinstance(raw, bool):
147 return None
148 if isinstance(raw, int):
149 return raw
150 if isinstance(raw, str):
151 try:
152 return int(raw)
153 except ValueError:
154 return None
155 return None
158def _parse_granted_scopes(raw: object) -> tuple[str, ...] | None:
159 return tuple(raw.split()) if isinstance(raw, str) and raw else None
162@dataclass(frozen=True, slots=True)
163class _PreparedGrant:
164 """A validated, ready-to-POST grant plus the identity key its token caches under."""
166 token_url: str
167 form: dict[str, str]
168 headers: dict[str, str]
169 identity_key: str
172class ClientCredentialsTokenSource:
173 """Cached M2M access tokens, one per ``(client identity, server)``.
175 ``get`` serves from the cache while the entry's TTL (derived from ``expires_in`` minus
176 ``expiry_skew_seconds``) holds, fetching under a per-server lock so concurrent misses
177 produce one grant. ``refetch`` is the 401-recovery path: it drops the failed token and
178 mints a fresh one, unless a concurrent caller already replaced it.
179 """
181 def __init__(
182 self,
183 post: M2MTokenEndpointPost = post_client_credentials_grant,
184 *,
185 backend: TokenCacheBackend | None = None,
186 default_ttl_seconds: float = 300.0,
187 expiry_skew_seconds: float = 60.0,
188 min_cache_seconds: float = 10.0,
189 max_locks: int = 1024,
190 clock: Callable[[], float] = time.time,
191 ) -> None:
192 self._post = post
193 self._backend: TokenCacheBackend = backend or InMemoryTokenCacheBackend(clock=clock)
194 self._default_ttl_seconds = default_ttl_seconds
195 self._expiry_skew_seconds = expiry_skew_seconds
196 self._min_cache_seconds = min_cache_seconds
197 self._max_locks = max_locks
198 self._clock = clock
199 self._locks: dict[str, asyncio.Lock] = {}
201 def _lock(self, server_id: str) -> asyncio.Lock:
202 """Per-server single-flight lock, bounded so ephemeral server ids (e.g. the REST tools
203 preview mints a fresh id per call) cannot grow the dict for the life of the process.
204 Evicting the oldest entry while a task still holds it only means a concurrent caller for
205 that server may run its own grant — single-flight is an optimization, not correctness.
206 """
207 if server_id not in self._locks and len(self._locks) >= self._max_locks:
208 self._locks.pop(next(iter(self._locks)), None)
209 return self._locks.setdefault(server_id, asyncio.Lock())
211 async def get(self, server_id: str, config: ClientCredentialsConfig) -> Result[OAuthToken, CredError]:
212 match _prepare_grant(config):
213 case Error(err):
214 return Error(err)
215 case Ok(grant):
216 cached = await self._backend.get(grant.identity_key, server_id)
217 if cached is not None:
218 return Ok(cached)
219 async with self._lock(server_id):
220 cached = await self._backend.get(grant.identity_key, server_id)
221 if cached is not None:
222 return Ok(cached)
223 return await self._fetch_and_cache(server_id, grant)
225 async def refetch(self, server_id: str, config: ClientCredentialsConfig, failed_access_token: str) -> str | None:
226 """Replace a token the upstream just 401'd; returns the fresh bearer value or ``None``.
228 Runs under the same per-server lock as ``get``: if a concurrent caller already replaced
229 the failed token, that replacement is returned without another grant, so a burst of 401s
230 yields one fetch. A failed refetch returns ``None`` and the caller surfaces the
231 upstream's original auth error (the contract's retry-once-then-give-up clause).
232 """
233 match _prepare_grant(config):
234 case Error(_):
235 return None
236 case Ok(grant):
237 async with self._lock(server_id):
238 cached: Final = await self._backend.get(grant.identity_key, server_id)
239 if cached is not None and cached.access_token != failed_access_token:
240 return cached.access_token
241 await self._backend.delete(grant.identity_key, server_id)
242 match await self._fetch_and_cache(server_id, grant):
243 case Ok(token):
244 return token.access_token
245 case Error(_):
246 return None
248 async def _fetch_and_cache(self, server_id: str, grant: _PreparedGrant) -> Result[OAuthToken, CredError]:
249 outcome: Final = await self._post(grant.token_url, grant.form, grant.headers)
250 match outcome:
251 case TokenEndpointUnreachable():
252 return Error(CredError.of_upstream_unavailable(f"OAuth2 token endpoint unreachable: {outcome.detail}"))
253 case TokenEndpointDenied():
254 if outcome.status_code >= 500:
255 return Error(CredError.of_upstream_unavailable(f"OAuth2 token endpoint failed: {outcome.detail}"))
256 return Error(CredError.of_misconfigured(f"OAuth2 client_credentials grant rejected: {outcome.detail}"))
257 case TokenEndpointSuccess():
258 return await self._cache_token(server_id, grant, outcome.body)
259 assert_never(outcome)
261 async def _cache_token(
262 self, server_id: str, grant: _PreparedGrant, body: dict[str, object]
263 ) -> Result[OAuthToken, CredError]:
264 access_token: Final = body.get("access_token")
265 if not isinstance(access_token, str) or not access_token:
266 return Error(CredError.of_misconfigured("OAuth2 token response is missing 'access_token'"))
267 expires_in: Final = _parse_expires_in(body.get("expires_in"))
268 token: Final = OAuthToken(
269 access_token=access_token,
270 expires_at=self._clock() + expires_in if expires_in is not None else None,
271 scopes=_parse_granted_scopes(body.get("scope")) or (),
272 )
273 # The min-cache floor is itself capped at the token's real lifetime, so a token whose
274 # expires_in is below the skew is never served past its actual expiry; a non-positive
275 # expires_in caches nothing (every request re-fetches, serialized by the per-server lock).
276 ttl: Final = (
277 max(expires_in - self._expiry_skew_seconds, min(float(expires_in), self._min_cache_seconds), 0.0)
278 if expires_in is not None
279 else self._default_ttl_seconds
280 )
281 if ttl > 0:
282 await self._backend.set(grant.identity_key, server_id, token, ttl)
283 return Ok(token)
286def _prepare_grant(config: ClientCredentialsConfig) -> Result[_PreparedGrant, CredError]:
287 if not config.client_id or not config.client_secret or not config.token_url:
288 missing: Final = ", ".join(
289 name
290 for name, present in (
291 ("client_id", bool(config.client_id)),
292 ("client_secret", bool(config.client_secret)),
293 ("token_url", bool(config.token_url)),
294 )
295 if not present
296 )
297 return Error(CredError.of_misconfigured(f"client_credentials config is missing: {missing}"))
299 from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( # noqa: PLC0415 # keep package v1-free at import time
300 build_token_endpoint_client_auth,
301 )
303 client_auth: Final = build_token_endpoint_client_auth(
304 auth_method=config.token_endpoint_auth_method,
305 client_id=config.client_id,
306 client_secret=config.client_secret.get_secret_value(),
307 )
308 form: Final = {
309 "grant_type": "client_credentials",
310 **client_auth.body,
311 **({"scope": " ".join(config.scopes)} if config.scopes else {}),
312 **({"audience": config.audience} if config.audience else {}),
313 **({"resource": config.upstream_resource} if config.upstream_resource else {}),
314 }
315 return Ok(
316 _PreparedGrant(
317 token_url=config.token_url,
318 form=form,
319 headers=client_auth.headers,
320 identity_key=_identity_key(config),
321 )
322 )
325def _identity_key(config: ClientCredentialsConfig) -> str:
326 """Hash of everything that names the client identity; any rotation yields a new key."""
327 material: Final = "\n".join(
328 (
329 config.token_url or "",
330 config.client_id or "",
331 config.client_secret.get_secret_value() if config.client_secret else "",
332 config.token_endpoint_auth_method or "",
333 " ".join(config.scopes),
334 config.audience or "",
335 config.upstream_resource or "",
336 )
337 )
338 return hashlib.sha256(material.encode("utf-8")).hexdigest()
341class ClientCredentialsBearerAuth(httpx2.Auth):
342 """Bearer auth that retries an upstream 401 exactly once with a freshly minted token.
344 The initial token was already resolved (so config/IdP failures surfaced as typed errors
345 before any upstream request); ``refetch`` is the source's 401-recovery callback. If the
346 refetch fails, or the retried request 401s again, the upstream's response stands.
347 """
349 def __init__(
350 self,
351 access_token: str,
352 refetch: Callable[[str], Awaitable[str | None]],
353 carrier: HeaderCarrier,
354 ) -> None:
355 self._carrier = carrier
356 self.header_name = carrier.header_name
357 self._access_token = SecretStr(access_token)
358 self._refetch = refetch
360 async def async_auth_flow(self, request: httpx2.Request) -> AsyncGenerator[httpx2.Request, httpx2.Response]:
361 token: Final = self._access_token.get_secret_value()
362 name, value = self._carrier.header(token)
363 request.headers[name] = value
364 response: Final = yield request
365 if response.status_code != 401:
366 return
367 fresh: Final = await self._refetch(token)
368 if fresh is None:
369 return
370 self._access_token = SecretStr(fresh)
371 fresh_name, fresh_value = self._carrier.header(fresh)
372 request.headers[fresh_name] = fresh_value
373 yield request
375 def sync_auth_flow(self, request: httpx2.Request) -> Generator[httpx2.Request, httpx2.Response, None]:
376 raise RuntimeError("ClientCredentialsBearerAuth only supports async httpx2 clients")