Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py: 17%

341 statements  

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

1""" 

2MCPJWTSigner — Built-in LiteLLM guardrail for zero trust MCP authentication. 

3 

4Signs outbound MCP requests with a LiteLLM-issued RS256 JWT so that MCP servers 

5can trust a single signing authority (liteLLM) instead of every upstream IdP. 

6 

7Usage in config.yaml: 

8 

9 guardrails: 

10 - guardrail_name: "mcp-jwt-signer" 

11 litellm_params: 

12 guardrail: mcp_jwt_signer 

13 mode: "pre_mcp_call" 

14 default_on: true 

15 

16 # Core signing config 

17 issuer: "https://my-litellm.example.com" # optional 

18 audience: "mcp" # optional 

19 ttl_seconds: 300 # optional 

20 

21 # FR-5: Verify + re-sign — validate incoming Bearer token before signing 

22 access_token_discovery_uri: "https://idp.example.com/.well-known/openid-configuration" 

23 token_introspection_endpoint: "https://idp.example.com/introspect" # opaque tokens 

24 verify_issuer: "https://idp.example.com" # expected iss in incoming JWT 

25 verify_audience: "api://my-app" # expected aud in incoming JWT 

26 

27 # FR-12: End-user identity mapping — ordered resolution chain 

28 # Supported: token:<claim>, litellm:user_id, litellm:email, 

29 # litellm:end_user_id, litellm:team_id 

30 end_user_claim_sources: 

31 - "token:sub" 

32 - "token:email" 

33 - "litellm:user_id" 

34 

35 # FR-13: Claim operations 

36 add_claims: # add if key not already present in the JWT 

37 deployment_id: "prod-001" 

38 set_claims: # always set (overrides computed value) 

39 env: "production" 

40 remove_claims: # remove from final JWT 

41 - "nbf" 

42 

43 # FR-14: Two-token model — issue a second JWT for the MCP transport channel 

44 channel_token_audience: "bedrock-gateway" 

45 channel_token_ttl: 60 

46 

47 # FR-15: Incoming claim validation — enforce required IdP claims 

48 required_claims: 

49 - "sub" 

50 - "email" 

51 optional_claims: # pass through from jwt_claims into outbound JWT 

52 - "groups" 

53 - "roles" 

54 

55 # FR-9: Debug headers 

56 debug_headers: false # emit x-litellm-mcp-debug header when true 

57 

58 # FR-10: Configurable scopes — explicit list replaces auto-generation 

59 allowed_scopes: 

60 - "mcp:tools/call" 

61 - "mcp:tools/list" 

62 

63MCP servers verify tokens via: 

64 GET /.well-known/openid-configuration → { jwks_uri: ".../.well-known/jwks.json" } 

65 GET /.well-known/jwks.json → RSA public key in JWKS format 

66 

67Optionally set MCP_JWT_SIGNING_KEY env var (PEM string or file:///path) to use 

68your own RSA keypair. If unset, an RSA-2048 keypair is auto-generated at startup. 

69""" 

70 

71import base64 

72import hashlib 

73import os 

74import re 

75import time 

76from collections.abc import Mapping, Sequence 

77from typing import TYPE_CHECKING, Any, Final, Optional 

78 

79import jwt 

80from cryptography.hazmat.primitives import serialization 

81from cryptography.hazmat.primitives.asymmetric import rsa 

82from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey, RSAPublicKey 

83from typing_extensions import NotRequired, ReadOnly, TypedDict 

84 

85from litellm._logging import verbose_proxy_logger 

86from litellm.caching import DualCache 

87from litellm.integrations.custom_guardrail import ( 

88 CustomGuardrail, 

89 log_guardrail_information, 

90) 

91from litellm.proxy._types import UserAPIKeyAuth 

92from litellm.types.guardrail_base_init import GuardrailBaseInitKwargs 

93from litellm.types.guardrails import GuardrailEventHooks 

94from litellm.types.utils import CallTypesLiteral 

95 

96if TYPE_CHECKING: 96 ↛ 97line 96 didn't jump to line 97 because the condition on line 96 was never true

97 from jwt.types import Options 

98 

99 

100class _OIDCDiscoveryDocument(TypedDict, total=False): 

101 jwks_uri: str 

102 

103 

104class _JWTDecodeKwargs(TypedDict): 

105 algorithms: Sequence[str] 

106 options: "Options" 

107 audience: NotRequired[str] 

108 issuer: NotRequired[str] 

109 

110 

111class _DebugHeaderClaims(TypedDict, total=False): 

112 sub: ReadOnly[object] 

113 iss: ReadOnly[object] 

114 exp: ReadOnly[object] 

115 scope: ReadOnly[str] 

116 

117 

118class _SignedClaimSummary(TypedDict): 

119 sub: ReadOnly[object] 

120 act: ReadOnly[Mapping[str, object]] 

121 exp: ReadOnly[object] 

122 

123 

124# Module-level singleton for the JWKS discovery endpoint to access. 

125_mcp_jwt_signer_instance: Optional["MCPJWTSigner"] = None 

126 

127_MCP_JWT_CALL_TYPES: Final = frozenset({"call_mcp_tool", "list_mcp_tools"}) 

