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

1"""The ``client_credentials`` (M2M) arm's token source and retrying bearer auth. 

2 

3Implements the client-credentials behavior contract for the v2 resolver: 

4 

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. 

20 

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""" 

26 

27from __future__ import annotations 

28 

29import asyncio 

30import hashlib 

31import time 

32from collections.abc import AsyncGenerator, Awaitable, Callable, Generator 

33from dataclasses import dataclass 

34from typing import Annotated, Final, Literal 

35 

36import httpx 

37import httpx2 

38from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError 

39from typing_extensions import assert_never 

40 

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) 

57 

58 

59class TokenEndpointSuccess(BaseModel): 

60 """The endpoint returned a JSON object; field validation is the caller's job.""" 

61 

62 model_config = ConfigDict(frozen=True) 

63 tag: Literal["success"] = "success" 

64 body: dict[str, object] 

65 

66 

67class TokenEndpointDenied(BaseModel): 

68 """The endpoint answered but did not grant a token (an HTTP error or a non-JSON body).""" 

69 

70 model_config = ConfigDict(frozen=True) 

71 tag: Literal["denied"] = "denied" 

72 status_code: int 

73 detail: str 

74 

75 

76class TokenEndpointUnreachable(BaseModel): 

77 """The endpoint could not be reached (DNS, TLS, connect/read failure).""" 

78 

79 model_config = ConfigDict(frozen=True) 

80 tag: Literal["unreachable"] = "unreachable" 

81 detail: str 

82 

83 

84TokenEndpointOutcome = Annotated[ 

85 TokenEndpointSuccess | TokenEndpointDenied | TokenEndpointUnreachable, 

86 Field(discriminator="tag"), 

87] 

88 

89M2MTokenEndpointPost = Callable[[str, "dict[str, str]", "dict[str, str]"], Awaitable[TokenEndpointOutcome]] 

90 

91 

92_TOKEN_BODY_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object]) 

93 

94 

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. 

99 

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 

112 

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) 

143 

144 

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 

156 

157 

158def _parse_granted_scopes(raw: object) -> tuple[str, ...] | None: 

159 return tuple(raw.split()) if isinstance(raw, str) and raw else None 

160 

161 

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

163class _PreparedGrant: 

164 """A validated, ready-to-POST grant plus the identity key its token caches under.""" 

165 

166 token_url: str 

167 form: dict[str, str] 

168 headers: dict[str, str] 

169 identity_key: str 

170 

171 

172class ClientCredentialsTokenSource: 

173 """Cached M2M access tokens, one per ``(client identity, server)``. 

174 

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 """ 

180 

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] = {} 

200 

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()) 

210 

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) 

224 

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``. 

227 

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 

247 

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) 

260 

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) 

284 

285 

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}")) 

298 

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 ) 

302 

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 ) 

323 

324 

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() 

339 

340 

341class ClientCredentialsBearerAuth(httpx2.Auth): 

342 """Bearer auth that retries an upstream 401 exactly once with a freshly minted token. 

343 

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 """ 

348 

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 

359 

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 

374 

375 def sync_auth_flow(self, request: httpx2.Request) -> Generator[httpx2.Request, httpx2.Response, None]: 

376 raise RuntimeError("ClientCredentialsBearerAuth only supports async httpx2 clients")