Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py: 22%

180 statements  

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

1"""Per-user OAuth identity binding: verify the upstream OIDC principal matches the LiteLLM caller. 

2 

3Closes the confused-deputy gap where a browser authenticated upstream as one principal produces a 

4token that the relay stores under a different, LiteLLM-authenticated principal: before the token 

5endpoint returns, stores, or caches an exchanged token for an identity-bound server, the id_token 

6is validated (signature via the pinned issuer's JWKS, issuer, audience, expiry) and its principal 

7claim is compared to the caller's trusted LiteLLM identity. Mismatches fail closed in enforce mode 

8and are logged in audit mode. 

9""" 

10 

11import hashlib 

12import hmac 

13import json 

14from collections.abc import Awaitable, Callable, Mapping, Sequence 

15from dataclasses import dataclass 

16from typing import Final, Literal, Protocol, TypeAlias 

17 

18import jwt 

19from fastapi import HTTPException 

20from jwt.types import Options 

21from typing_extensions import assert_never 

22 

23from litellm._logging import verbose_logger 

24from litellm.caching.in_memory_cache import InMemoryCache 

25from litellm.llms.custom_httpx.http_handler import get_async_httpx_client 

26from litellm.types.llms.custom_http import httpxSpecialProvider 

27from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer 

28 

29_ALLOWED_ID_TOKEN_ALGORITHMS: Final = ( 

30 "RS256", 

31 "RS384", 

32 "RS512", 

33 "ES256", 

34 "ES384", 

35 "ES512", 

36 "PS256", 

37 "PS384", 

38 "PS512", 

39) 

40_JWKS_CACHE_TTL_SECONDS: Final = 3600 

41_jwks_cache: Final = InMemoryCache(default_ttl=_JWKS_CACHE_TTL_SECONDS) 

42 

43JwksFetcher: TypeAlias = Callable[ 

44 [MCPOAuthIdentityBinding], # mutable-ok: Callable parameter syntax requires a list 

45 Awaitable[Sequence[Mapping[str, object]]], 

46] 

47CallerPrincipalLoader: TypeAlias = Callable[ 

48 [str, MCPOAuthIdentityBinding], # mutable-ok: Callable parameter syntax requires a list 

49 Awaitable[str | None], 

50] 

51 

52 

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

54class VerifiedRefreshToken: 

55 refresh_token: str 

56 binding_proof: str 

57 

58 

59StoredRefreshTokenLoader: TypeAlias = Callable[ 

60 [str, str, MCPOAuthIdentityBinding], # mutable-ok: Callable parameter syntax requires a list 

61 Awaitable[VerifiedRefreshToken | None], 

62] 

63 

64_RejectionCode: TypeAlias = Literal["oauth_principal_mismatch", "oauth_identity_binding_failed"] 

65 

66 

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

68class _BindingRejection: 

69 code: _RejectionCode 

70 description: str 

71 

72 

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

74class RefreshOwnershipProven: 

75 """The gateway itself unwrapped the upstream refresh token from a sealed per-user envelope.""" 

76 

77 

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

79class RefreshTokenPresented: 

80 refresh_token: str 

81 

82 

83RefreshOwnership: TypeAlias = RefreshOwnershipProven | RefreshTokenPresented | None 

84 

85 

86class BindingValidator(Protocol): 

87 async def __call__( 87 ↛ exitline 87 didn't return from function '__call__' because

88 self, 

89 *, 

90 server: MCPServer, 

91 token_response: Mapping[str, object], 

92 litellm_user_id: str | None, 

93 grant_type: str, 

94 refresh_ownership: RefreshOwnership, 

95 ) -> str | None: ... 

96 

97 

98async def _fetch_issuer_jwks(binding: MCPOAuthIdentityBinding) -> Sequence[Mapping[str, object]]: 

99 jwks_url: Final[str] = binding.jwks_url or await _discover_jwks_url(binding.issuer) 

100 cached: Final = await _jwks_cache.async_get_cache(jwks_url) 

101 if isinstance(cached, list): 

102 return cached 

103 client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) 

104 response: Final = await client.get(jwks_url) 

105 response.raise_for_status() 

106 document: Final = response.json() 

107 keys: Final = document.get("keys") if isinstance(document, dict) else None 

108 if not isinstance(keys, list): 

109 raise TypeError(f"JWKS document at {jwks_url} has no 'keys' array") 

110 await _jwks_cache.async_set_cache(jwks_url, keys, ttl=_JWKS_CACHE_TTL_SECONDS) 