128 

129# Simple in-memory JWKS cache: keyed by JWKS URI → (keys_list, fetched_at). 

130_jwks_cache: Final[dict[str, tuple[Sequence[Mapping[str, object]], float]]] = {} 

131_JWKS_CACHE_TTL: Final = 3600 # 1 hour 

132 

133 

134def get_mcp_jwt_signer() -> Optional["MCPJWTSigner"]: 

135 """Return the active MCPJWTSigner singleton, or None if not initialized.""" 

136 return _mcp_jwt_signer_instance 

137 

138 

139def _load_private_key_from_env(env_var: str) -> RSAPrivateKey: 

140 """Load an RSA private key from an env var (PEM string or file:// path).""" 

141 key_material: Final = os.environ.get(env_var, "") 

142 if not key_material: 

143 raise ValueError(f"MCPJWTSigner: environment variable '{env_var}' is set but empty.") 

144 if key_material.startswith("file://"): 

145 path: Final = key_material[len("file://") :] 

146 with open(path, "rb") as f: 

147 key_bytes = f.read() 

148 else: 

149 key_bytes = key_material.encode("utf-8") 

150 return serialization.load_pem_private_key(key_bytes, password=None) 

151 

152 

153def _generate_rsa_key_pair() -> RSAPrivateKey: 

154 """Generate a new RSA-2048 private key.""" 

155 return rsa.generate_private_key( 

156 public_exponent=65537, 

157 key_size=2048, 

158 ) 

159 

160 

161def _int_to_base64url(n: int) -> str: 

162 """Encode an integer as a base64url string (no padding).""" 

163 byte_length: Final = (n.bit_length() + 7) // 8 

164 return base64.urlsafe_b64encode(n.to_bytes(byte_length, byteorder="big")).rstrip(b"=").decode("ascii") 

165 

166 

167def _compute_kid(public_key: RSAPublicKey) -> str: 

168 """Derive a key ID from the public key's DER encoding (SHA-256, first 16 hex chars).""" 

169 der_bytes: Final = public_key.public_bytes( 

170 encoding=serialization.Encoding.DER, 

171 format=serialization.PublicFormat.SubjectPublicKeyInfo, 

172 ) 

173 return hashlib.sha256(der_bytes).hexdigest()[:16] 

174 

175 

176async def _fetch_jwks(jwks_uri: str) -> Sequence[Mapping[str, object]]: 

177 """ 

178 Fetch and cache a JWKS from the given URI. 

179 

180 Results are cached for _JWKS_CACHE_TTL seconds to avoid hammering the IdP. 

181 """ 

182 now: Final = time.time() 

183 cached: Final = _jwks_cache.get(jwks_uri) 

184 if cached is not None: 

185 keys, fetched_at = cached 

186 if now - fetched_at < _JWKS_CACHE_TTL: 

187 return keys 

188 

189 from litellm.llms.custom_httpx.http_handler import ( 

190 get_async_httpx_client, 

191 httpxSpecialProvider, 

192 ) 

193 

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

195 resp: Final = await client.get(jwks_uri, headers={"Accept": "application/json"}) 

196 resp.raise_for_status() 

197 jwks_body: Final[Mapping[str, Sequence[Mapping[str, object]]]] = resp.json() 

198 fetched_keys: Final = jwks_body.get("keys", []) 

199 _jwks_cache[jwks_uri] = (fetched_keys, now) 

200 return fetched_keys 

201 

202 

203async def _fetch_oidc_discovery(discovery_uri: str) -> _OIDCDiscoveryDocument: 

204 """Fetch an OIDC discovery document and return its parsed JSON.""" 

205 from litellm.llms.custom_httpx.http_handler import ( 

206 get_async_httpx_client, 

207 httpxSpecialProvider, 

208 ) 

209 

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

211 resp: Final = await client.get(discovery_uri, headers={"Accept": "application/json"}) 

212 resp.raise_for_status() 

213 document: Final[_OIDCDiscoveryDocument] = resp.json() 

214 return document 

215 

216 

217class MCPJWTSigner(CustomGuardrail): 

218 """ 

219 Built-in LiteLLM guardrail that signs outbound MCP requests with a 

220 LiteLLM-issued RS256 JWT, enabling zero trust authentication. 

221 

222 MCP servers verify tokens using liteLLM's OIDC discovery endpoint and 

223 JWKS endpoint rather than trusting each upstream IdP directly. 

224 

225 The signed JWT carries: 

226 - iss: LiteLLM issuer identifier 

227 - aud: MCP audience (configurable) 

228 - sub: End-user identity (resolved via end_user_claim_sources, RFC 8693) 

229 - act: Actor/agent identity (team_id or org_id, RFC 8693 delegation) 

230 - scope: Tool-level access scopes (configurable via allowed_scopes) 

231 - iat, exp, nbf: Standard timing claims 

232 

233 Feature set: 

234 FR-5: Verify + re-sign (access_token_discovery_uri, token_introspection_endpoint) 

235 FR-9: Debug headers (debug_headers) 

236 FR-10: Configurable scopes (allowed_scopes) 

237 FR-12: Configurable end-user identity mapping (end_user_claim_sources) 

238 FR-13: Claim operations (add_claims, set_claims, remove_claims) 

239 FR-14: Two-token model (channel_token_audience, channel_token_ttl) 

240 FR-15: Incoming claim validation (required_claims, optional_claims) 

241 """ 

