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
« 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.
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"""
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
18import jwt
19from fastapi import HTTPException
20from jwt.types import Options
21from typing_extensions import assert_never
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
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)
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]
53@dataclass(frozen=True, slots=True)
54class VerifiedRefreshToken:
55 refresh_token: str
56 binding_proof: str
59StoredRefreshTokenLoader: TypeAlias = Callable[
60 [str, str, MCPOAuthIdentityBinding], # mutable-ok: Callable parameter syntax requires a list
61 Awaitable[VerifiedRefreshToken | None],
62]
64_RejectionCode: TypeAlias = Literal["oauth_principal_mismatch", "oauth_identity_binding_failed"]
67@dataclass(frozen=True, slots=True)
68class _BindingRejection:
69 code: _RejectionCode
70 description: str
73@dataclass(frozen=True, slots=True)
74class RefreshOwnershipProven:
75 """The gateway itself unwrapped the upstream refresh token from a sealed per-user envelope."""
78@dataclass(frozen=True, slots=True)
79class RefreshTokenPresented:
80 refresh_token: str
83RefreshOwnership: TypeAlias = RefreshOwnershipProven | RefreshTokenPresented | None
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: ...
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
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
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 )
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 )
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
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 )
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
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
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
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)
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()
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)
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
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
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)
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.
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 )