111 return keys 

112 

113 

114async def _discover_jwks_url(issuer: str) -> str: 

115 discovery_url: Final = f"{issuer.rstrip('/')}/.well-known/openid-configuration" 

116 client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) 

117 response: Final = await client.get(discovery_url) 

118 response.raise_for_status() 

119 metadata: Final = response.json() 

120 jwks_uri: Final = metadata.get("jwks_uri") if isinstance(metadata, dict) else None 

121 if not isinstance(jwks_uri, str) or not jwks_uri: 

122 raise ValueError(f"OIDC discovery at {discovery_url} returned no jwks_uri") 

123 return jwks_uri 

124 

125 

126def _select_signing_key(id_token: str, keys: Sequence[Mapping[str, object]]) -> "jwt.PyJWK | _BindingRejection": 

127 header: Final = jwt.get_unverified_header(id_token) 

128 kid: Final = header.get("kid") 

129 for key in keys: 

130 if kid is None or key.get("kid") == kid: 

131 return jwt.PyJWK(dict(key)) # mutable-ok: PyJWT requires a concrete JWK dictionary 

132 return _BindingRejection( 

133 code="oauth_identity_binding_failed", 

134 description=f"id_token signing key (kid={kid!r}) not found in the issuer's JWKS", 

135 ) 

136 

137 

138def _decode_id_token( 

139 id_token: str, 

140 binding: MCPOAuthIdentityBinding, 

141 signing_key: "jwt.PyJWK", 

142) -> "Mapping[str, object] | _BindingRejection": 

143 try: 

144 decode_options: Final[Options] = {"require": ("iss", "exp", "aud", "sub", "iat")} 

145 return jwt.decode( 

146 id_token, 

147 signing_key.key, 

148 algorithms=_ALLOWED_ID_TOKEN_ALGORITHMS, 

149 issuer=binding.issuer, 

150 audience=binding.audiences, 

151 options=decode_options, 

152 ) 

153 except jwt.InvalidTokenError as exc: 

154 return _BindingRejection( 

155 code="oauth_identity_binding_failed", 

156 description=f"id_token validation failed: {exc}", 

157 ) 

158 

159 

160def _upstream_principal( 

161 claims: Mapping[str, object], 

162 binding: MCPOAuthIdentityBinding, 

163) -> "str | _BindingRejection": 

164 principal: Final = claims.get(binding.principal_claim) 

165 if not isinstance(principal, str) or not principal: 

166 return _BindingRejection( 

167 code="oauth_identity_binding_failed", 

168 description=f"id_token has no usable '{binding.principal_claim}' claim", 

169 ) 

170 if ( 

171 binding.principal_claim == "email" 

172 and binding.require_email_verified 

173 and claims.get("email_verified") is not True 

174 ): 

175 return _BindingRejection( 

176 code="oauth_identity_binding_failed", 

177 description="id_token email is not verified (email_verified is not true)", 

178 ) 

179 return principal 

180 

181 

182async def _load_caller_principal(litellm_user_id: str, binding: MCPOAuthIdentityBinding) -> str | None: 

183 if binding.caller_field == "user_id": 

184 return litellm_user_id 

185 from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( # noqa: PLC0415 # inline import avoids a module-load circular import 

186 load_active_user_by_id, 

187 ) 

188 

189 loaded: Final = await load_active_user_by_id(litellm_user_id) 

190 if isinstance(loaded, str): 

191 return None 

192 return loaded.user_email 

193 

194 

195async def _load_stored_refresh_token( 

196 litellm_user_id: str, server_id: str, binding: MCPOAuthIdentityBinding 

197) -> VerifiedRefreshToken | None: 

198 try: 

199 from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # keep database imports lazy 

200 get_user_oauth_credential, 

201 ) 

202 from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 # keep database imports lazy 

203 

204 prisma_client: Final = get_prisma_client_or_throw( 

205 "Database not connected. Cannot verify OAuth refresh token ownership." 

206 ) 

207 cred: Final = await get_user_oauth_credential( 

208 prisma_client=prisma_client, 

209 user_id=litellm_user_id, 

210 server_id=server_id, 

211 ) 

212 if not cred or not await credential_binding_matches(binding, litellm_user_id, server_id, cred): 

213 return None 

214 refresh_token: Final = cred.get("refresh_token") 

215 proof: Final = cred.get("identity_binding_proof") 

216 return VerifiedRefreshToken(refresh_token, proof) if refresh_token and proof else None 