242 

243 ALGORITHM = "RS256" 

244 DEFAULT_TTL = 300 

245 DEFAULT_AUDIENCE = "mcp" 

246 SIGNING_KEY_ENV = "MCP_JWT_SIGNING_KEY" 

247 

248 @classmethod 

249 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: 

250 return [GuardrailEventHooks.pre_mcp_call] 

251 

252 def __init__( 

253 self, 

254 # Core signing config 

255 issuer: str | None = None, 

256 audience: str | None = None, 

257 ttl_seconds: int | None = None, 

258 # FR-5: Verify + re-sign 

259 access_token_discovery_uri: str | None = None, 

260 token_introspection_endpoint: str | None = None, 

261 verify_issuer: str | None = None, 

262 verify_audience: str | None = None, 

263 # FR-12: End-user identity mapping 

264 end_user_claim_sources: list[str] | None = None, 

265 # FR-13: Claim operations 

266 add_claims: Mapping[str, object] | None = None, 

267 set_claims: Mapping[str, object] | None = None, 

268 remove_claims: list[str] | None = None, 

269 # FR-14: Two-token model 

270 channel_token_audience: str | None = None, 

271 channel_token_ttl: int | None = None, 

272 # FR-15: Incoming claim validation 

273 required_claims: list[str] | None = None, 

274 optional_claims: list[str] | None = None, 

275 # FR-9: Debug headers 

276 debug_headers: bool = False, 

277 # FR-10: Configurable scopes 

278 allowed_scopes: list[str] | None = None, 

279 **kwargs: Any, 

280 ) -> None: 

281 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) 

282 base_kwargs: Final[GuardrailBaseInitKwargs] = kwargs 

283 super().__init__(**base_kwargs) 

284 

285 # --- Signing key setup --- 

286 key_material: Final = os.environ.get(self.SIGNING_KEY_ENV) 

287 if key_material: 

288 self._private_key = _load_private_key_from_env(self.SIGNING_KEY_ENV) 

289 self._persistent_key: bool = True 

290 verbose_proxy_logger.info("MCPJWTSigner: loaded RSA key from env var %s", self.SIGNING_KEY_ENV) 

291 else: 

292 self._private_key = _generate_rsa_key_pair() 

293 self._persistent_key = False 

294 verbose_proxy_logger.info( 

295 "MCPJWTSigner: auto-generated RSA-2048 keypair (set %s to use your own key)", 

296 self.SIGNING_KEY_ENV, 

297 ) 

298 

299 self._public_key = self._private_key.public_key() 

300 self._kid = _compute_kid(self._public_key) 

301 

302 # --- Core config --- 

303 self.issuer: str = ( 

304 issuer or os.environ.get("MCP_JWT_ISSUER") or os.environ.get("LITELLM_EXTERNAL_URL") or "litellm" 

305 ) 

306 self.audience: str = audience or os.environ.get("MCP_JWT_AUDIENCE") or self.DEFAULT_AUDIENCE 

307 resolved_ttl: Final = int( 

308 ttl_seconds if ttl_seconds is not None else os.environ.get("MCP_JWT_TTL_SECONDS", str(self.DEFAULT_TTL)) 

309 ) 

310 if resolved_ttl <= 0: 

311 raise ValueError(f"MCPJWTSigner: ttl_seconds must be > 0, got {resolved_ttl}") 

312 self.ttl_seconds: int = resolved_ttl 

313 

314 # --- FR-5: Verify + re-sign --- 

315 self.access_token_discovery_uri: str | None = access_token_discovery_uri 

316 self.token_introspection_endpoint: str | None = token_introspection_endpoint 

317 self.verify_issuer: str | None = verify_issuer 

318 self.verify_audience: str | None = verify_audience 

319 # Cached OIDC discovery document (fetched lazily, TTL = 24 h) 

320 self._oidc_discovery_doc: _OIDCDiscoveryDocument | None = None 

321 self._oidc_discovery_fetched_at: float = 0.0 

322 

323 # --- FR-12: End-user identity mapping --- 

324 # Default chain: try incoming JWT sub, fall back to litellm user_id 

325 self.end_user_claim_sources: list[str] = end_user_claim_sources or [ 

326 "token:sub", 

327 "litellm:user_id", 

328 ] 

329 

330 # --- FR-13: Claim operations --- 

331 self.add_claims: Mapping[str, object] = add_claims or {} 

332 self.set_claims: Mapping[str, object] = set_claims or {} 

333 self.remove_claims: list[str] = remove_claims or [] 

334 

335 # --- FR-14: Two-token model --- 

336 self.channel_token_audience: str | None = channel_token_audience 

337 self.channel_token_ttl: int = channel_token_ttl if channel_token_ttl is not None else self.ttl_seconds 

338 

339 # --- FR-15: Incoming claim validation --- 

340 self.required_claims: list[str] = required_claims or [] 

