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
« 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.
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.
7Usage in config.yaml:
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
16 # Core signing config
17 issuer: "https://my-litellm.example.com" # optional
18 audience: "mcp" # optional
19 ttl_seconds: 300 # optional
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
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"
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"
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
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"
55 # FR-9: Debug headers
56 debug_headers: false # emit x-litellm-mcp-debug header when true
58 # FR-10: Configurable scopes — explicit list replaces auto-generation
59 allowed_scopes:
60 - "mcp:tools/call"
61 - "mcp:tools/list"
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
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"""
71import base64
72import hashlib
73import os
74import re
75import time
76from collections.abc import Mapping, Sequence
77from typing import TYPE_CHECKING, Any, Final, Optional
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
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
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
100class _OIDCDiscoveryDocument(TypedDict, total=False):
101 jwks_uri: str
104class _JWTDecodeKwargs(TypedDict):
105 algorithms: Sequence[str]
106 options: "Options"
107 audience: NotRequired[str]
108 issuer: NotRequired[str]
111class _DebugHeaderClaims(TypedDict, total=False):
112 sub: ReadOnly[object]
113 iss: ReadOnly[object]
114 exp: ReadOnly[object]
115 scope: ReadOnly[str]
118class _SignedClaimSummary(TypedDict):
119 sub: ReadOnly[object]
120 act: ReadOnly[Mapping[str, object]]
121 exp: ReadOnly[object]
124# Module-level singleton for the JWKS discovery endpoint to access.
125_mcp_jwt_signer_instance: Optional["MCPJWTSigner"] = None
127_MCP_JWT_CALL_TYPES: Final = frozenset({"call_mcp_tool", "list_mcp_tools"})
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
134def get_mcp_jwt_signer() -> Optional["MCPJWTSigner"]:
135 """Return the active MCPJWTSigner singleton, or None if not initialized."""
136 return _mcp_jwt_signer_instance
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)
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 )
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")
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]
176async def _fetch_jwks(jwks_uri: str) -> Sequence[Mapping[str, object]]:
177 """
178 Fetch and cache a JWKS from the given URI.
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
189 from litellm.llms.custom_httpx.http_handler import (
190 get_async_httpx_client,
191 httpxSpecialProvider,
192 )
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
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 )
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
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.
222 MCP servers verify tokens using liteLLM's OIDC discovery endpoint and
223 JWKS endpoint rather than trusting each upstream IdP directly.
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
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 """
243 ALGORITHM = "RS256"
244 DEFAULT_TTL = 300
245 DEFAULT_AUDIENCE = "mcp"
246 SIGNING_KEY_ENV = "MCP_JWT_SIGNING_KEY"
248 @classmethod
249 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
250 return [GuardrailEventHooks.pre_mcp_call]
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)
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 )
299 self._public_key = self._private_key.public_key()
300 self._kid = _compute_kid(self._public_key)
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
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
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 ]
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 []
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
339 # --- FR-15: Incoming claim validation ---
340 self.required_claims: list[str] = required_claims or []
341 self.optional_claims: list[str] = optional_claims or []
343 # --- FR-9: Debug headers ---
344 self.debug_headers: bool = debug_headers
346 # --- FR-10: Configurable scopes ---
347 self.allowed_scopes: list[str] | None = allowed_scopes
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
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 )
370 # ------------------------------------------------------------------
371 # Public helpers (used by /.well-known/jwks.json endpoint)
372 # ------------------------------------------------------------------
374 @property
375 def jwks_max_age(self) -> int:
376 """
377 Recommended Cache-Control max-age for the JWKS response (seconds).
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
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 }
403 # ------------------------------------------------------------------
404 # FR-5: Verify + re-sign helpers
405 # ------------------------------------------------------------------
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
411 async def _get_oidc_discovery(self) -> _OIDCDiscoveryDocument:
412 """Fetch and cache the OIDC discovery document with a 24-hour TTL.
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 {}
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.
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 )
443 jwks_keys: Final = await _fetch_jwks(jwks_uri)
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")
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
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
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
467 if signing_jwk is None:
468 raise jwt.exceptions.PyJWKSetError(f"No JWKS key matching kid={kid!r} at {jwks_uri!r}")
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"
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
485 if self.verify_issuer:
486 decode_kwargs["issuer"] = self.verify_issuer
488 payload: Final[dict[str, object]] = jwt.decode(raw_token, signing_jwk.key, **decode_kwargs)
489 return payload
491 async def _introspect_opaque_token(self, token: str) -> dict[str, object]:
492 """
493 Perform RFC 7662 token introspection for opaque (non-JWT) tokens.
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 )
504 from litellm.llms.custom_httpx.http_handler import (
505 get_async_httpx_client,
506 httpxSpecialProvider,
507 )
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
523 # ------------------------------------------------------------------
524 # FR-15: Incoming claim validation
525 # ------------------------------------------------------------------
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
538 from fastapi import HTTPException
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 )
552 # ------------------------------------------------------------------
553 # FR-12: End-user identity mapping
554 # ------------------------------------------------------------------
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.
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
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
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
581 elif source == "litellm:user_id":
582 uid = user_api_key_dict.user_id
583 value = str(uid) if uid else None
585 elif source == "litellm:email":
586 email = user_api_key_dict.user_email
587 value = str(email) if email else None
589 elif source == "litellm:end_user_id":
590 eid = user_api_key_dict.end_user_id
591 value = str(eid) if eid else None
593 elif source == "litellm:team_id":
594 tid = user_api_key_dict.team_id
595 value = str(tid) if tid else None
597 else:
598 verbose_proxy_logger.warning("MCPJWTSigner: unknown end_user_claim_source %r — skipping", source)
599 continue
601 if value:
602 return value
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"
610 # ------------------------------------------------------------------
611 # FR-10: Scope building
612 # ------------------------------------------------------------------
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.
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
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)
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)
648 # ------------------------------------------------------------------
649 # FR-13: Claim operations
650 # ------------------------------------------------------------------
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
659 # set_claims: always override (highest priority)
660 claims = {**claims, **self.set_claims}
662 # remove_claims: delete listed keys
663 for k in self.remove_claims:
664 claims.pop(k, None)
666 return claims
668 # ------------------------------------------------------------------
669 # FR-15: optional_claims passthrough
670 # ------------------------------------------------------------------
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
685 # ------------------------------------------------------------------
686 # Core JWT builder
687 # ------------------------------------------------------------------
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.
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 }
714 # sub — resolved via ordered claim sources (FR-12)
715 claims["sub"] = self._resolve_end_user_identity(user_api_key_dict, jwt_claims)
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
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}
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
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)
737 # optional_claims passthrough (FR-15)
738 claims = self._passthrough_optional_claims(claims, jwt_claims)
740 # Claim operations — applied last so admin overrides take effect (FR-13)
741 claims = self._apply_claim_operations(claims)
743 return claims
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).
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 }
765 # ------------------------------------------------------------------
766 # FR-9: Debug header
767 # ------------------------------------------------------------------
769 @staticmethod
770 def _build_debug_header(claims: _DebugHeaderClaims, kid: str) -> str:
771 """
772 Build the x-litellm-mcp-debug header value.
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}"
785 # ------------------------------------------------------------------
786 # Guardrail hook
787 # ------------------------------------------------------------------
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.
801 Signs outbound MCP tool calls and tools/list requests.
802 """
803 if call_type not in _MCP_JWT_CALL_TYPES:
804 return data
806 hook_data: Final = dict(data)
807 if call_type == "list_mcp_tools":
808 hook_data["mcp_tool_name"] = ""
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")
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
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 )
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
849 # ------------------------------------------------------------------
850 # FR-15: Validate required claims
851 # ------------------------------------------------------------------
852 self._validate_required_claims(jwt_claims)
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)
859 signed_token: Final = jwt.encode(
860 claims,
861 self._private_key,
862 algorithm=self.ALGORITHM,
863 headers={"kid": self._kid},
864 )
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 }
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}"
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)
894 hook_data["extra_headers"] = new_headers
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 )
908 return hook_data
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.
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
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 ") :]
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 )
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()
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