217 except Exception: # noqa: BLE001 # a credential lookup failure must fail closed 

218 return None 

219 

220 

221async def current_binding_proof( 

222 binding: MCPOAuthIdentityBinding, 

223 user_id: str, 

224 server_id: str, 

225 caller_principal_loader: CallerPrincipalLoader = _load_caller_principal, 

226) -> str | None: 

227 principal: Final = await caller_principal_loader(user_id, binding) 

228 if not principal: 

229 return None 

230 return _binding_proof(binding, user_id, server_id, principal) 

231 

232 

233def _binding_proof(binding: MCPOAuthIdentityBinding, user_id: str, server_id: str, principal: str) -> str: 

234 payload: Final = json.dumps( 

235 ("oidc-nonce-v1", server_id, user_id, principal, binding.model_dump(mode="json")), 

236 sort_keys=True, 

237 ) 

238 return hashlib.sha256(payload.encode()).hexdigest() 

239 

240 

241async def credential_binding_matches( 

242 binding: MCPOAuthIdentityBinding, 

243 user_id: str, 

244 server_id: str, 

245 credential: Mapping[str, object], 

246 caller_principal_loader: CallerPrincipalLoader = _load_caller_principal, 

247) -> bool: 

248 stored: Final = credential.get("identity_binding_proof") 

249 if not isinstance(stored, str) or not stored: 

250 return False 

251 expected: Final = await current_binding_proof(binding, user_id, server_id, caller_principal_loader) 

252 return expected is not None and hmac.compare_digest(stored, expected) 

253 

254 

255def _principals_match(upstream: str, caller: str, binding: MCPOAuthIdentityBinding) -> bool: 

256 if binding.principal_claim == "email" or binding.caller_field == "user_email": 

257 return upstream.strip().casefold() == caller.strip().casefold() 

258 return upstream == caller 

259 

260 

261async def _evaluate_refresh_ownership( 

262 binding: MCPOAuthIdentityBinding, 

263 litellm_user_id: str | None, 

264 server_id: str, 

265 refresh_ownership: RefreshOwnership, 

266 stored_refresh_token_loader: StoredRefreshTokenLoader, 

267) -> _BindingRejection | str: 

268 match refresh_ownership: 

269 case RefreshOwnershipProven(): 

270 return _BindingRejection( 

271 code="oauth_identity_binding_failed", 

272 description="an identity envelope alone does not prove upstream principal binding", 

273 ) 

274 case None: 

275 return _BindingRejection( 

276 code="oauth_identity_binding_failed", 

277 description="refresh_token grant without an id_token carries no refresh token to prove ownership of", 

278 ) 

279 case RefreshTokenPresented(refresh_token): 

280 if not litellm_user_id: 

281 return _BindingRejection( 

282 code="oauth_identity_binding_failed", 

283 description="the request carries no resolvable LiteLLM user identity to bind the credential to", 

284 ) 

285 stored: Final = await stored_refresh_token_loader(litellm_user_id, server_id, binding) 

286 if stored is None or not hmac.compare_digest(stored.refresh_token, refresh_token): 

287 return _BindingRejection( 

288 code="oauth_identity_binding_failed", 

289 description="the presented refresh_token is not the caller's stored credential for this server", 

290 ) 

291 return stored.binding_proof 

292 assert_never(refresh_ownership) # pragma: no cover 

293 

294 

295async def _evaluate_binding( 

296 binding: MCPOAuthIdentityBinding, 

297 token_response: Mapping[str, object], 

298 litellm_user_id: str | None, 

299 grant_type: str, 

300 server_id: str, 

301 refresh_ownership: RefreshOwnership, 

302 jwks_fetcher: JwksFetcher, 

303 caller_principal_loader: CallerPrincipalLoader, 

304 stored_refresh_token_loader: StoredRefreshTokenLoader, 

305 expected_nonce: str | None, 

306) -> _BindingRejection | str: 

307 id_token: Final = token_response.get("id_token") 

308 if not isinstance(id_token, str) or not id_token: 

309 if grant_type != "refresh_token": 

310 return _BindingRejection( 

311 code="oauth_identity_binding_failed", 

312 description="the upstream token response carries no id_token to bind the credential to a principal", 

313 ) 

314 return await _evaluate_refresh_ownership( 

315 binding, 

316 litellm_user_id, 

317 server_id, 

318 refresh_ownership, 

319 stored_refresh_token_loader, 

320 ) 

321 if not litellm_user_id: 