341 self.optional_claims: list[str] = optional_claims or [] 

342 

343 # --- FR-9: Debug headers --- 

344 self.debug_headers: bool = debug_headers 

345 

346 # --- FR-10: Configurable scopes --- 

347 self.allowed_scopes: list[str] | None = allowed_scopes 

348 

349 # Register singleton for JWKS/OIDC discovery endpoints. 

350 global _mcp_jwt_signer_instance 

351 if _mcp_jwt_signer_instance is not None: 

352 verbose_proxy_logger.warning( 

353 "MCPJWTSigner: replacing existing singleton — previously issued tokens " 

354 "signed with the old key will fail JWKS verification. " 

355 "Avoid configuring multiple mcp_jwt_signer guardrails." 

356 ) 

357 _mcp_jwt_signer_instance = self 

358 

359 verbose_proxy_logger.info( 

360 "MCPJWTSigner initialized: issuer=%s audience=%s ttl=%ds kid=%s verify=%s channel_token=%s debug=%s", 

361 self.issuer, 

362 self.audience, 

363 self.ttl_seconds, 

364 self._kid, 

365 bool(self.access_token_discovery_uri), 

366 bool(self.channel_token_audience), 

367 self.debug_headers, 

368 ) 

369 

370 # ------------------------------------------------------------------ 

371 # Public helpers (used by /.well-known/jwks.json endpoint) 

372 # ------------------------------------------------------------------ 

373 

374 @property 

375 def jwks_max_age(self) -> int: 

376 """ 

377 Recommended Cache-Control max-age for the JWKS response (seconds). 

378 

379 1 hour for persistent keys; 5 minutes for auto-generated keys so MCP 

380 servers re-fetch quickly after a proxy restart. 

381 """ 

382 return 3600 if self._persistent_key else 300 

383 

384 def get_jwks(self) -> Mapping[str, Sequence[Mapping[str, str]]]: 

385 """ 

386 Return the JWKS for the RSA public key. 

387 Used by GET /.well-known/jwks.json so MCP servers can verify tokens. 

388 """ 

389 public_numbers: Final = self._public_key.public_numbers() 

390 return { 

391 "keys": [ 

392 { 

393 "kty": "RSA", 

394 "alg": self.ALGORITHM, 

395 "use": "sig", 

396 "kid": self._kid, 

397 "n": _int_to_base64url(public_numbers.n), 

398 "e": _int_to_base64url(public_numbers.e), 

399 } 

400 ] 

401 } 

402 

403 # ------------------------------------------------------------------ 

404 # FR-5: Verify + re-sign helpers 

405 # ------------------------------------------------------------------ 

406 

407 # 24-hour TTL for the OIDC discovery doc — long enough to avoid hammering 

408 # the IdP, short enough to pick up jwks_uri changes after key rotation. 

409 _OIDC_DISCOVERY_TTL = 86400 

410 

411 async def _get_oidc_discovery(self) -> _OIDCDiscoveryDocument: 

412 """Fetch and cache the OIDC discovery document with a 24-hour TTL. 

413 

414 Only caches when the doc contains a 'jwks_uri' so that a transient or 

415 malformed response doesn't permanently disable JWT verification. 

416 """ 

417 now: Final = time.time() 

418 cache_expired: Final = (now - self._oidc_discovery_fetched_at) >= self._OIDC_DISCOVERY_TTL 

419 if (self._oidc_discovery_doc is None or cache_expired) and self.access_token_discovery_uri: 

420 doc: Final = await _fetch_oidc_discovery(self.access_token_discovery_uri) 

421 if "jwks_uri" in doc: 

422 self._oidc_discovery_doc = doc 

423 self._oidc_discovery_fetched_at = now 

424 else: 

425 return doc 

426 return self._oidc_discovery_doc or {} 

427 

428 async def _verify_incoming_jwt(self, raw_token: str) -> dict[str, object]: 

429 """ 

430 Verify an incoming Bearer JWT against the configured IdP's JWKS. 

431 

432 Returns the verified payload claims dict. 

433 Raises jwt.PyJWTError (or subclass) if verification fails. 

434 """ 

435 discovery: Final = await self._get_oidc_discovery() 

436 jwks_uri: Final = discovery.get("jwks_uri") 

437 if not jwks_uri: 

438 raise ValueError( 

439 "MCPJWTSigner: access_token_discovery_uri discovery document " 

440 f"at {self.access_token_discovery_uri!r} has no 'jwks_uri'." 

441 ) 

442 

443 jwks_keys: Final = await _fetch_jwks(jwks_uri) 

444 

445 # Only read `kid` from the unverified header — never `alg`. 

446 # Reading `alg` from an attacker-controlled header enables algorithm 

447 # confusion attacks (e.g. alg:none, HS256 with the public key as secret). 

448 # The algorithm is determined from the JWKS key entry instead. 

449 unverified_header: Final = jwt.get_unverified_header(raw_token) 

450 kid: Final = unverified_header.get("kid") 

451 

452 # Build a JWKS object and pick the matching key. 

453 # PyJWT's PyJWKSet handles key-type parsing and kid matching correctly. 