322 return _BindingRejection( 

323 code="oauth_identity_binding_failed", 

324 description="the request carries no resolvable LiteLLM user identity to bind the credential to", 

325 ) 

326 try: 

327 keys: Final = await jwks_fetcher(binding) 

328 except Exception as exc: # noqa: BLE001 # a JWKS fetch failure must fail closed, not surface as a 500 

329 return _BindingRejection( 

330 code="oauth_identity_binding_failed", 

331 description=f"could not fetch the issuer's JWKS: {exc}", 

332 ) 

333 try: 

334 signing_key: Final = _select_signing_key(id_token, keys) 

335 except (jwt.PyJWTError, ValueError, TypeError): 

336 return _BindingRejection( 

337 code="oauth_identity_binding_failed", 

338 description="invalid id_token header or issuer signing key", 

339 ) 

340 if isinstance(signing_key, _BindingRejection): 

341 return signing_key 

342 claims: Final = _decode_id_token(id_token, binding, signing_key) 

343 if isinstance(claims, _BindingRejection): 

344 return claims 

345 if grant_type == "authorization_code" and (binding.mode == "enforce" or expected_nonce is not None): 

346 nonce: Final = claims.get("nonce") 

347 if not expected_nonce or not isinstance(nonce, str) or not hmac.compare_digest(nonce, expected_nonce): 

348 return _BindingRejection( 

349 code="oauth_identity_binding_failed", 

350 description="id_token nonce does not match the authenticated authorization transaction", 

351 ) 

352 upstream: Final = _upstream_principal(claims, binding) 

353 if isinstance(upstream, _BindingRejection): 

354 return upstream 

355 caller: Final = await caller_principal_loader(litellm_user_id, binding) 

356 if not caller: 

357 return _BindingRejection( 

358 code="oauth_identity_binding_failed", 

359 description=f"the LiteLLM user has no '{binding.caller_field}' to compare the upstream principal against", 

360 ) 

361 if not _principals_match(upstream, caller, binding): 

362 return _BindingRejection( 

363 code="oauth_principal_mismatch", 

364 description="The browser account does not match the selected credential owner.", 

365 ) 

366 return _binding_proof(binding, litellm_user_id, server_id, caller) 

367 

368 

369async def enforce_oauth_identity_binding( 

370 server: MCPServer, 

371 token_response: Mapping[str, object], 

372 litellm_user_id: str | None, 

373 grant_type: str, 

374 refresh_ownership: RefreshOwnership, 

375 jwks_fetcher: JwksFetcher = _fetch_issuer_jwks, 

376 caller_principal_loader: CallerPrincipalLoader = _load_caller_principal, 

377 stored_refresh_token_loader: StoredRefreshTokenLoader = _load_stored_refresh_token, 

378 expected_nonce: str | None = None, 

379) -> str | None: 

380 """Validate the exchanged token's upstream principal against the LiteLLM caller. 

381 

382 No-op when the server has no binding or it is disabled. In enforce mode a failure raises 403 

383 before the caller returns, stores, or caches the token; in audit mode failures are logged only. 

384 A refresh_token grant without an id_token is allowed only when the presented refresh token matches 

385 the caller's stored credential and that credential was previously identity-validated. 

386 """ 

387 binding: Final = server.oauth_identity_binding 

388 if binding is None or binding.mode not in ("audit", "enforce"): 

389 return 

390 rejection: Final = await _evaluate_binding( 

391 binding=binding, 

392 token_response=token_response, 

393 litellm_user_id=litellm_user_id, 

394 grant_type=grant_type, 

395 server_id=server.server_id, 

396 refresh_ownership=refresh_ownership, 

397 jwks_fetcher=jwks_fetcher, 

398 caller_principal_loader=caller_principal_loader, 

399 stored_refresh_token_loader=stored_refresh_token_loader, 

400 expected_nonce=expected_nonce, 

401 ) 

402 if isinstance(rejection, str): 

403 return rejection if binding.mode == "enforce" else None 

404 if binding.mode == "audit": 

405 verbose_logger.warning( 

406 "oauth_identity_binding audit: server=%s user=%s grant=%s rejected=%s (%s)", 

407 server.server_id, 

408 litellm_user_id, 

409 grant_type, 

410 rejection.code, 

411 rejection.description, 

412 ) 

413 return 

414 raise HTTPException( 

415 status_code=403, 

416 detail={ 

417 "error": rejection.code, 

418 "error_description": rejection.description, 

419 "server_id": server.server_id, 

420 "credential_owner": "caller", 

421 "credential_stored": False, 

422 }, 

423 )