454 from jwt import PyJWKSet 

455 

456 try: 

457 jwks_set: Final = PyJWKSet.from_dict({"keys": jwks_keys}) 

458 except Exception as exc: 

459 raise jwt.exceptions.PyJWKSetError(f"Failed to parse JWKS from {jwks_uri!r}: {exc}") from exc 

460 

461 signing_jwk = None 

462 for jwk_obj in jwks_set.keys: 

463 if not kid or jwk_obj.key_id == kid: 

464 signing_jwk = jwk_obj 

465 break 

466 

467 if signing_jwk is None: 

468 raise jwt.exceptions.PyJWKSetError(f"No JWKS key matching kid={kid!r} at {jwks_uri!r}") 

469 

470 # Use the algorithm declared by the JWKS key entry, not the token header. 

471 # PyJWT populates algorithm_name from the key's `alg` field; when absent 

472 # it infers from the key type (RSAPublicKey → RS256). 

473 alg: Final = getattr(signing_jwk, "algorithm_name", None) or "RS256" 

474 

475 decode_options: Final[Options] = {"verify_exp": True} 

476 decode_kwargs: Final[_JWTDecodeKwargs] = { 

477 "algorithms": [alg], 

478 "options": decode_options, 

479 } 

480 if self.verify_audience: 

481 decode_kwargs["audience"] = self.verify_audience 

482 else: 

483 decode_options["verify_aud"] = False 

484 

485 if self.verify_issuer: 

486 decode_kwargs["issuer"] = self.verify_issuer 

487 

488 payload: Final[dict[str, object]] = jwt.decode(raw_token, signing_jwk.key, **decode_kwargs) 

489 return payload 

490 

491 async def _introspect_opaque_token(self, token: str) -> dict[str, object]: 

492 """ 

493 Perform RFC 7662 token introspection for opaque (non-JWT) tokens. 

494 

495 Returns the introspection response dict. Raises on HTTP error or 

496 inactive token. 

497 """ 

498 if not self.token_introspection_endpoint: 

499 raise ValueError( 

500 "MCPJWTSigner: token_introspection_endpoint is required for " 

501 "opaque token verification but is not configured." 

502 ) 

503 

504 from litellm.llms.custom_httpx.http_handler import ( 

505 get_async_httpx_client, 

506 httpxSpecialProvider, 

507 ) 

508 

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

510 resp: Final = await client.post( 

511 self.token_introspection_endpoint, 

512 data={"token": token}, 

513 headers={"Accept": "application/json"}, 

514 ) 

515 resp.raise_for_status() 

516 result: Final[dict[str, object]] = resp.json() 

517 if not result.get("active", False): 

518 raise jwt.exceptions.ExpiredSignatureError( 

519 "MCPJWTSigner: incoming token is inactive (introspection returned active=false)" 

520 ) 

521 return result 

522 

523 # ------------------------------------------------------------------ 

524 # FR-15: Incoming claim validation 

525 # ------------------------------------------------------------------ 

526 

527 def _validate_required_claims( 

528 self, 

529 jwt_claims: Mapping[str, object] | None, 

530 ) -> None: 

531 """ 

532 Raise HTTP 403 if any required_claims are absent from the verified 

533 incoming token claims. 

534 """ 

535 if not self.required_claims: 

536 return 

537 

538 from fastapi import HTTPException 

539 

540 missing: Final = [c for c in self.required_claims if not (jwt_claims or {}).get(c)] 

541 if missing: 

542 raise HTTPException( 

543 status_code=403, 

544 detail={ 

545 "error": ( 

546 f"MCPJWTSigner: incoming token is missing required claims: " 

547 f"{missing}. Configure the IdP to include these claims." 

548 ) 

549 }, 

550 ) 

551 

552 # ------------------------------------------------------------------ 

553 # FR-12: End-user identity mapping 

554 # ------------------------------------------------------------------ 

555 

556 def _resolve_end_user_identity( 

557 self, 

558 user_api_key_dict: UserAPIKeyAuth, 

559 jwt_claims: Mapping[str, object] | None, 

560 ) -> str: 

561 """ 

562 Resolve the outbound JWT 'sub' using the ordered end_user_claim_sources list. 

563 

564 Supported source prefixes: 

565 token:<claim> — from verified incoming JWT / introspection claims 

566 litellm:user_id — from UserAPIKeyAuth.user_id 

567 litellm:email — from UserAPIKeyAuth.user_email 

568 litellm:end_user_id — from UserAPIKeyAuth.end_user_id 

569 litellm:team_id — from UserAPIKeyAuth.team_id 

570 

571 Falls back to a stable hash of the API token for service-account callers. 

572 """ 

573 for source in self.end_user_claim_sources: 

574 value: str | None = None 

575 

576 if source.startswith("token:"): 

577 claim_name = source[len("token:") :] 

578 raw = (jwt_claims or {}).get(claim_name) 

579 value = str(raw) if raw else None 

580 

581 elif source == "litellm:user_id": 

582 uid = user_api_key_dict.user_id 

583 value = str(uid) if uid else None 

584 

585 elif source == "litellm:email": 

586 email = user_api_key_dict.user_email 

587 value = str(email) if email else None 

588 

589 elif source == "litellm:end_user_id": 

590 eid = user_api_key_dict.end_user_id 

591 value = str(eid) if eid else None 

592 

593 elif source == "litellm:team_id": 

594 tid = user_api_key_dict.team_id 

595 value = str(tid) if tid else None 

596 

597 else: 

598 verbose_proxy_logger.warning("MCPJWTSigner: unknown end_user_claim_source %r — skipping", source) 

599 continue 

600 

601 if value: 

602 return value 

603 

604 # Final fallback for service accounts with no user identity 

605 token: Final = user_api_key_dict.token or user_api_key_dict.api_key 

606 if token: 

607 return "apikey:" + hashlib.sha256(str(token).encode()).hexdigest()[:16] 

608 return "litellm-proxy" 

609 

610 # ------------------------------------------------------------------ 

611 # FR-10: Scope building 

612 # ------------------------------------------------------------------ 

613 

614 def _build_scope( 

615 self, 

616 raw_tool_name: str, 

617 call_type: CallTypesLiteral | None = None, 

618 ) -> str: 

619 """ 

620 Build the JWT scope string. 

621 

622 When allowed_scopes is configured: join them verbatim. 

623 Otherwise auto-generate minimal, least-privilege scopes: 

624 - Tool call → mcp:tools/call mcp:tools/<name>:call 

625 - No tool → mcp:tools/list 

626 

627 NOTE: tools/list is intentionally NOT granted on tool-call JWTs to 

628 prevent callers from enumerating tools they didn't ask to use. 

629 Conversely, tools/call is NOT granted on tools/list-only JWTs so an 

630 intercepted list token cannot be replayed to invoke tools. 

631 """ 

632 if self.allowed_scopes is not None: 

633 return " ".join(self.allowed_scopes) 

634 

635 tool_name: Final = re.sub(r"[^a-zA-Z0-9_\-]", "_", raw_tool_name) if raw_tool_name else "" 

636 if tool_name: 

637 scopes = ["mcp:tools/call", f"mcp:tools/{tool_name}:call"] 

638 elif call_type == "call_mcp_tool": 

639 # Tool-call request reached the signer without a tool name (e.g. 

640 # missing mcp_tool_name in hook data). Fall back to a generic 

641 # tools/call scope so the upstream server still accepts the 

642 # invocation rather than rejecting it as a tools/list-only token. 

643 scopes = ["mcp:tools/call"] 

644 else: 

645 scopes = ["mcp:tools/list"] 

646 return " ".join(scopes) 

647 

648 # ------------------------------------------------------------------ 

649 # FR-13: Claim operations 

650 # ------------------------------------------------------------------ 

651 

652 def _apply_claim_operations(self, claims: dict[str, object]) -> dict[str, object]: 

653 """Apply add_claims, set_claims, and remove_claims to the claim dict.""" 

654 # add_claims: insert only when key is absent 

655 for k, v in self.add_claims.items(): 

656 if k not in claims: 

657 claims[k] = v 

658 

659 # set_claims: always override (highest priority) 

660 claims = {**claims, **self.set_claims} 

661 

662 # remove_claims: delete listed keys 

663 for k in self.remove_claims: 

664 claims.pop(k, None) 

665 

666 return claims 

667 

668 # ------------------------------------------------------------------ 

669 # FR-15: optional_claims passthrough 

670 # ------------------------------------------------------------------ 

671 

672 def _passthrough_optional_claims( 

673 self, 

674 claims: dict[str, object], 

675 jwt_claims: Mapping[str, object] | None, 

676 ) -> dict[str, object]: 

677 """Forward optional_claims from verified incoming token into the outbound JWT.""" 

678 if not self.optional_claims or not jwt_claims: 

679 return claims 

680 for claim in self.optional_claims: 

681 if claim in jwt_claims and claim not in claims: 

682 claims[claim] = jwt_claims[claim] 

683 return claims 

684 

685 # ------------------------------------------------------------------ 

686 # Core JWT builder 

687 # ------------------------------------------------------------------ 

688 

689 def _build_claims( 

690 self, 

691 user_api_key_dict: UserAPIKeyAuth, 

692 data: dict, 

693 jwt_claims: Mapping[str, object] | None = None, 

694 call_type: CallTypesLiteral | None = None, 

695 ) -> dict[str, object]: 

696 """ 

697 Build JWT claims for the outbound MCP access token. 

698 

699 Args: 

700 user_api_key_dict: LiteLLM auth context for the current request. 

701 data: Pre-call hook data dict (contains mcp_tool_name etc.). 

702 jwt_claims: Verified incoming IdP claims (FR-5), or LiteLLM-decoded 

703 jwt_claims if available. None for pure API-key requests. 

704 """ 

705 now: Final = int(time.time()) 

706 claims: dict[str, object] = { 

707 "iss": self.issuer, 

708 "aud": self.audience, 

709 "iat": now, 

710 "exp": now + self.ttl_seconds, 

711 "nbf": now, 

712 } 

713 

714 # sub — resolved via ordered claim sources (FR-12) 

715 claims["sub"] = self._resolve_end_user_identity(user_api_key_dict, jwt_claims) 

716 

717 # email passthrough when available from LiteLLM context 

718 user_email: Final = user_api_key_dict.user_email 

719 if user_email: 

720 claims["email"] = user_email 

721 

722 # act — RFC 8693 delegation claim (team/org context) 

723 team_id: Final = user_api_key_dict.team_id 

724 org_id: Final = user_api_key_dict.org_id 

725 act_sub: Final = team_id or org_id or "litellm-proxy" 

726 claims["act"] = {"sub": act_sub} 

727 

728 # end_user_id when set separately from user_id 

729 end_user_id: Final = user_api_key_dict.end_user_id 

730 if end_user_id: 

731 claims["end_user_id"] = end_user_id 

732 

733 # scope (FR-10) 

734 raw_tool_name: Final[str] = data.get("mcp_tool_name", "") 

735 claims["scope"] = self._build_scope(raw_tool_name, call_type=call_type) 

736 

737 # optional_claims passthrough (FR-15) 

738 claims = self._passthrough_optional_claims(claims, jwt_claims) 

739 

740 # Claim operations — applied last so admin overrides take effect (FR-13) 

741 claims = self._apply_claim_operations(claims) 

742 

743 return claims 

744 

745 def _build_channel_token_claims( 

746 self, 

747 base_claims: Mapping[str, object], 

748 ) -> dict[str, object]: 

749 """ 

750 Build claims for the channel token (FR-14 two-token model). 

751 

752 Inherits sub/act/scope from the access token but uses a separate 

753 audience and TTL so the transport layer and resource layer receive 

754 purpose-bound credentials. 

755 """ 

756 now: Final = int(time.time()) 

757 return { 

758 **base_claims, 

759 "aud": self.channel_token_audience, 

760 "iat": now, 

761 "exp": now + self.channel_token_ttl, 

762 "nbf": now, 

763 } 

764 

765 # ------------------------------------------------------------------ 

766 # FR-9: Debug header 

767 # ------------------------------------------------------------------ 

768 

769 @staticmethod 

770 def _build_debug_header(claims: _DebugHeaderClaims, kid: str) -> str: 

771 """ 

772 Build the x-litellm-mcp-debug header value. 

773 

774 Format: v=1; kid=<kid>; sub=<sub>; iss=<iss>; exp=<exp>; scope=<scope> 

775 Scope is truncated to 80 chars for header safety. 

776 """ 

777 sub: Final = claims.get("sub", "") 

778 iss: Final = claims.get("iss", "") 

779 exp: Final = claims.get("exp", 0) 

780 scope = claims.get("scope", "") 

781 if len(scope) > 80: 

782 scope = scope[:77] + "..." 

783 return f"v=1; kid={kid}; sub={sub}; iss={iss}; exp={exp}; scope={scope}" 

784 

785 # ------------------------------------------------------------------ 

786 # Guardrail hook 

787 # ------------------------------------------------------------------ 

788 

789 @log_guardrail_information 

790 async def async_pre_call_hook( 

791 self, 

792 user_api_key_dict: UserAPIKeyAuth, 

793 cache: DualCache, 

794 data: dict, 

795 call_type: CallTypesLiteral, 

796 ) -> Exception | str | dict | None: 

797 """ 

798 Verifies the incoming token (when configured), validates required claims, 

799 then signs an outbound JWT and injects it as the Authorization header. 

800 

801 Signs outbound MCP tool calls and tools/list requests. 

802 """ 

803 if call_type not in _MCP_JWT_CALL_TYPES: 

804 return data 

805 

806 hook_data: Final = dict(data) 

807 if call_type == "list_mcp_tools": 

808 hook_data["mcp_tool_name"] = "" 

809 

810 # ------------------------------------------------------------------ 

811 # FR-5: Verify incoming token before re-signing 

812 # ------------------------------------------------------------------ 

813 jwt_claims: dict[str, object] | None = None 

814 raw_token: Final[str | None] = hook_data.get("incoming_bearer_token") 

815 

816 if self.access_token_discovery_uri and raw_token: 

817 # Three-dot pattern → JWT; otherwise opaque. 

818 is_jwt: Final = raw_token.count(".") == 2 

819 try: 

820 if is_jwt: 

821 jwt_claims = await self._verify_incoming_jwt(raw_token) 

822 elif self.token_introspection_endpoint: 

823 jwt_claims = await self._introspect_opaque_token(raw_token) 

824 else: 

825 verbose_proxy_logger.warning( 

826 "MCPJWTSigner: access_token_discovery_uri is set but the " 

827 "incoming token appears to be opaque and no " 

828 "token_introspection_endpoint is configured. " 

829 "Proceeding without incoming token verification." 

830 ) 

831 except Exception as exc: 

832 verbose_proxy_logger.error("MCPJWTSigner: incoming token verification failed: %s", exc) 

833 from fastapi import HTTPException 

834 

835 raise HTTPException( 

836 status_code=401, 

837 detail={"error": (f"MCPJWTSigner: incoming token verification failed: {exc}")}, 

838 ) 

839 elif not raw_token and self.access_token_discovery_uri: 

840 verbose_proxy_logger.debug( 

841 "MCPJWTSigner: access_token_discovery_uri configured but no Bearer " 

842 "token found in request (API-key auth request — skipping verification)." 

843 ) 

844 

845 # Fall back to LiteLLM-decoded JWT claims (available when proxy uses JWT auth). 

846 if jwt_claims is None: 

847 jwt_claims = user_api_key_dict.jwt_claims 

848 

849 # ------------------------------------------------------------------ 

850 # FR-15: Validate required claims 

851 # ------------------------------------------------------------------ 

852 self._validate_required_claims(jwt_claims) 

853 

854 # ------------------------------------------------------------------ 

855 # Build outbound access token 

856 # ------------------------------------------------------------------ 

857 claims: Final = self._build_claims(user_api_key_dict, hook_data, jwt_claims, call_type=call_type) 

858 

859 signed_token: Final = jwt.encode( 

860 claims, 

861 self._private_key, 

862 algorithm=self.ALGORITHM, 

863 headers={"kid": self._kid}, 

864 ) 

865 

866 # Merge into existing extra_headers — a prior guardrail in the chain may 

867 # have already injected tracing headers or correlation IDs. 

868 existing_headers: Final[dict[str, str]] = hook_data.get("extra_headers") or {} 

869 new_headers: Final[dict[str, str]] = { 

870 **existing_headers, 

871 "Authorization": f"Bearer {signed_token}", 

872 } 

873 

874 # ------------------------------------------------------------------ 

875 # FR-14: Two-token model — channel token 

876 # ------------------------------------------------------------------ 

877 if self.channel_token_audience: 

878 channel_claims: Final = self._build_channel_token_claims(claims) 

879 channel_token: Final = jwt.encode( 

880 channel_claims, 

881 self._private_key, 

882 algorithm=self.ALGORITHM, 

883 headers={"kid": self._kid}, 

884 ) 

885 new_headers["x-mcp-channel-token"] = f"Bearer {channel_token}" 

886 

887 # ------------------------------------------------------------------ 

888 # FR-9: Debug header 

889 # ------------------------------------------------------------------ 

890 if self.debug_headers: 

891 debug_claims: Final[_DebugHeaderClaims] = claims 

892 new_headers["x-litellm-mcp-debug"] = self._build_debug_header(debug_claims, self._kid) 

893 

894 hook_data["extra_headers"] = new_headers 

895 

896 logged_claims: Final[_SignedClaimSummary] = claims 

897 verbose_proxy_logger.debug( 

898 "MCPJWTSigner: signed JWT sub=%s act=%s tool=%s exp=%d verified=%s channel=%s call_type=%s", 

899 logged_claims.get("sub"), 

900 logged_claims.get("act", {}).get("sub"), 

901 hook_data.get("mcp_tool_name"), 

902 logged_claims["exp"], 

903 jwt_claims is not None, 

904 bool(self.channel_token_audience), 

905 call_type, 

906 ) 

907 

908 return hook_data 

909 

910 

911async def inject_mcp_jwt_headers_for_upstream( 

912 user_api_key_dict: UserAPIKeyAuth | None, 

913 extra_headers: dict[str, str] | None = None, 

914 raw_headers: dict[str, str] | None = None, 

915 *, 

916 for_list_tools: bool = False, 

917 mcp_tool_name: str = "", 

918) -> dict[str, str]: 

919 """ 

920 Sign outbound MCP headers when MCPJWTSigner is configured. 

921 

922 Used by tools/list paths that do not go through proxy pre_call_hook. 

923 """ 

924 merged: Final = dict(extra_headers or {}) 

925 signer: Final = get_mcp_jwt_signer() 

926 if signer is None or user_api_key_dict is None: 

927 return merged 

928 

929 normalized_raw: Final = {k.lower(): v for k, v in (raw_headers or {}).items()} 

930 incoming_bearer_token: str | None = None 

931 auth_hdr: Final = normalized_raw.get("authorization", "") 

932 if auth_hdr.lower().startswith("bearer "): 

933 incoming_bearer_token = auth_hdr[len("bearer ") :] 

934 

935 hook_data: Final = { 

936 "mcp_tool_name": "" if for_list_tools else mcp_tool_name, 

937 "incoming_bearer_token": incoming_bearer_token, 

938 "extra_headers": merged, 

939 } 

940 call_type: Final[CallTypesLiteral] = "list_mcp_tools" if for_list_tools else "call_mcp_tool" 

941 try: 

942 from litellm.proxy.proxy_server import ( # noqa: PLC0415 

943 proxy_logging_obj as _proxy_logging, 

944 ) 

945 

946 shared_cache = _proxy_logging.internal_usage_cache.dual_cache if _proxy_logging is not None else DualCache() 

947 except Exception: 

948 shared_cache = DualCache() 

949 

950 result: Final = await signer.async_pre_call_hook( 

951 user_api_key_dict=user_api_key_dict, 

952 cache=shared_cache, 

953 data=hook_data, 

954 call_type=call_type, 

955 ) 

956 if isinstance(result, dict) and result.get("extra_headers"): 

957 merged.update(result["extra_headers"]) 

958 return merged