diff --git a/docs/my-website/docs/mcp_zero_trust.md b/docs/my-website/docs/mcp_zero_trust.md
new file mode 100644
index 00000000000..282d8b29e9e
--- /dev/null
+++ b/docs/my-website/docs/mcp_zero_trust.md
@@ -0,0 +1,165 @@
+import Tabs from '@theme/Tabs';
+import TabItem from '@theme/TabItem';
+
+# MCP Zero Trust Auth (JWT Signer)
+
+The `MCPJWTSigner` guardrail signs every outbound MCP tool call with a LiteLLM-issued RS256 JWT. MCP servers validate tokens against LiteLLM's JWKS endpoint instead of trusting each upstream IdP directly.
+
+## Architecture
+
+```mermaid
+sequenceDiagram
+ participant Client
+ participant LiteLLM
+ participant JWKS as LiteLLM JWKS
/.well-known/jwks.json
+ participant MCP as MCP Server
+
+ Client->>LiteLLM: tool call (Bearer API key / JWT)
+ Note over LiteLLM: MCPJWTSigner.async_pre_call_hook()
builds RS256 JWT:
sub=user_id, act=team_id,
scope=mcp:tools/{name}:call
+
+ LiteLLM->>MCP: call_tool(args)
Authorization: Bearer
+ MCP->>JWKS: GET /.well-known/jwks.json
+ JWKS-->>MCP: RSA public key (JWKS)
+ MCP->>MCP: verify JWT signature + claims
+ MCP-->>LiteLLM: tool result
+ LiteLLM-->>Client: response
+```
+
+### OIDC Discovery
+
+LiteLLM publishes standard OIDC discovery so MCP servers can find the signing key automatically:
+
+```
+GET /.well-known/openid-configuration
+→ { "jwks_uri": "https:///.well-known/jwks.json", ... }
+
+GET /.well-known/jwks.json
+→ { "keys": [{ "kty": "RSA", "alg": "RS256", "kid": "...", "n": "...", "e": "..." }] }
+```
+
+## Setup
+
+### 1. Enable in `config.yaml`
+
+```yaml title="config.yaml"
+guardrails:
+ - guardrail_name: "mcp-jwt-signer"
+ litellm_params:
+ guardrail: mcp_jwt_signer
+ mode: pre_mcp_call
+ default_on: true
+ issuer: "https://my-litellm.example.com" # optional — defaults to request base URL
+ audience: "mcp" # optional — default: "mcp"
+ ttl_seconds: 300 # optional — default: 300
+```
+
+### 2. (Optional) Bring your own RSA key
+
+If unset, LiteLLM auto-generates an RSA-2048 keypair at startup (lost on restart).
+
+```bash
+# PEM string
+export MCP_JWT_SIGNING_KEY="-----BEGIN RSA PRIVATE KEY-----\n..."
+
+# Or point to a file
+export MCP_JWT_SIGNING_KEY="file:///secrets/mcp-signing-key.pem"
+```
+
+### 3. Build a verified MCP server with FastMCP
+
+[FastMCP](https://gofastmcp.com) has a built-in `JWTVerifier` that fetches LiteLLM's JWKS automatically, handles key rotation, and enforces `iss`/`aud`/`exp` — zero boilerplate.
+
+**Install:**
+```bash
+pip install fastmcp PyJWT cryptography
+```
+
+**`weather_server.py`:**
+```python
+from fastmcp import FastMCP, Context
+from fastmcp.server.auth.providers.jwt import JWTVerifier
+
+LITELLM_BASE_URL = "https://my-litellm.example.com"
+
+# Point JWTVerifier at LiteLLM's JWKS endpoint.
+# It auto-fetches and caches the RSA public key — no key material to manage.
+auth = JWTVerifier(
+ jwks_uri=f"{LITELLM_BASE_URL}/.well-known/jwks.json",
+ issuer=LITELLM_BASE_URL, # must match MCPJWTSigner `issuer:` in config.yaml
+ audience="mcp", # must match MCPJWTSigner `audience:`
+ algorithm="RS256",
+)
+
+mcp = FastMCP("weather-server", auth=auth)
+
+
+@mcp.tool()
+async def get_weather(city: str, ctx: Context) -> str:
+ """Return weather for a city. Caller identity comes from the verified JWT."""
+ caller = ctx.client_id # = JWT `sub` claim (user_id or apikey hash)
+ await ctx.info(f"Request from {caller}")
+ return f"Weather in {city}: sunny, 72°F"
+
+
+if __name__ == "__main__":
+ mcp.run(transport="http", host="0.0.0.0", port=8000)
+```
+
+`ctx.client_id` is populated from the JWT `sub` claim after verification — you get the caller's identity for free with no extra code.
+
+**Wire it into LiteLLM `config.yaml`:**
+```yaml title="config.yaml"
+mcp_servers:
+ - server_name: weather
+ url: http://localhost:8000/mcp
+ transport: http
+
+guardrails:
+ - guardrail_name: mcp-jwt-signer
+ litellm_params:
+ guardrail: mcp_jwt_signer
+ mode: pre_mcp_call
+ default_on: true
+ issuer: "https://my-litellm.example.com"
+ audience: "mcp"
+```
+
+**Run and test:**
+```bash
+# Terminal 1 — start the MCP server
+python weather_server.py
+
+# Terminal 2 — start LiteLLM
+litellm --config config.yaml
+
+# Terminal 3 — call through LiteLLM (JWT is injected automatically)
+curl -X POST http://localhost:4000/mcp/weather/call_tool \
+ -H "Authorization: Bearer $LITELLM_API_KEY" \
+ -H "Content-Type: application/json" \
+ -d '{"name": "get_weather", "arguments": {"city": "San Francisco"}}'
+```
+
+LiteLLM signs the JWT, sends it to the weather server, and FastMCP verifies it in one round-trip. A request without a valid token gets a `401` back from FastMCP before any tool code runs.
+
+## JWT Claims
+
+| Claim | Value | RFC |
+|-------|-------|-----|
+| `iss` | LiteLLM issuer URL | RFC 7519 |
+| `aud` | configured `audience` | RFC 7519 |
+| `sub` | `user_api_key_dict.user_id` | RFC 8693 |
+| `act.sub` | `team_id` → `org_id` → `"litellm-proxy"` | RFC 8693 delegation |
+| `email` | `user_api_key_dict.user_email` (if set) | — |
+| `scope` | `mcp:tools/call mcp:tools/list mcp:tools/{name}:call` | — |
+| `iat`, `exp`, `nbf` | standard timing | RFC 7519 |
+
+## Limitations
+
+- **OpenAPI-backed MCP servers** (`spec_path` set) do not support hook header injection. When `MCPJWTSigner` is active, calls to these servers log a warning and the JWT header is skipped. Use SSE/HTTP transport MCP servers to get full JWT injection.
+- The keypair is **in-memory by default** — rotated on every restart unless `MCP_JWT_SIGNING_KEY` is set. FastMCP's `JWTVerifier` automatically re-fetches JWKS on key ID miss, so rotation is handled transparently.
+
+## Related
+
+- [MCP Guardrails](./mcp_guardrail) — PII masking and blocking for MCP calls
+- [MCP OAuth](./mcp_oauth) — upstream OAuth2 for MCP server access
+- [MCP AWS SigV4](./mcp_aws_sigv4) — AWS-signed requests to MCP servers
diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js
index 1362745a91f..46e1bd041e6 100644
--- a/docs/my-website/sidebars.js
+++ b/docs/my-website/sidebars.js
@@ -636,6 +636,7 @@ const sidebars = {
"mcp_control",
"mcp_cost",
"mcp_guardrail",
+ "mcp_zero_trust",
"mcp_troubleshoot",
]
},
diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py
index af3a715051b..3385e7feef6 100644
--- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py
+++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py
@@ -677,7 +677,60 @@ async def oauth_authorization_server_mcp(
# Alias for standard OpenID discovery
@router.get("/.well-known/openid-configuration")
async def openid_configuration(request: Request):
- return await oauth_authorization_server_mcp(request)
+ response = await oauth_authorization_server_mcp(request)
+
+ # If MCPJWTSigner is active, augment the discovery doc with JWKS fields so
+ # MCP servers and gateways (e.g. AWS Bedrock AgentCore Gateway) can resolve
+ # the signing keys and verify liteLLM-issued tokens.
+ try:
+ from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
+ get_mcp_jwt_signer,
+ )
+
+ signer = get_mcp_jwt_signer()
+ if signer is not None:
+ request_base_url = get_request_base_url(request)
+ if isinstance(response, dict):
+ response = {
+ **response,
+ "jwks_uri": f"{request_base_url}/.well-known/jwks.json",
+ "id_token_signing_alg_values_supported": ["RS256"],
+ }
+ except ImportError:
+ pass
+
+ return response
+
+
+@router.get("/.well-known/jwks.json")
+async def jwks_json(request: Request):
+ """
+ JSON Web Key Set endpoint.
+
+ Returns the RSA public key used by MCPJWTSigner to sign outbound MCP tokens.
+ MCP servers and gateways use this endpoint to verify liteLLM-issued JWTs.
+
+ Returns an empty key set if MCPJWTSigner is not configured.
+ """
+ try:
+ from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
+ get_mcp_jwt_signer,
+ )
+
+ signer = get_mcp_jwt_signer()
+ if signer is not None:
+ return JSONResponse(
+ content=signer.get_jwks(),
+ headers={"Cache-Control": f"public, max-age={signer.jwks_max_age}"},
+ )
+ except ImportError:
+ pass
+
+ # No signer active — return empty key set; short cache so activation is picked up quickly.
+ return JSONResponse(
+ content={"keys": []},
+ headers={"Cache-Control": "public, max-age=60"},
+ )
# Additional legacy pattern support
diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
index 43fe54fdfb7..295e05bc9be 100644
--- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
+++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
@@ -1908,7 +1908,21 @@ class MCPServerManager:
user_api_key_auth: Optional[UserAPIKeyAuth],
proxy_logging_obj: ProxyLogging,
server: MCPServer,
- ):
+ raw_headers: Optional[Dict[str, str]] = None,
+ ) -> Dict[str, Any]:
+ """
+ Run pre-call checks and guardrail hooks for an MCP tool call.
+
+ Returns a dict that may contain:
+ - "arguments": hook-modified tool arguments (only if changed)
+ - "extra_headers": headers injected by pre_mcp_call guardrail hooks
+
+ Args:
+ raw_headers: Raw inbound HTTP headers from the client request.
+ Passed through to guardrail hooks via mcp_raw_headers so that
+ hooks (e.g. MCPJWTSigner) can read headers like X-Channel-Token
+ for the two-token model (FR-14).
+ """
## check if the tool is allowed or banned for the given server
if not self.check_allowed_or_banned_tools(name, server):
raise HTTPException(
@@ -1957,6 +1971,10 @@ class MCPServerManager:
if user_api_key_auth
else None
),
+ # Raw inbound headers — passed through to guardrail hooks so that
+ # hook implementations (e.g. MCPJWTSigner) can read request headers
+ # such as X-Channel-Token for the two-token model (FR-14).
+ "raw_headers": raw_headers,
}
# Create MCP request object for processing
@@ -1969,6 +1987,7 @@ class MCPServerManager:
mcp_request_obj, pre_hook_kwargs
)
+ hook_result: Dict[str, Any] = {}
try:
# Use standard pre_call_hook
modified_data = await proxy_logging_obj.pre_call_hook(
@@ -1984,7 +2003,9 @@ class MCPServerManager:
)
)
if modified_kwargs.get("arguments") != arguments:
- arguments = modified_kwargs["arguments"]
+ hook_result["arguments"] = modified_kwargs["arguments"]
+ if modified_kwargs.get("extra_headers"):
+ hook_result["extra_headers"] = modified_kwargs["extra_headers"]
except (
BlockedPiiEntityError,
@@ -1995,6 +2016,8 @@ class MCPServerManager:
verbose_logger.error(f"Guardrail blocked MCP tool call pre call: {str(e)}")
raise e
+ return hook_result
+
def _create_during_hook_task(
self,
name: str,
@@ -2047,6 +2070,7 @@ class MCPServerManager:
raw_headers: Optional[Dict[str, str]],
proxy_logging_obj: Optional[ProxyLogging],
host_progress_callback: Optional[Callable] = None,
+ hook_extra_headers: Optional[Dict[str, str]] = None,
) -> CallToolResult:
"""
Call a regular MCP tool using the MCP client.
@@ -2061,6 +2085,9 @@ class MCPServerManager:
oauth2_headers: Optional OAuth2 headers
raw_headers: Optional raw headers from the request
proxy_logging_obj: Optional ProxyLogging object for hook integration
+ host_progress_callback: Optional callback for progress updates
+ hook_extra_headers: Optional headers injected by pre_mcp_call guardrail
+ hooks. Merged last (highest priority) into outbound request headers.
Returns:
CallToolResult from the MCP server
@@ -2116,6 +2143,17 @@ class MCPServerManager:
extra_headers = {}
extra_headers.update(mcp_server.static_headers)
+ if hook_extra_headers:
+ if extra_headers is None:
+ extra_headers = {}
+ if "Authorization" in extra_headers and "Authorization" in hook_extra_headers:
+ verbose_logger.warning(
+ "MCPServerManager: hook_extra_headers contains 'Authorization' which will "
+ "overwrite the existing Authorization header set by static_headers or server "
+ "authentication. The hook JWT will take precedence."
+ )
+ extra_headers.update(hook_extra_headers)
+
stdio_env = self._build_stdio_env(mcp_server, raw_headers)
client = await self._create_mcp_client(
@@ -2201,15 +2239,19 @@ class MCPServerManager:
# Allow validation and modification of tool calls before execution
# Using standard pre_call_hook
#########################################################
+ hook_result: Dict[str, Any] = {}
if proxy_logging_obj:
- await self.pre_call_tool_check(
+ hook_result = await self.pre_call_tool_check(
name=name,
arguments=arguments,
server_name=server_name,
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
server=mcp_server,
+ raw_headers=raw_headers,
)
+ if "arguments" in hook_result:
+ arguments = hook_result["arguments"]
# Prepare tasks for during hooks
tasks = []
@@ -2227,8 +2269,16 @@ class MCPServerManager:
# For OpenAPI servers, call the tool handler directly instead of via MCP client
if mcp_server.spec_path:
verbose_logger.debug(
- f"Calling OpenAPI tool {name} directly via HTTP handler"
+ "Calling OpenAPI tool %s directly via HTTP handler", name
)
+ if hook_result.get("extra_headers"):
+ verbose_logger.warning(
+ "pre_mcp_call hook returned extra_headers for OpenAPI-backed "
+ "MCP server '%s' — header injection is not supported for "
+ "OpenAPI servers; headers will be ignored. Use SSE/HTTP "
+ "transport to enable hook header injection.",
+ server_name,
+ )
tasks.append(
asyncio.create_task(
self._call_openapi_tool_handler(mcp_server, name, arguments)
@@ -2247,6 +2297,7 @@ class MCPServerManager:
raw_headers=raw_headers,
proxy_logging_obj=proxy_logging_obj,
host_progress_callback=host_progress_callback,
+ hook_extra_headers=hook_result.get("extra_headers"),
)
# For OpenAPI tools, await outside the client context
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index ecbd7314cd7..240123ba6df 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -2471,6 +2471,10 @@ class UserAPIKeyAuth(
Any
] = None # Expanded created_by user when expand=user is used
end_user_object_permission: Optional[LiteLLM_ObjectPermissionTable] = None
+ # TODO: jwt_claims carries decoded upstream IdP claims (groups, roles, etc.) so
+ # guardrails can forward them into outbound tokens (e.g. MCPJWTSigner). Currently
+ # populated but not yet consumed — forward-compat hook for a follow-up PR.
+ jwt_claims: Optional[Dict] = None
model_config = ConfigDict(arbitrary_types_allowed=True)
diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py
index 376048e7a13..451ed56339d 100644
--- a/litellm/proxy/auth/user_api_key_auth.py
+++ b/litellm/proxy/auth/user_api_key_auth.py
@@ -700,6 +700,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
)
if valid_token is not None:
api_key = valid_token.token or ""
+ valid_token.jwt_claims = jwt_claims
do_standard_jwt_auth = False
# Fall through to virtual key checks
@@ -729,6 +730,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
team_membership: Optional[LiteLLM_TeamMembership] = result.get(
"team_membership", None
)
+ jwt_claims: Optional[dict] = result.get("jwt_claims", None)
global_proxy_spend = await get_global_proxy_spend(
litellm_proxy_admin_name=litellm_proxy_admin_name,
@@ -757,6 +759,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
org_id=org_id,
end_user_id=end_user_id,
parent_otel_span=parent_otel_span,
+ jwt_claims=jwt_claims,
)
valid_token = UserAPIKeyAuth(
@@ -803,6 +806,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
team_metadata=(
team_object.metadata if team_object is not None else None
),
+ jwt_claims=jwt_claims,
)
# Check if model has zero cost - if so, skip all budget checks
diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/__init__.py
new file mode 100644
index 00000000000..94d1e1c043b
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/__init__.py
@@ -0,0 +1,94 @@
+"""MCP JWT Signer guardrail — built-in LiteLLM guardrail for zero trust MCP auth."""
+
+from typing import TYPE_CHECKING, Any
+
+from litellm.types.guardrails import SupportedGuardrailIntegrations
+
+from .mcp_jwt_signer import MCPJWTSigner, get_mcp_jwt_signer
+
+if TYPE_CHECKING:
+ from litellm.types.guardrails import Guardrail, LitellmParams
+
+
+def _get_param(litellm_params: "LitellmParams", key: str, default: Any = None) -> Any:
+ """
+ Extract a config param from litellm_params, checking optional_params first
+ (where YAML extras land) then the top-level object.
+ """
+ optional_params = getattr(litellm_params, "optional_params", None)
+ if optional_params is not None:
+ v = getattr(optional_params, key, None)
+ if v is not None:
+ return v
+ v = getattr(litellm_params, key, None)
+ return v if v is not None else default
+
+
+def initialize_guardrail(
+ litellm_params: "LitellmParams", guardrail: "Guardrail"
+) -> MCPJWTSigner:
+ import litellm
+
+ guardrail_name = guardrail.get("guardrail_name")
+ if not guardrail_name:
+ raise ValueError("MCPJWTSigner guardrail requires a guardrail_name")
+
+ mode = litellm_params.mode
+ if mode != "pre_mcp_call":
+ raise ValueError(
+ f"MCPJWTSigner guardrail '{guardrail_name}' has mode='{mode}' but must use "
+ "mode='pre_mcp_call'. JWT injection only fires for MCP tool calls."
+ )
+
+ signer = MCPJWTSigner(
+ guardrail_name=guardrail_name,
+ event_hook=litellm_params.mode,
+ default_on=litellm_params.default_on,
+ # Core claims
+ issuer=_get_param(litellm_params, "issuer"),
+ audience=_get_param(litellm_params, "audience"),
+ ttl_seconds=_get_param(litellm_params, "ttl_seconds"),
+ # FR-5: inbound token verification
+ access_token_discovery_uri=_get_param(
+ litellm_params, "access_token_discovery_uri"
+ ),
+ access_token_introspection_endpoint=_get_param(
+ litellm_params, "access_token_introspection_endpoint"
+ ),
+ # FR-12: end-user identity mapping
+ end_user_claim_sources=_get_param(litellm_params, "end_user_claim_sources"),
+ # FR-13: claim operations
+ add_claims=_get_param(litellm_params, "add_claims"),
+ set_claims=_get_param(litellm_params, "set_claims"),
+ remove_claims=_get_param(litellm_params, "remove_claims"),
+ # FR-14: two-token model
+ channel_token_header=_get_param(litellm_params, "channel_token_header"),
+ channel_token_discovery_uri=_get_param(
+ litellm_params, "channel_token_discovery_uri"
+ ),
+ channel_token_jwks_uri=_get_param(litellm_params, "channel_token_jwks_uri"),
+ # FR-15: claim validation
+ required_claims=_get_param(litellm_params, "required_claims"),
+ optional_claims=_get_param(litellm_params, "optional_claims"),
+ # FR-9: debug headers
+ debug_header=_get_param(litellm_params, "debug_header", default=True),
+ # FR-10: configurable scope
+ allowed_tools=_get_param(litellm_params, "allowed_tools"),
+ )
+ litellm.logging_callback_manager.add_litellm_callback(signer)
+ return signer
+
+
+guardrail_initializer_registry = {
+ SupportedGuardrailIntegrations.MCP_JWT_SIGNER.value: initialize_guardrail,
+}
+
+guardrail_class_registry = {
+ SupportedGuardrailIntegrations.MCP_JWT_SIGNER.value: MCPJWTSigner,
+}
+
+__all__ = [
+ "MCPJWTSigner",
+ "initialize_guardrail",
+ "get_mcp_jwt_signer",
+]
diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py
new file mode 100644
index 00000000000..ab637dff935
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py
@@ -0,0 +1,758 @@
+"""
+MCPJWTSigner — Built-in LiteLLM guardrail for zero trust MCP authentication.
+
+Signs outbound MCP requests with a LiteLLM-issued RS256 JWT so that MCP servers
+can trust a single signing authority (liteLLM) instead of every upstream IdP.
+
+Full feature config (all params optional):
+
+ guardrails:
+ - guardrail_name: "mcp-jwt-signer"
+ litellm_params:
+ guardrail: mcp_jwt_signer
+ mode: "pre_mcp_call"
+ default_on: true
+ issuer: "https://my-litellm.example.com" # optional
+ audience: "mcp" # optional
+ ttl_seconds: 300 # optional
+
+ # FR-5: Inbound token verification
+ access_token_discovery_uri: "https://login.example.com/.well-known/openid-configuration"
+ access_token_introspection_endpoint: null # for opaque tokens (future)
+
+ # FR-12: End-user identity mapping (ordered; first non-empty wins)
+ end_user_claim_sources: ["sub", "sso_id", "preferred_username", "email"]
+
+ # FR-13: Claim operations (Kong parity)
+ add_claims: {} # add new claims if not already present
+ set_claims: {} # override/set claims
+ remove_claims: [] # strip claims from output JWT
+
+ # FR-14: Two-token model (access token + channel/agent token)
+ channel_token_header: "X-Channel-Token"
+ channel_token_discovery_uri: null # OIDC discovery for channel token IdP
+ channel_token_jwks_uri: null # fallback JWKS URI for channel token IdP
+
+ # FR-15: Incoming claim validation
+ required_claims: [] # claims that MUST be present in incoming JWT
+ optional_claims: [] # informational; no rejection if absent
+
+ # FR-9: Debug headers
+ debug_header: true # emit x-litellm-mcp-debug on outbound requests
+
+ # FR-10: Configurable scope (admin-defined fine-grained tool control)
+ allowed_tools: [] # if non-empty, restricts scope to these tools only
+
+MCP servers verify tokens via:
+ GET /.well-known/openid-configuration → { jwks_uri: ".../.well-known/jwks.json" }
+ GET /.well-known/jwks.json → RSA public key in JWKS format
+
+Optionally set MCP_JWT_SIGNING_KEY env var (PEM string or file:///path) to use
+your own RSA keypair. If unset, an RSA-2048 keypair is auto-generated at startup.
+"""
+
+import base64
+import hashlib
+import json
+import os
+import re
+import time
+from typing import Any, Dict, List, Optional, Union
+
+import jwt
+from cryptography.hazmat.primitives import serialization
+from cryptography.hazmat.primitives.asymmetric import rsa
+from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey
+
+from litellm._logging import verbose_proxy_logger
+from litellm.caching import DualCache
+from litellm.integrations.custom_guardrail import (
+ CustomGuardrail,
+ log_guardrail_information,
+)
+from litellm.proxy._types import UserAPIKeyAuth
+from litellm.types.utils import CallTypesLiteral
+
+# Module-level singleton for the JWKS discovery endpoint to access.
+_mcp_jwt_signer_instance: Optional["MCPJWTSigner"] = None
+
+# Default ordered list of sources to resolve end-user identity (FR-12).
+_DEFAULT_END_USER_CLAIM_SOURCES: List[str] = [
+ "sub",
+ "preferred_username",
+ "email",
+ "user_id",
+]
+
+# UserAPIKeyAuth attribute names that are valid sources for end-user identity.
+# Only these field names trigger step 3 (attribute lookup on the auth object)
+# in _resolve_end_user_identity — prevents spurious matches on mock objects or
+# subclasses that happen to have extra properties with the same name as a JWT claim.
+_USER_API_KEY_IDENTITY_FIELDS: frozenset = frozenset(
+ {
+ "user_id",
+ "end_user_id",
+ "user_email",
+ "org_id",
+ "team_id",
+ }
+)
+
+
+def get_mcp_jwt_signer() -> Optional["MCPJWTSigner"]:
+ """Return the active MCPJWTSigner singleton, or None if not initialized."""
+ return _mcp_jwt_signer_instance
+
+
+# ---------------------------------------------------------------------------
+# Key helpers
+# ---------------------------------------------------------------------------
+
+
+def _load_private_key_from_env(env_var: str) -> RSAPrivateKey:
+ """Load an RSA private key from an env var (PEM string or file:// path)."""
+ key_material = os.environ.get(env_var, "")
+ if not key_material:
+ raise ValueError(
+ f"MCPJWTSigner: environment variable '{env_var}' is set but empty."
+ )
+ if key_material.startswith("file://"):
+ path = key_material[len("file://"):]
+ with open(path, "rb") as f:
+ key_bytes = f.read()
+ else:
+ key_bytes = key_material.encode("utf-8")
+ return serialization.load_pem_private_key(key_bytes, password=None) # type: ignore[return-value]
+
+
+def _generate_rsa_key_pair() -> RSAPrivateKey:
+ """Generate a new RSA-2048 private key."""
+ return rsa.generate_private_key(
+ public_exponent=65537,
+ key_size=2048,
+ )
+
+
+def _int_to_base64url(n: int) -> str:
+ """Encode an integer as a base64url string (no padding)."""
+ byte_length = (n.bit_length() + 7) // 8
+ return (
+ base64.urlsafe_b64encode(n.to_bytes(byte_length, byteorder="big"))
+ .rstrip(b"=")
+ .decode("ascii")
+ )
+
+
+def _compute_kid(public_key: Any) -> str:
+ """Derive a key ID from the public key's DER encoding (SHA-256, first 16 hex chars)."""
+ der_bytes = public_key.public_bytes(
+ encoding=serialization.Encoding.DER,
+ format=serialization.PublicFormat.SubjectPublicKeyInfo,
+ )
+ return hashlib.sha256(der_bytes).hexdigest()[:16]
+
+
+# ---------------------------------------------------------------------------
+# FR-5: OIDC discovery cache (NFR-2: avoid per-request fetches)
+# ---------------------------------------------------------------------------
+
+
+class _OIDCDiscoveryCache:
+ """
+ Cache for OIDC discovery documents and PyJWKClient instances.
+
+ Avoids per-request OIDC discovery + JWKS fetches (NFR-2).
+ Discovery docs are cached indefinitely (they rarely change).
+ PyJWKClient handles its own JWKS refresh internally.
+ """
+
+ def __init__(self) -> None:
+ self._discovery_docs: Dict[str, Dict] = {}
+ self._jwks_clients: Dict[str, Any] = {} # uri -> PyJWKClient
+
+ async def _fetch_discovery_doc(self, discovery_uri: str) -> Dict:
+ if discovery_uri in self._discovery_docs:
+ return self._discovery_docs[discovery_uri]
+
+ from litellm.llms.custom_httpx.http_handler import (
+ get_async_httpx_client,
+ httpxSpecialProvider,
+ )
+
+ http_client = get_async_httpx_client(
+ llm_provider=httpxSpecialProvider.Oauth2Check
+ )
+ response = await http_client.get(discovery_uri)
+ response.raise_for_status()
+ doc: Dict = response.json()
+ self._discovery_docs[discovery_uri] = doc
+ return doc
+
+ async def get_jwks_client(self, discovery_uri: str) -> Any:
+ """Return a PyJWKClient for the given OIDC discovery URI."""
+ from jwt import PyJWKClient # type: ignore[attr-defined]
+
+ if discovery_uri not in self._jwks_clients:
+ doc = await self._fetch_discovery_doc(discovery_uri)
+ jwks_uri = doc.get("jwks_uri")
+ if not jwks_uri:
+ raise ValueError(
+ f"MCPJWTSigner: OIDC discovery doc at '{discovery_uri}' "
+ "does not contain a 'jwks_uri' field."
+ )
+ self._jwks_clients[discovery_uri] = PyJWKClient(jwks_uri)
+ return self._jwks_clients[discovery_uri]
+
+ async def get_issuer(self, discovery_uri: str) -> Optional[str]:
+ """Return the issuer from the OIDC discovery doc."""
+ doc = await self._fetch_discovery_doc(discovery_uri)
+ return doc.get("issuer")
+
+
+# Module-level OIDC cache shared across signer instances.
+_oidc_cache = _OIDCDiscoveryCache()
+
+
+# ---------------------------------------------------------------------------
+# FR-12: End-user identity resolution
+# ---------------------------------------------------------------------------
+
+
+def _resolve_end_user_identity(
+ sources: List[str],
+ jwt_claims: Optional[Dict],
+ user_api_key_dict: UserAPIKeyAuth,
+ raw_headers: Optional[Dict[str, str]],
+) -> Optional[str]:
+ """
+ Resolve end-user identity from an ordered list of sources.
+
+ For each source, tries (in order):
+ 1. The JWT claim with that name (from the incoming token)
+ 2. The request header with that name (case-insensitive)
+ 3. The UserAPIKeyAuth attribute with that name
+
+ Returns the first non-empty string value found, or None.
+ """
+ normalized_headers: Dict[str, str] = {
+ k.lower(): v for k, v in (raw_headers or {}).items()
+ }
+
+ for source in sources:
+ # 1. JWT claim
+ if jwt_claims:
+ value = jwt_claims.get(source)
+ if value:
+ return str(value)
+
+ # 2. Request header (case-insensitive)
+ header_val = normalized_headers.get(source.lower())
+ if header_val:
+ return header_val
+
+ # 3. UserAPIKeyAuth attribute — only for known identity-related fields.
+ # Using an explicit whitelist prevents spurious matches on MagicMock
+ # attributes or extra subclass properties that share a name with a
+ # JWT claim (e.g. mock.sub would return a truthy MagicMock).
+ if source in _USER_API_KEY_IDENTITY_FIELDS:
+ attr_val = getattr(user_api_key_dict, source, None)
+ if attr_val:
+ return str(attr_val)
+
+ return None
+
+
+# ---------------------------------------------------------------------------
+# FR-15: Incoming claim validation
+# ---------------------------------------------------------------------------
+
+
+def _validate_required_claims(
+ jwt_claims: Optional[Dict],
+ required_claims: List[str],
+) -> None:
+ """
+ Validate that all required_claims are present in the incoming JWT claims.
+
+ Raises ValueError if any required claim is missing or if jwt_claims is
+ None/empty but required_claims are configured (not a JWT auth request).
+ """
+ if not required_claims:
+ return
+
+ if not jwt_claims:
+ raise ValueError(
+ f"MCPJWTSigner: required_claims {required_claims} are configured but "
+ "the incoming request has no JWT claims. Required claims are only "
+ "satisfiable from JWT-authenticated requests (not virtual key auth)."
+ )
+
+ missing = [c for c in required_claims if c not in jwt_claims]
+ if missing:
+ raise ValueError(
+ f"MCPJWTSigner: incoming JWT is missing required_claims: {missing}. "
+ f"Present claims: {list(jwt_claims.keys())}"
+ )
+
+
+# ---------------------------------------------------------------------------
+# Main class
+# ---------------------------------------------------------------------------
+
+
+class MCPJWTSigner(CustomGuardrail):
+ """
+ Built-in LiteLLM guardrail that signs outbound MCP requests with a
+ LiteLLM-issued RS256 JWT, enabling zero trust MCP authentication.
+
+ Features:
+ - FR-1/FR-2: RS256 JWT signing
+ - FR-3: Configurable issuer/audience
+ - FR-5: Verify + re-sign using upstream JWT claims (access_token_discovery_uri)
+ - FR-9: Debug headers (x-litellm-mcp-debug)
+ - FR-10: Configurable fine-grained scope (allowed_tools)
+ - FR-11: act claim (RFC 8693 delegation)
+ - FR-12: Configurable end-user identity mapping (end_user_claim_sources)
+ - FR-13: Claim operations (add_claims, set_claims, remove_claims)
+ - FR-14: Two-token model (access + channel token, channel_token_header)
+ - FR-15: Required/optional claim validation
+ """
+
+ ALGORITHM = "RS256"
+ DEFAULT_TTL = 300
+ DEFAULT_AUDIENCE = "mcp"
+ SIGNING_KEY_ENV = "MCP_JWT_SIGNING_KEY"
+
+ def __init__(
+ self,
+ issuer: Optional[str] = None,
+ audience: Optional[str] = None,
+ ttl_seconds: Optional[int] = None,
+ # FR-5: inbound token verification
+ access_token_discovery_uri: Optional[str] = None,
+ access_token_introspection_endpoint: Optional[str] = None,
+ # FR-12: end-user identity mapping
+ end_user_claim_sources: Optional[List[str]] = None,
+ # FR-13: claim operations
+ add_claims: Optional[Dict[str, Any]] = None,
+ set_claims: Optional[Dict[str, Any]] = None,
+ remove_claims: Optional[List[str]] = None,
+ # FR-14: two-token model
+ channel_token_header: Optional[str] = None,
+ channel_token_discovery_uri: Optional[str] = None,
+ channel_token_jwks_uri: Optional[str] = None,
+ # FR-15: claim validation
+ required_claims: Optional[List[str]] = None,
+ optional_claims: Optional[List[str]] = None,
+ # FR-9: debug headers
+ debug_header: bool = True,
+ # FR-10: configurable scope
+ allowed_tools: Optional[List[str]] = None,
+ **kwargs: Any,
+ ) -> None:
+ super().__init__(**kwargs)
+
+ # Key setup
+ key_material = os.environ.get(self.SIGNING_KEY_ENV)
+ if key_material:
+ self._private_key = _load_private_key_from_env(self.SIGNING_KEY_ENV)
+ self._persistent_key: bool = True
+ verbose_proxy_logger.info(
+ "MCPJWTSigner: loaded RSA key from env var %s", self.SIGNING_KEY_ENV
+ )
+ else:
+ self._private_key = _generate_rsa_key_pair()
+ self._persistent_key = False
+ verbose_proxy_logger.info(
+ "MCPJWTSigner: auto-generated RSA-2048 keypair (set %s to use your own key)",
+ self.SIGNING_KEY_ENV,
+ )
+
+ self._public_key = self._private_key.public_key()
+ self._kid = _compute_kid(self._public_key)
+
+ # Core claims
+ self.issuer: str = (
+ issuer
+ or os.environ.get("MCP_JWT_ISSUER")
+ or os.environ.get("LITELLM_EXTERNAL_URL")
+ or "litellm"
+ )
+ self.audience: str = (
+ audience
+ or os.environ.get("MCP_JWT_AUDIENCE")
+ or self.DEFAULT_AUDIENCE
+ )
+ resolved_ttl = int(
+ ttl_seconds
+ if ttl_seconds is not None
+ else os.environ.get("MCP_JWT_TTL_SECONDS", str(self.DEFAULT_TTL))
+ )
+ if resolved_ttl <= 0:
+ raise ValueError(
+ f"MCPJWTSigner: ttl_seconds must be > 0, got {resolved_ttl}"
+ )
+ self.ttl_seconds: int = resolved_ttl
+
+ # FR-5
+ self.access_token_discovery_uri: Optional[str] = access_token_discovery_uri
+ self.access_token_introspection_endpoint: Optional[str] = (
+ access_token_introspection_endpoint
+ )
+
+ # FR-12
+ self.end_user_claim_sources: List[str] = (
+ end_user_claim_sources
+ if end_user_claim_sources is not None
+ else _DEFAULT_END_USER_CLAIM_SOURCES
+ )
+
+ # FR-13
+ self.add_claims: Dict[str, Any] = add_claims or {}
+ self.set_claims: Dict[str, Any] = set_claims or {}
+ self.remove_claims: List[str] = remove_claims or []
+
+ # FR-14
+ self.channel_token_header: str = (
+ channel_token_header or "X-Channel-Token"
+ ).lower()
+ self.channel_token_discovery_uri: Optional[str] = channel_token_discovery_uri
+ self.channel_token_jwks_uri: Optional[str] = channel_token_jwks_uri
+
+ # FR-15
+ self.required_claims: List[str] = required_claims or []
+ self.optional_claims: List[str] = optional_claims or []
+
+ # FR-9
+ self.debug_header: bool = debug_header
+
+ # FR-10
+ self.allowed_tools: List[str] = allowed_tools or []
+
+ # Register singleton so the JWKS endpoint can access it.
+ global _mcp_jwt_signer_instance
+ if _mcp_jwt_signer_instance is not None:
+ verbose_proxy_logger.warning(
+ "MCPJWTSigner: replacing existing singleton — previously issued tokens "
+ "signed with the old key will fail JWKS verification. "
+ "Avoid configuring multiple mcp_jwt_signer guardrails."
+ )
+ _mcp_jwt_signer_instance = self
+
+ verbose_proxy_logger.info(
+ "MCPJWTSigner initialized: issuer=%s audience=%s ttl=%ds kid=%s "
+ "access_token_discovery_uri=%s channel_token_header=%s",
+ self.issuer,
+ self.audience,
+ self.ttl_seconds,
+ self._kid,
+ self.access_token_discovery_uri or "(none — M2M mode)",
+ self.channel_token_header,
+ )
+
+ # ------------------------------------------------------------------
+ # Public helpers (used by /.well-known/jwks.json endpoint)
+ # ------------------------------------------------------------------
+
+ @property
+ def jwks_max_age(self) -> int:
+ """
+ Recommended Cache-Control max-age for the JWKS response (seconds).
+
+ Use 1 hour for persistent keys (loaded from env var) — safe to cache long.
+ Use 5 minutes for auto-generated keys — key rotates on every restart, so
+ MCP servers must re-fetch quickly to avoid verifying with a stale key.
+ """
+ return 3600 if self._persistent_key else 300
+
+ def get_jwks(self) -> Dict[str, Any]:
+ """
+ Return the JWKS (JSON Web Key Set) for the RSA public key.
+ Used by GET /.well-known/jwks.json so MCP servers can verify tokens.
+ """
+ public_numbers = self._public_key.public_numbers()
+ return {
+ "keys": [
+ {
+ "kty": "RSA",
+ "alg": self.ALGORITHM,
+ "use": "sig",
+ "kid": self._kid,
+ "n": _int_to_base64url(public_numbers.n),
+ "e": _int_to_base64url(public_numbers.e),
+ }
+ ]
+ }
+
+ # ------------------------------------------------------------------
+ # FR-14: Channel token verification
+ # ------------------------------------------------------------------
+
+ async def _verify_channel_token(self, channel_token: str) -> Dict[str, Any]:
+ """
+ Decode and verify the channel token JWT. Returns decoded claims.
+
+ If channel_token_discovery_uri or channel_token_jwks_uri is configured,
+ the signature is verified. Otherwise the token is decoded without
+ signature verification (with a warning logged).
+ """
+ if self.channel_token_discovery_uri:
+ jwks_client = await _oidc_cache.get_jwks_client(
+ self.channel_token_discovery_uri
+ )
+ signing_key = jwks_client.get_signing_key_from_jwt(channel_token)
+ return jwt.decode(
+ channel_token,
+ signing_key.key,
+ algorithms=["RS256", "RS384", "RS512", "ES256", "ES384", "ES512"],
+ options={"verify_aud": False},
+ )
+
+ if self.channel_token_jwks_uri:
+ from jwt import PyJWKClient # type: ignore[attr-defined]
+
+ jwks_client = PyJWKClient(self.channel_token_jwks_uri)
+ signing_key = jwks_client.get_signing_key_from_jwt(channel_token)
+ return jwt.decode(
+ channel_token,
+ signing_key.key,
+ algorithms=["RS256", "RS384", "RS512", "ES256", "ES384", "ES512"],
+ options={"verify_aud": False},
+ )
+
+ verbose_proxy_logger.warning(
+ "MCPJWTSigner: channel token present but no channel_token_discovery_uri "
+ "or channel_token_jwks_uri configured — decoding without signature "
+ "verification. Configure one to enable zero trust channel token auth."
+ )
+ return jwt.decode(channel_token, options={"verify_signature": False})
+
+ # ------------------------------------------------------------------
+ # Internal claim building
+ # ------------------------------------------------------------------
+
+ def _build_scope(self, tool_name: str) -> str:
+ """
+ Build the scope claim. FR-10: uses allowed_tools if configured;
+ otherwise auto-generates least-privilege tool-scoped access.
+ """
+ if self.allowed_tools:
+ # Admin-defined fine-grained scope: only allowed tools get scopes.
+ scope_parts = []
+ for allowed_tool in self.allowed_tools:
+ sanitized = re.sub(r"[^a-zA-Z0-9_\-]", "_", allowed_tool)
+ scope_parts.append(f"mcp:tools/{sanitized}:call")
+ scope_parts.append(f"mcp:tools/{sanitized}:list")
+ # tools/list only granted when not in the middle of a specific tool call.
+ if not tool_name:
+ scope_parts.append("mcp:tools/list")
+ return " ".join(sorted(set(scope_parts)))
+
+ # Auto-generated least-privilege scope (original behaviour).
+ if tool_name:
+ return f"mcp:tools/call mcp:tools/{tool_name}:call"
+ return "mcp:tools/call mcp:tools/list"
+
+ def _build_claims(
+ self,
+ user_api_key_dict: UserAPIKeyAuth,
+ data: dict,
+ channel_token_claims: Optional[Dict] = None,
+ ) -> Dict[str, Any]:
+ """
+ Build JWT claims from the authenticated user context and MCP request data.
+ Follows RFC 8693 (OAuth 2.0 Token Exchange) for sub/act semantics.
+
+ When access_token_discovery_uri is configured (FR-5), upstream jwt_claims
+ from the already-validated incoming token are used as the source of truth
+ for end-user identity, rather than the LiteLLM virtual-key profile.
+ """
+ now = int(time.time())
+ claims: Dict[str, Any] = {
+ "iss": self.issuer,
+ "aud": self.audience,
+ "iat": now,
+ "exp": now + self.ttl_seconds,
+ "nbf": now,
+ }
+
+ # FR-5: Use upstream jwt_claims when available (verify + re-sign mode).
+ jwt_claims: Optional[Dict] = getattr(user_api_key_dict, "jwt_claims", None)
+
+ # Raw request headers (for FR-12 header-based identity resolution).
+ raw_headers: Optional[Dict[str, str]] = data.get("mcp_raw_headers")
+
+ # FR-12: Resolve sub (end-user identity) from ordered sources.
+ end_user = _resolve_end_user_identity(
+ self.end_user_claim_sources,
+ jwt_claims,
+ user_api_key_dict,
+ raw_headers,
+ )
+ if end_user:
+ claims["sub"] = end_user
+ else:
+ token = getattr(user_api_key_dict, "token", None) or getattr(
+ user_api_key_dict, "api_key", None
+ )
+ if token:
+ claims["sub"] = "apikey:" + hashlib.sha256(
+ str(token).encode()
+ ).hexdigest()[:16]
+ else:
+ claims["sub"] = "litellm-proxy"
+
+ # Email: prefer jwt_claims, fall back to user profile.
+ email = (jwt_claims or {}).get("email") or getattr(
+ user_api_key_dict, "user_email", None
+ )
+ if email:
+ claims["email"] = email
+
+ # FR-14: Two-token model — act reflects the requester/agent identity.
+ if channel_token_claims:
+ # Channel token present: its sub (or client_id) is the actor.
+ channel_sub = (
+ channel_token_claims.get("sub")
+ or channel_token_claims.get("client_id")
+ or "unknown-agent"
+ )
+ act: Dict[str, Any] = {"sub": channel_sub}
+ if channel_token_claims.get("client_id"):
+ act["client_id"] = channel_token_claims["client_id"]
+ claims["act"] = act
+ else:
+ # Fallback: team_id or org_id as the acting entity (RFC 8693).
+ team_id = getattr(user_api_key_dict, "team_id", None)
+ org_id = getattr(user_api_key_dict, "org_id", None)
+ claims["act"] = {"sub": team_id or org_id or "litellm-proxy"}
+
+ # FR-10: Scope claim — tool-level least-privilege access.
+ raw_tool_name: str = data.get("mcp_tool_name", "")
+ tool_name = (
+ re.sub(r"[^a-zA-Z0-9_\-]", "_", raw_tool_name) if raw_tool_name else ""
+ )
+ claims["scope"] = self._build_scope(tool_name)
+
+ # FR-13: Claim operations — applied in Kong order: add → set → remove.
+ for k, v in self.add_claims.items():
+ if k not in claims:
+ claims[k] = v
+ claims.update(self.set_claims)
+ for k in self.remove_claims:
+ claims.pop(k, None)
+
+ return claims
+
+ # ------------------------------------------------------------------
+ # Guardrail hook
+ # ------------------------------------------------------------------
+
+ @log_guardrail_information
+ async def async_pre_call_hook(
+ self,
+ user_api_key_dict: UserAPIKeyAuth,
+ cache: DualCache,
+ data: dict,
+ call_type: CallTypesLiteral,
+ ) -> Optional[Union[Exception, str, dict]]:
+ """
+ Signs a JWT and injects it as the outbound Authorization header for
+ MCP tool calls. All other call types pass through unchanged.
+
+ Also handles:
+ - FR-5: Uses upstream jwt_claims for re-signing when available
+ - FR-9: Emits x-litellm-mcp-debug header
+ - FR-14: Reads and verifies channel token for two-token model
+ - FR-15: Validates required_claims against incoming JWT claims
+ """
+ if call_type != "call_mcp_tool":
+ return data
+
+ jwt_claims: Optional[Dict] = getattr(user_api_key_dict, "jwt_claims", None)
+
+ # FR-5: When access_token_discovery_uri is set, the guardrail operates in
+ # "verify + re-sign" mode. Verification of the incoming JWT is already
+ # performed by liteLLM's JWT auth handler (configured independently via
+ # general_settings.litellm_jwtauth). The decoded claims land in
+ # user_api_key_dict.jwt_claims and are used directly for re-signing.
+ if self.access_token_discovery_uri and not jwt_claims:
+ verbose_proxy_logger.debug(
+ "MCPJWTSigner: access_token_discovery_uri is configured but incoming "
+ "request has no JWT claims (virtual key auth). Proceeding with "
+ "M2M signing from user profile."
+ )
+
+ # FR-15: Validate required claims against the incoming token.
+ try:
+ _validate_required_claims(jwt_claims, self.required_claims)
+ except ValueError as exc:
+ raise exc # Propagate to block the MCP call
+
+ # FR-14: Resolve channel token for two-token model.
+ channel_token_claims: Optional[Dict] = None
+ raw_headers: Optional[Dict[str, str]] = data.get("mcp_raw_headers")
+ if raw_headers:
+ normalized_headers = {k.lower(): v for k, v in raw_headers.items()}
+ channel_token_raw = normalized_headers.get(self.channel_token_header)
+ if channel_token_raw:
+ try:
+ channel_token_claims = await self._verify_channel_token(
+ channel_token_raw
+ )
+ verbose_proxy_logger.debug(
+ "MCPJWTSigner: channel token resolved — act.sub=%s",
+ channel_token_claims.get("sub") or channel_token_claims.get("client_id"),
+ )
+ except Exception as exc:
+ verbose_proxy_logger.warning(
+ "MCPJWTSigner: channel token verification failed (%s). "
+ "Falling back to single-token mode.",
+ exc,
+ )
+
+ claims = self._build_claims(user_api_key_dict, data, channel_token_claims)
+
+ signed_token = jwt.encode(
+ claims,
+ self._private_key,
+ algorithm=self.ALGORITHM,
+ headers={"kid": self._kid},
+ )
+
+ # Merge into existing extra_headers rather than replacing — a prior guardrail
+ # in the chain may have already injected headers (e.g. tracing, correlation IDs).
+ # MCPJWTSigner sets Authorization last so its JWT takes precedence.
+ existing_headers: Dict[str, str] = data.get("extra_headers") or {}
+ outbound_headers: Dict[str, str] = {
+ **existing_headers,
+ "Authorization": f"Bearer {signed_token}",
+ }
+
+ # FR-9: Debug header — tells downstream what auth resolution was used.
+ if self.debug_header:
+ debug_info: Dict[str, Any] = {
+ "signer": "mcp_jwt_signer",
+ "kid": self._kid,
+ "issuer": self.issuer,
+ "sub": claims.get("sub"),
+ "act": claims.get("act"),
+ "mode": "re-sign" if jwt_claims else "sign",
+ "channel_token": channel_token_claims is not None,
+ }
+ outbound_headers["x-litellm-mcp-debug"] = json.dumps(
+ debug_info, separators=(",", ":")
+ )
+
+ data["extra_headers"] = outbound_headers
+
+ verbose_proxy_logger.debug(
+ "MCPJWTSigner: signed JWT sub=%s act=%s tool=%s exp=%d mode=%s",
+ claims.get("sub"),
+ claims.get("act", {}).get("sub"),
+ data.get("mcp_tool_name"),
+ claims["exp"],
+ "re-sign" if jwt_claims else "sign",
+ )
+
+ return data
diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py
index 01a0f55aac7..762eda818fb 100644
--- a/litellm/proxy/utils.py
+++ b/litellm/proxy/utils.py
@@ -454,8 +454,6 @@ class ProxyLogging:
for hook in PROXY_HOOKS:
proxy_hook = get_proxy_hook(hook)
- import inspect
-
expected_args = inspect.getfullargspec(proxy_hook).args
passed_in_args: Dict[str, Any] = {}
if "internal_usage_cache" in expected_args:
@@ -559,6 +557,11 @@ class ProxyLogging:
"user_api_key_request_route": kwargs.get("user_api_key_request_route"),
"mcp_tool_name": request_obj.tool_name, # Keep original for reference
"mcp_arguments": request_obj.arguments, # Keep original for reference
+ # Raw inbound headers from the client request. Passed through so that
+ # guardrail hooks (e.g. MCPJWTSigner) can read request headers such as
+ # X-Channel-Token for the two-token model (FR-14). May be None when
+ # called from contexts that don't have raw headers available.
+ "mcp_raw_headers": kwargs.get("raw_headers"),
}
return synthetic_data
@@ -824,17 +827,22 @@ class ProxyLogging:
) -> dict:
"""
Helper function to convert pre_call_hook response back to kwargs for MCP usage.
+
+ Supports:
+ - modified_arguments: Override tool call arguments
+ - extra_headers: Inject custom headers into the outbound MCP request
"""
if not response_data:
return original_kwargs
- # Apply any argument modifications from the hook response
modified_kwargs = original_kwargs.copy()
- # If the response contains modified arguments, apply them
if response_data.get("modified_arguments"):
modified_kwargs["arguments"] = response_data["modified_arguments"]
+ if response_data.get("extra_headers"):
+ modified_kwargs["extra_headers"] = response_data["extra_headers"]
+
return modified_kwargs
async def process_pre_call_hook_response(self, response, data, call_type):
diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py
index 27fa27e6da3..f798f05380d 100644
--- a/litellm/types/guardrails.py
+++ b/litellm/types/guardrails.py
@@ -79,6 +79,7 @@ class SupportedGuardrailIntegrations(Enum):
SEMANTIC_GUARD = "semantic_guard"
MCP_END_USER_PERMISSION = "mcp_end_user_permission"
BLOCK_CODE_EXECUTION = "block_code_execution"
+ MCP_JWT_SIGNER = "mcp_jwt_signer"
class Role(Enum):
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py
new file mode 100644
index 00000000000..32f3a340855
--- /dev/null
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py
@@ -0,0 +1,707 @@
+"""
+Tests for pre_mcp_call guardrail hook header mutation support.
+
+Validates that:
+1. _convert_mcp_hook_response_to_kwargs extracts extra_headers from hook response
+2. pre_call_tool_check returns hook-provided extra_headers AND modified arguments
+3. call_tool flows hook headers and modified arguments downstream
+4. Hook-provided headers take highest priority (merge after static_headers)
+5. OpenAPI-backed servers log a warning and continue (skip injection) when hook headers are present
+6. JWT claims are propagated in both standard and virtual-key fast paths
+7. Backward compatibility: hooks without extra_headers continue to work
+"""
+
+import asyncio
+from typing import Any, Dict, Optional
+from unittest.mock import AsyncMock, MagicMock, patch
+
+import pytest
+from fastapi import HTTPException
+
+from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
+from litellm.proxy._types import UserAPIKeyAuth
+from litellm.proxy.utils import ProxyLogging
+from litellm.types.mcp import MCPAuth, MCPTransport
+from litellm.types.mcp_server.mcp_server_manager import MCPServer
+
+
+class TestConvertMcpHookResponseToKwargs:
+ """Tests for ProxyLogging._convert_mcp_hook_response_to_kwargs"""
+
+ def setup_method(self):
+ self.proxy_logging = ProxyLogging(user_api_key_cache=MagicMock())
+
+ def test_returns_original_kwargs_when_response_is_none(self):
+ original = {"arguments": {"key": "val"}, "name": "tool"}
+ result = self.proxy_logging._convert_mcp_hook_response_to_kwargs(
+ None, original
+ )
+ assert result == original
+
+ def test_returns_original_kwargs_when_response_is_empty_dict(self):
+ original = {"arguments": {"key": "val"}}
+ result = self.proxy_logging._convert_mcp_hook_response_to_kwargs({}, original)
+ assert result == original
+
+ def test_extracts_modified_arguments(self):
+ original = {"arguments": {"old": "value"}}
+ response = {"modified_arguments": {"new": "value"}}
+ result = self.proxy_logging._convert_mcp_hook_response_to_kwargs(
+ response, original
+ )
+ assert result["arguments"] == {"new": "value"}
+
+ def test_extracts_extra_headers(self):
+ original = {"arguments": {"key": "val"}}
+ response = {"extra_headers": {"Authorization": "Bearer signed-jwt"}}
+ result = self.proxy_logging._convert_mcp_hook_response_to_kwargs(
+ response, original
+ )
+ assert result["extra_headers"] == {"Authorization": "Bearer signed-jwt"}
+
+ def test_extracts_both_arguments_and_headers(self):
+ original = {"arguments": {"old": "value"}}
+ response = {
+ "modified_arguments": {"new": "value"},
+ "extra_headers": {"X-Custom": "header-val"},
+ }
+ result = self.proxy_logging._convert_mcp_hook_response_to_kwargs(
+ response, original
+ )
+ assert result["arguments"] == {"new": "value"}
+ assert result["extra_headers"] == {"X-Custom": "header-val"}
+
+ def test_no_extra_headers_key_preserves_original(self):
+ """Backward compat: hooks that only return modified_arguments still work."""
+ original = {"arguments": {"key": "val"}}
+ response = {"modified_arguments": {"key": "new_val"}}
+ result = self.proxy_logging._convert_mcp_hook_response_to_kwargs(
+ response, original
+ )
+ assert "extra_headers" not in result
+ assert result["arguments"] == {"key": "new_val"}
+
+ def test_empty_extra_headers_not_set(self):
+ """Empty dict for extra_headers is falsy and should not be set."""
+ original = {"arguments": {"key": "val"}}
+ response = {"extra_headers": {}}
+ result = self.proxy_logging._convert_mcp_hook_response_to_kwargs(
+ response, original
+ )
+ assert "extra_headers" not in result
+
+
+class TestPreCallToolCheckReturnsHeaders:
+ """Tests that pre_call_tool_check returns hook-provided headers."""
+
+ def _make_server(self, name="test_server"):
+ return MCPServer(
+ server_id="test-id",
+ name=name,
+ server_name=name,
+ url="https://example.com",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.none,
+ )
+
+ @pytest.mark.asyncio
+ async def test_returns_empty_dict_when_hook_has_no_headers(self):
+ manager = MCPServerManager()
+ server = self._make_server()
+
+ proxy_logging = MagicMock(spec=ProxyLogging)
+ proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
+ return_value=MagicMock()
+ )
+ proxy_logging._convert_mcp_to_llm_format = MagicMock(
+ return_value={"model": "fake"}
+ )
+ proxy_logging.pre_call_hook = AsyncMock(
+ return_value={"modified_arguments": {"key": "val"}}
+ )
+ proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(
+ return_value={"arguments": {"key": "val"}}
+ )
+
+ with patch.object(manager, "check_allowed_or_banned_tools", return_value=True):
+ with patch.object(
+ manager,
+ "check_tool_permission_for_key_team",
+ new_callable=AsyncMock,
+ ):
+ with patch.object(manager, "validate_allowed_params"):
+ result = await manager.pre_call_tool_check(
+ name="test_tool",
+ arguments={"key": "val"},
+ server_name="test_server",
+ user_api_key_auth=None,
+ proxy_logging_obj=proxy_logging,
+ server=server,
+ )
+
+ assert result == {}
+
+ @pytest.mark.asyncio
+ async def test_returns_extra_headers_from_hook(self):
+ manager = MCPServerManager()
+ server = self._make_server()
+
+ hook_headers = {"Authorization": "Bearer signed-jwt", "X-Trace-Id": "abc123"}
+
+ proxy_logging = MagicMock(spec=ProxyLogging)
+ proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
+ return_value=MagicMock()
+ )
+ proxy_logging._convert_mcp_to_llm_format = MagicMock(
+ return_value={"model": "fake"}
+ )
+ proxy_logging.pre_call_hook = AsyncMock(
+ return_value={"extra_headers": hook_headers}
+ )
+ proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(
+ return_value={"arguments": {"key": "val"}, "extra_headers": hook_headers}
+ )
+
+ with patch.object(manager, "check_allowed_or_banned_tools", return_value=True):
+ with patch.object(
+ manager,
+ "check_tool_permission_for_key_team",
+ new_callable=AsyncMock,
+ ):
+ with patch.object(manager, "validate_allowed_params"):
+ result = await manager.pre_call_tool_check(
+ name="test_tool",
+ arguments={"key": "val"},
+ server_name="test_server",
+ user_api_key_auth=None,
+ proxy_logging_obj=proxy_logging,
+ server=server,
+ )
+
+ assert result["extra_headers"] == hook_headers
+
+ @pytest.mark.asyncio
+ async def test_returns_empty_dict_when_hook_returns_none(self):
+ manager = MCPServerManager()
+ server = self._make_server()
+
+ proxy_logging = MagicMock(spec=ProxyLogging)
+ proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
+ return_value=MagicMock()
+ )
+ proxy_logging._convert_mcp_to_llm_format = MagicMock(
+ return_value={"model": "fake"}
+ )
+ proxy_logging.pre_call_hook = AsyncMock(return_value=None)
+
+ with patch.object(manager, "check_allowed_or_banned_tools", return_value=True):
+ with patch.object(
+ manager,
+ "check_tool_permission_for_key_team",
+ new_callable=AsyncMock,
+ ):
+ with patch.object(manager, "validate_allowed_params"):
+ result = await manager.pre_call_tool_check(
+ name="test_tool",
+ arguments={"key": "val"},
+ server_name="test_server",
+ user_api_key_auth=None,
+ proxy_logging_obj=proxy_logging,
+ server=server,
+ )
+
+ assert result == {}
+
+ @pytest.mark.asyncio
+ async def test_returns_modified_arguments_from_hook(self):
+ """Modified arguments from the hook must be returned so the caller can use them."""
+ manager = MCPServerManager()
+ server = self._make_server()
+
+ original_args = {"key": "original"}
+ modified_args = {"key": "modified", "extra": "added"}
+
+ proxy_logging = MagicMock(spec=ProxyLogging)
+ proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
+ return_value=MagicMock()
+ )
+ proxy_logging._convert_mcp_to_llm_format = MagicMock(
+ return_value={"model": "fake"}
+ )
+ proxy_logging.pre_call_hook = AsyncMock(
+ return_value={"modified_arguments": modified_args}
+ )
+ proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(
+ return_value={"arguments": modified_args}
+ )
+
+ with patch.object(manager, "check_allowed_or_banned_tools", return_value=True):
+ with patch.object(
+ manager,
+ "check_tool_permission_for_key_team",
+ new_callable=AsyncMock,
+ ):
+ with patch.object(manager, "validate_allowed_params"):
+ result = await manager.pre_call_tool_check(
+ name="test_tool",
+ arguments=original_args,
+ server_name="test_server",
+ user_api_key_auth=None,
+ proxy_logging_obj=proxy_logging,
+ server=server,
+ )
+
+ assert result["arguments"] == modified_args
+
+ @pytest.mark.asyncio
+ async def test_returns_both_modified_arguments_and_headers(self):
+ """Hook can modify both arguments and inject headers simultaneously."""
+ manager = MCPServerManager()
+ server = self._make_server()
+
+ modified_args = {"key": "modified"}
+ hook_headers = {"Authorization": "Bearer jwt"}
+
+ proxy_logging = MagicMock(spec=ProxyLogging)
+ proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
+ return_value=MagicMock()
+ )
+ proxy_logging._convert_mcp_to_llm_format = MagicMock(
+ return_value={"model": "fake"}
+ )
+ proxy_logging.pre_call_hook = AsyncMock(return_value={"dummy": True})
+ proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(
+ return_value={"arguments": modified_args, "extra_headers": hook_headers}
+ )
+
+ with patch.object(manager, "check_allowed_or_banned_tools", return_value=True):
+ with patch.object(
+ manager,
+ "check_tool_permission_for_key_team",
+ new_callable=AsyncMock,
+ ):
+ with patch.object(manager, "validate_allowed_params"):
+ result = await manager.pre_call_tool_check(
+ name="test_tool",
+ arguments={"key": "original"},
+ server_name="test_server",
+ user_api_key_auth=None,
+ proxy_logging_obj=proxy_logging,
+ server=server,
+ )
+
+ assert result["arguments"] == modified_args
+ assert result["extra_headers"] == hook_headers
+
+
+class TestCallToolFlowsHookHeaders:
+ """Tests that call_tool passes hook_extra_headers to _call_regular_mcp_tool."""
+
+ def _make_server(self, name="test_server"):
+ return MCPServer(
+ server_id="test-id",
+ name=name,
+ server_name=name,
+ url="https://example.com",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.none,
+ )
+
+ @pytest.mark.asyncio
+ async def test_hook_headers_passed_to_call_regular_mcp_tool(self):
+ """Verify that hook_extra_headers kwarg is forwarded."""
+ manager = MCPServerManager()
+ server = self._make_server()
+
+ hook_headers = {"Authorization": "Bearer signed-jwt"}
+
+ with patch.object(
+ manager,
+ "_get_mcp_server_from_tool_name",
+ return_value=server,
+ ):
+ with patch.object(
+ manager,
+ "pre_call_tool_check",
+ new_callable=AsyncMock,
+ return_value={"extra_headers": hook_headers},
+ ):
+ with patch.object(
+ manager,
+ "_create_during_hook_task",
+ return_value=asyncio.create_task(asyncio.sleep(0)),
+ ):
+ with patch.object(
+ manager,
+ "_call_regular_mcp_tool",
+ new_callable=AsyncMock,
+ return_value=MagicMock(),
+ ) as mock_call:
+ proxy_logging = MagicMock(spec=ProxyLogging)
+
+ await manager.call_tool(
+ server_name="test_server",
+ name="test_tool",
+ arguments={"key": "val"},
+ proxy_logging_obj=proxy_logging,
+ )
+
+ mock_call.assert_called_once()
+ call_kwargs = mock_call.call_args
+ assert call_kwargs.kwargs.get("hook_extra_headers") == hook_headers
+
+ @pytest.mark.asyncio
+ async def test_no_hook_headers_when_no_proxy_logging(self):
+ """Without proxy_logging_obj, no pre_call_tool_check runs."""
+ manager = MCPServerManager()
+ server = self._make_server()
+
+ with patch.object(
+ manager,
+ "_get_mcp_server_from_tool_name",
+ return_value=server,
+ ):
+ with patch.object(
+ manager,
+ "_call_regular_mcp_tool",
+ new_callable=AsyncMock,
+ return_value=MagicMock(),
+ ) as mock_call:
+ await manager.call_tool(
+ server_name="test_server",
+ name="test_tool",
+ arguments={"key": "val"},
+ proxy_logging_obj=None,
+ )
+
+ mock_call.assert_called_once()
+ call_kwargs = mock_call.call_args
+ assert call_kwargs.kwargs.get("hook_extra_headers") is None
+
+ @pytest.mark.asyncio
+ async def test_modified_arguments_passed_to_downstream(self):
+ """Hook-modified arguments must be used for the actual tool call."""
+ manager = MCPServerManager()
+ server = self._make_server()
+
+ modified_args = {"key": "modified_by_hook"}
+
+ with patch.object(
+ manager,
+ "_get_mcp_server_from_tool_name",
+ return_value=server,
+ ):
+ with patch.object(
+ manager,
+ "pre_call_tool_check",
+ new_callable=AsyncMock,
+ return_value={"arguments": modified_args},
+ ):
+ with patch.object(
+ manager,
+ "_create_during_hook_task",
+ return_value=asyncio.create_task(asyncio.sleep(0)),
+ ):
+ with patch.object(
+ manager,
+ "_call_regular_mcp_tool",
+ new_callable=AsyncMock,
+ return_value=MagicMock(),
+ ) as mock_call:
+ proxy_logging = MagicMock(spec=ProxyLogging)
+
+ await manager.call_tool(
+ server_name="test_server",
+ name="test_tool",
+ arguments={"key": "original"},
+ proxy_logging_obj=proxy_logging,
+ )
+
+ mock_call.assert_called_once()
+ call_kwargs = mock_call.call_args
+ assert call_kwargs.kwargs.get("arguments") == modified_args
+
+ @pytest.mark.asyncio
+ async def test_openapi_server_warns_and_continues_on_hook_headers(self):
+ """OpenAPI-backed servers log a warning and continue when hook injects headers."""
+ manager = MCPServerManager()
+ server = MCPServer(
+ server_id="test-id",
+ name="openapi_server",
+ server_name="openapi_server",
+ url="https://example.com",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.none,
+ spec_path="/path/to/spec.yaml",
+ )
+
+ with patch.object(
+ manager, "_get_mcp_server_from_tool_name", return_value=server
+ ):
+ with patch.object(
+ manager,
+ "pre_call_tool_check",
+ new_callable=AsyncMock,
+ return_value={"extra_headers": {"Authorization": "Bearer jwt"}},
+ ):
+ with patch.object(
+ manager,
+ "_create_during_hook_task",
+ return_value=asyncio.create_task(asyncio.sleep(0)),
+ ):
+ with patch.object(
+ manager,
+ "_call_openapi_tool_handler",
+ new_callable=AsyncMock,
+ return_value=MagicMock(),
+ ):
+ import litellm.proxy._experimental.mcp_server.mcp_server_manager as mgr_mod
+
+ proxy_logging = MagicMock(spec=ProxyLogging)
+
+ with patch.object(mgr_mod, "verbose_logger") as mock_logger:
+ # Should NOT raise — just warn and proceed
+ await manager.call_tool(
+ server_name="openapi_server",
+ name="test_tool",
+ arguments={},
+ proxy_logging_obj=proxy_logging,
+ )
+ mock_logger.warning.assert_called_once()
+ assert "header injection is not supported" in mock_logger.warning.call_args[0][0]
+
+ @pytest.mark.asyncio
+ async def test_openapi_server_no_error_without_hook_headers(self):
+ """No exception when OpenAPI server has no hook-injected headers."""
+ manager = MCPServerManager()
+ server = MCPServer(
+ server_id="test-id",
+ name="openapi_server",
+ server_name="openapi_server",
+ url="https://example.com",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.none,
+ spec_path="/path/to/spec.yaml",
+ )
+
+ with patch.object(
+ manager, "_get_mcp_server_from_tool_name", return_value=server
+ ):
+ with patch.object(
+ manager,
+ "pre_call_tool_check",
+ new_callable=AsyncMock,
+ return_value={},
+ ):
+ with patch.object(
+ manager,
+ "_create_during_hook_task",
+ return_value=asyncio.create_task(asyncio.sleep(0)),
+ ):
+ with patch.object(
+ manager,
+ "_call_openapi_tool_handler",
+ new_callable=AsyncMock,
+ return_value=MagicMock(),
+ ):
+ proxy_logging = MagicMock(spec=ProxyLogging)
+
+ await manager.call_tool(
+ server_name="openapi_server",
+ name="test_tool",
+ arguments={},
+ proxy_logging_obj=proxy_logging,
+ )
+
+
+class TestHookHeaderMergePriority:
+ """Tests that hook-provided headers have highest priority in _call_regular_mcp_tool."""
+
+ def _make_server(
+ self,
+ static_headers: Optional[Dict[str, str]] = None,
+ extra_headers_config: Optional[list] = None,
+ ):
+ return MCPServer(
+ server_id="test-id",
+ name="Test Server",
+ server_name="test_server",
+ url="https://example.com",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.none,
+ static_headers=static_headers,
+ extra_headers=extra_headers_config,
+ )
+
+ @pytest.mark.asyncio
+ async def test_hook_headers_override_static_headers(self):
+ """Hook headers should take precedence over static_headers."""
+ manager = MCPServerManager()
+ server = self._make_server(
+ static_headers={"Authorization": "Bearer static-token", "X-Static": "yes"}
+ )
+
+ hook_headers = {"Authorization": "Bearer hook-signed-jwt"}
+
+ captured_extra_headers: Dict[str, Any] = {}
+
+ async def fake_create_mcp_client(
+ server, mcp_auth_header=None, extra_headers=None, stdio_env=None
+ ):
+ captured_extra_headers["value"] = extra_headers
+ mock_client = MagicMock()
+ mock_client.call_tool = AsyncMock(return_value=MagicMock())
+ return mock_client
+
+ with patch.object(
+ manager, "_create_mcp_client", side_effect=fake_create_mcp_client
+ ):
+ with patch.object(manager, "_build_stdio_env", return_value=None):
+ try:
+ await manager._call_regular_mcp_tool(
+ mcp_server=server,
+ original_tool_name="test_tool",
+ arguments={"key": "val"},
+ tasks=[],
+ mcp_auth_header=None,
+ mcp_server_auth_headers=None,
+ oauth2_headers=None,
+ raw_headers=None,
+ proxy_logging_obj=None,
+ hook_extra_headers=hook_headers,
+ )
+ except Exception:
+ pass
+
+ headers = captured_extra_headers.get("value", {})
+ assert headers["Authorization"] == "Bearer hook-signed-jwt"
+ assert headers["X-Static"] == "yes"
+
+ @pytest.mark.asyncio
+ async def test_no_hook_headers_preserves_existing_behavior(self):
+ """When hook_extra_headers is None, existing header logic is unchanged."""
+ manager = MCPServerManager()
+ server = self._make_server(
+ static_headers={"X-Static": "static-value"}
+ )
+
+ captured_extra_headers: Dict[str, Any] = {}
+
+ async def fake_create_mcp_client(
+ server, mcp_auth_header=None, extra_headers=None, stdio_env=None
+ ):
+ captured_extra_headers["value"] = extra_headers
+ mock_client = MagicMock()
+ mock_client.call_tool = AsyncMock(return_value=MagicMock())
+ return mock_client
+
+ with patch.object(
+ manager, "_create_mcp_client", side_effect=fake_create_mcp_client
+ ):
+ with patch.object(manager, "_build_stdio_env", return_value=None):
+ try:
+ await manager._call_regular_mcp_tool(
+ mcp_server=server,
+ original_tool_name="test_tool",
+ arguments={"key": "val"},
+ tasks=[],
+ mcp_auth_header=None,
+ mcp_server_auth_headers=None,
+ oauth2_headers=None,
+ raw_headers=None,
+ proxy_logging_obj=None,
+ hook_extra_headers=None,
+ )
+ except Exception:
+ pass
+
+ headers = captured_extra_headers.get("value", {})
+ assert headers == {"X-Static": "static-value"}
+
+ @pytest.mark.asyncio
+ async def test_hook_headers_merge_with_oauth2(self):
+ """Hook headers merge on top of OAuth2 headers."""
+ manager = MCPServerManager()
+ server = MCPServer(
+ server_id="test-id",
+ name="Test Server",
+ server_name="test_server",
+ url="https://example.com",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ )
+
+ captured_extra_headers: Dict[str, Any] = {}
+
+ async def fake_create_mcp_client(
+ server, mcp_auth_header=None, extra_headers=None, stdio_env=None
+ ):
+ captured_extra_headers["value"] = extra_headers
+ mock_client = MagicMock()
+ mock_client.call_tool = AsyncMock(return_value=MagicMock())
+ return mock_client
+
+ with patch.object(
+ manager, "_create_mcp_client", side_effect=fake_create_mcp_client
+ ):
+ with patch.object(manager, "_build_stdio_env", return_value=None):
+ try:
+ await manager._call_regular_mcp_tool(
+ mcp_server=server,
+ original_tool_name="test_tool",
+ arguments={"key": "val"},
+ tasks=[],
+ mcp_auth_header=None,
+ mcp_server_auth_headers=None,
+ oauth2_headers={
+ "Authorization": "Bearer oauth2-token",
+ "X-OAuth": "yes",
+ },
+ raw_headers=None,
+ proxy_logging_obj=None,
+ hook_extra_headers={
+ "Authorization": "Bearer hook-jwt",
+ "X-Trace-Id": "trace-123",
+ },
+ )
+ except Exception:
+ pass
+
+ headers = captured_extra_headers.get("value", {})
+ assert headers["Authorization"] == "Bearer hook-jwt"
+ assert headers["X-OAuth"] == "yes"
+ assert headers["X-Trace-Id"] == "trace-123"
+
+
+class TestUserAPIKeyAuthJwtClaims:
+ """Tests that UserAPIKeyAuth correctly carries jwt_claims."""
+
+ def test_jwt_claims_field_defaults_to_none(self):
+ auth = UserAPIKeyAuth(api_key="test-key")
+ assert auth.jwt_claims is None
+
+ def test_jwt_claims_field_accepts_dict(self):
+ claims = {"sub": "user-123", "iss": "litellm", "exp": 9999999999}
+ auth = UserAPIKeyAuth(api_key="test-key", jwt_claims=claims)
+ assert auth.jwt_claims == claims
+ assert auth.jwt_claims["sub"] == "user-123"
+
+ def test_jwt_claims_backward_compatible_without_field(self):
+ """Existing code that doesn't pass jwt_claims should still work."""
+ auth = UserAPIKeyAuth(
+ api_key="test-key",
+ user_id="user-1",
+ team_id="team-1",
+ )
+ assert auth.jwt_claims is None
+ assert auth.user_id == "user-1"
+
+ def test_jwt_claims_set_after_construction(self):
+ """Virtual-key fast path sets jwt_claims after the object is created."""
+ auth = UserAPIKeyAuth(api_key="test-key")
+ assert auth.jwt_claims is None
+
+ claims = {"sub": "user-456", "iss": "okta", "groups": ["admin"]}
+ auth.jwt_claims = claims
+ assert auth.jwt_claims == claims
+ assert auth.jwt_claims["groups"] == ["admin"]
diff --git a/tests/test_litellm/proxy/guardrails/test_mcp_jwt_signer.py b/tests/test_litellm/proxy/guardrails/test_mcp_jwt_signer.py
new file mode 100644
index 00000000000..6f17f9d0f23
--- /dev/null
+++ b/tests/test_litellm/proxy/guardrails/test_mcp_jwt_signer.py
@@ -0,0 +1,810 @@
+"""
+Tests for the MCPJWTSigner built-in guardrail.
+
+Tests cover:
+ - RSA key generation and loading
+ - JWT signing and JWKS format
+ - Claim building (sub, act, scope)
+ - Hook fires for call_mcp_tool, skips other call types
+ - get_mcp_jwt_signer() singleton pattern
+"""
+
+import base64
+import json
+import time
+from typing import Any, Dict, Optional
+from unittest.mock import MagicMock, patch
+
+import jwt
+import pytest
+from cryptography.hazmat.primitives import serialization
+from cryptography.hazmat.primitives.asymmetric import rsa
+
+# ---------------------------------------------------------------------------
+# Helpers
+# ---------------------------------------------------------------------------
+
+
+def _make_user_api_key_dict(
+ user_id: str = "user-123",
+ team_id: str = "team-abc",
+ user_email: str = "user@example.com",
+ end_user_id: Optional[str] = None,
+) -> MagicMock:
+ mock = MagicMock()
+ mock.user_id = user_id
+ mock.team_id = team_id
+ mock.user_email = user_email
+ mock.end_user_id = end_user_id
+ mock.org_id = None
+ # jwt_claims must default to None so the mock doesn't pretend to have
+ # upstream JWT claims when none were configured.
+ mock.jwt_claims = None
+ return mock
+
+
+def _decode_unverified(token: str) -> Dict[str, Any]:
+ return jwt.decode(token, options={"verify_signature": False})
+
+
+# ---------------------------------------------------------------------------
+# Import target (inline so we can reset the singleton between tests)
+# ---------------------------------------------------------------------------
+
+
+def _make_signer(**kwargs: Any):
+ # Reset singleton before each signer creation to avoid cross-test pollution
+ import litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer as mod
+
+ mod._mcp_jwt_signer_instance = None
+
+ from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
+ MCPJWTSigner,
+ )
+
+ return MCPJWTSigner(
+ guardrail_name="test-jwt-signer",
+ event_hook="pre_mcp_call",
+ default_on=True,
+ **kwargs,
+ )
+
+
+# ---------------------------------------------------------------------------
+# Key generation tests
+# ---------------------------------------------------------------------------
+
+
+def test_auto_generates_rsa_keypair():
+ """MCPJWTSigner auto-generates an RSA-2048 keypair when env var is unset."""
+ signer = _make_signer()
+ assert signer._private_key is not None
+ assert signer._public_key is not None
+ assert signer._kid is not None and len(signer._kid) == 16
+
+
+def test_kid_is_deterministic():
+ """Two signers built from the same key have the same kid."""
+ signer1 = _make_signer()
+ private_pem = signer1._private_key.private_bytes(
+ encoding=serialization.Encoding.PEM,
+ format=serialization.PrivateFormat.TraditionalOpenSSL,
+ encryption_algorithm=serialization.NoEncryption(),
+ ).decode("utf-8")
+
+ with patch.dict("os.environ", {"MCP_JWT_SIGNING_KEY": private_pem}):
+ signer2 = _make_signer()
+
+ assert signer1._kid == signer2._kid
+
+
+def test_load_key_from_env_var():
+ """MCPJWTSigner loads a user-provided RSA key from the env var."""
+ private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
+ pem = private_key.private_bytes(
+ encoding=serialization.Encoding.PEM,
+ format=serialization.PrivateFormat.TraditionalOpenSSL,
+ encryption_algorithm=serialization.NoEncryption(),
+ ).decode("utf-8")
+
+ with patch.dict("os.environ", {"MCP_JWT_SIGNING_KEY": pem}):
+ signer = _make_signer()
+
+ assert signer._kid is not None
+
+
+# ---------------------------------------------------------------------------
+# JWKS tests
+# ---------------------------------------------------------------------------
+
+
+def test_get_jwks_format():
+ """get_jwks() returns a valid JWKS dict with RSA fields."""
+ signer = _make_signer()
+ jwks = signer.get_jwks()
+
+ assert "keys" in jwks
+ assert len(jwks["keys"]) == 1
+ key = jwks["keys"][0]
+
+ assert key["kty"] == "RSA"
+ assert key["alg"] == "RS256"
+ assert key["use"] == "sig"
+ assert key["kid"] == signer._kid
+ assert "n" in key and len(key["n"]) > 0
+ assert "e" in key and key["e"] == "AQAB" # 65537 in base64url
+
+
+def test_jwks_public_key_can_verify_signed_jwt():
+ """A JWT signed by MCPJWTSigner can be verified using the JWKS public key."""
+ signer = _make_signer(issuer="https://litellm.example.com", audience="mcp")
+ now = int(time.time())
+ claims = {"iss": "https://litellm.example.com", "aud": "mcp", "iat": now, "exp": now + 300}
+
+ token = jwt.encode(claims, signer._private_key, algorithm="RS256")
+
+ # Reconstruct public key from JWKS
+ jwks = signer.get_jwks()
+ key_data = jwks["keys"][0]
+ n = int.from_bytes(base64.urlsafe_b64decode(key_data["n"] + "=="), byteorder="big")
+ e = int.from_bytes(base64.urlsafe_b64decode(key_data["e"] + "=="), byteorder="big")
+ from cryptography.hazmat.primitives.asymmetric.rsa import RSAPublicNumbers
+ pub_key = RSAPublicNumbers(e=e, n=n).public_key()
+
+ decoded = jwt.decode(
+ token,
+ pub_key,
+ algorithms=["RS256"],
+ audience="mcp",
+ issuer="https://litellm.example.com",
+ )
+ assert decoded["iss"] == "https://litellm.example.com"
+
+
+# ---------------------------------------------------------------------------
+# Claim building tests
+# ---------------------------------------------------------------------------
+
+
+def test_build_claims_standard_fields():
+ """_build_claims() populates iss, aud, iat, exp, nbf."""
+ signer = _make_signer(issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300)
+ user_dict = _make_user_api_key_dict()
+ data = {"mcp_tool_name": "get_weather"}
+
+ claims = signer._build_claims(user_dict, data)
+
+ assert claims["iss"] == "https://litellm.example.com"
+ assert claims["aud"] == "mcp"
+ assert "iat" in claims
+ assert "exp" in claims
+ assert claims["exp"] - claims["iat"] == 300
+ assert "nbf" in claims
+
+
+def test_build_claims_identity():
+ """_build_claims() sets sub from user_id and act from team_id (RFC 8693)."""
+ signer = _make_signer()
+ user_dict = _make_user_api_key_dict(user_id="user-xyz", team_id="team-eng")
+ data: Dict[str, Any] = {}
+
+ claims = signer._build_claims(user_dict, data)
+
+ assert claims["sub"] == "user-xyz"
+ assert claims["act"]["sub"] == "team-eng"
+ assert claims["email"] == "user@example.com"
+
+
+def test_build_claims_scope_with_tool():
+ """_build_claims() encodes tool-specific scope when mcp_tool_name is set."""
+ signer = _make_signer()
+ user_dict = _make_user_api_key_dict()
+ data = {"mcp_tool_name": "search_web"}
+
+ claims = signer._build_claims(user_dict, data)
+
+ scopes = set(claims["scope"].split())
+ assert "mcp:tools/call" in scopes
+ assert "mcp:tools/search_web:call" in scopes
+ # Tool-call JWTs must NOT carry mcp:tools/list — least-privilege
+ assert "mcp:tools/list" not in scopes
+
+
+def test_build_claims_scope_without_tool():
+ """_build_claims() includes mcp:tools/list when no specific tool is called."""
+ signer = _make_signer()
+ user_dict = _make_user_api_key_dict()
+ data: Dict[str, Any] = {}
+
+ claims = signer._build_claims(user_dict, data)
+
+ scopes = set(claims["scope"].split())
+ assert "mcp:tools/call" in scopes
+ assert "mcp:tools/list" in scopes
+ # No per-tool call scope when no tool name was given
+ assert not any(s.endswith(":call") and s != "mcp:tools/call" for s in scopes)
+
+
+def test_build_claims_act_fallback_to_litellm_proxy():
+ """_build_claims() falls back to 'litellm-proxy' when team_id and org_id are absent."""
+ signer = _make_signer()
+ user_dict = _make_user_api_key_dict()
+ user_dict.team_id = None
+ user_dict.org_id = None
+
+ claims = signer._build_claims(user_dict, {})
+
+ assert claims["act"]["sub"] == "litellm-proxy"
+
+
+def test_build_claims_sub_fallback_to_token_hash():
+ """_build_claims() sets sub to an apikey: hash when user_id is absent."""
+ signer = _make_signer()
+ user_dict = _make_user_api_key_dict(user_id="")
+ user_dict.user_id = None
+ user_dict.token = "sk-test-api-key-abc123"
+
+ claims = signer._build_claims(user_dict, {})
+
+ assert claims["sub"].startswith("apikey:")
+ assert len(claims["sub"]) == len("apikey:") + 16 # sha256 hex[:16]
+
+
+def test_build_claims_sub_fallback_to_litellm_proxy_when_no_token():
+ """_build_claims() falls back to 'litellm-proxy' when user_id and token are both absent."""
+ signer = _make_signer()
+ user_dict = _make_user_api_key_dict(user_id="")
+ user_dict.user_id = None
+ user_dict.token = None
+ user_dict.api_key = None
+
+ claims = signer._build_claims(user_dict, {})
+
+ assert claims["sub"] == "litellm-proxy"
+
+
+def test_init_raises_on_zero_ttl():
+ """MCPJWTSigner raises ValueError when ttl_seconds is 0."""
+ with pytest.raises(ValueError, match="ttl_seconds must be > 0"):
+ _make_signer(ttl_seconds=0)
+
+
+def test_init_raises_on_negative_ttl():
+ """MCPJWTSigner raises ValueError when ttl_seconds is negative."""
+ with pytest.raises(ValueError, match="ttl_seconds must be > 0"):
+ _make_signer(ttl_seconds=-60)
+
+
+def test_jwks_max_age_persistent_key():
+ """jwks_max_age is 3600 when key loaded from env var."""
+ from cryptography.hazmat.primitives import serialization
+ from cryptography.hazmat.primitives.asymmetric import rsa as crsa
+
+ private_key = crsa.generate_private_key(public_exponent=65537, key_size=2048)
+ pem = private_key.private_bytes(
+ encoding=serialization.Encoding.PEM,
+ format=serialization.PrivateFormat.TraditionalOpenSSL,
+ encryption_algorithm=serialization.NoEncryption(),
+ ).decode("utf-8")
+
+ with patch.dict("os.environ", {"MCP_JWT_SIGNING_KEY": pem}):
+ signer = _make_signer()
+
+ assert signer.jwks_max_age == 3600
+
+
+def test_jwks_max_age_auto_generated_key():
+ """jwks_max_age is 300 for auto-generated (ephemeral) keys."""
+ signer = _make_signer()
+ assert signer.jwks_max_age == 300
+
+
+# ---------------------------------------------------------------------------
+# Hook dispatch tests
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.asyncio
+async def test_hook_fires_for_call_mcp_tool():
+ """async_pre_call_hook() injects Authorization header for call_mcp_tool."""
+ signer = _make_signer(issuer="https://litellm.example.com", audience="mcp")
+ user_dict = _make_user_api_key_dict()
+ data = {"mcp_tool_name": "do_thing"}
+
+ result = await signer.async_pre_call_hook(
+ user_api_key_dict=user_dict,
+ cache=MagicMock(),
+ data=data,
+ call_type="call_mcp_tool",
+ )
+
+ assert isinstance(result, dict)
+ assert "extra_headers" in result
+ assert result["extra_headers"]["Authorization"].startswith("Bearer ")
+
+
+@pytest.mark.asyncio
+async def test_hook_skips_non_mcp_call_types():
+ """async_pre_call_hook() leaves data unchanged for non-MCP call types."""
+ signer = _make_signer()
+ user_dict = _make_user_api_key_dict()
+ data = {"messages": [{"role": "user", "content": "hello"}]}
+
+ for call_type in ("completion", "acompletion", "embedding", "list_mcp_tools"):
+ original_data = {**data}
+ result = await signer.async_pre_call_hook(
+ user_api_key_dict=user_dict,
+ cache=MagicMock(),
+ data=original_data,
+ call_type=call_type, # type: ignore[arg-type]
+ )
+ assert "extra_headers" not in (result or {}), f"extra_headers should not be set for {call_type}"
+
+
+@pytest.mark.asyncio
+async def test_signed_token_is_verifiable():
+ """The JWT injected by the hook can be verified against the JWKS public key."""
+ signer = _make_signer(issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300)
+ user_dict = _make_user_api_key_dict(user_id="alice", team_id="backend")
+ data = {"mcp_tool_name": "search"}
+
+ result = await signer.async_pre_call_hook(
+ user_api_key_dict=user_dict,
+ cache=MagicMock(),
+ data=data,
+ call_type="call_mcp_tool",
+ )
+
+ assert isinstance(result, dict)
+ token = result["extra_headers"]["Authorization"].removeprefix("Bearer ")
+
+ decoded = _decode_unverified(token)
+ assert decoded["sub"] == "alice"
+ assert decoded["act"]["sub"] == "backend"
+ assert "mcp:tools/search:call" in decoded["scope"]
+ assert decoded["iss"] == "https://litellm.example.com"
+ assert decoded["aud"] == "mcp"
+
+
+# ---------------------------------------------------------------------------
+# Singleton tests
+# ---------------------------------------------------------------------------
+
+
+def test_get_mcp_jwt_signer_returns_none_before_init():
+ """get_mcp_jwt_signer() returns None before any MCPJWTSigner is created."""
+ import litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer as mod
+
+ mod._mcp_jwt_signer_instance = None
+
+ from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
+ get_mcp_jwt_signer,
+ )
+
+ assert get_mcp_jwt_signer() is None
+
+
+def test_get_mcp_jwt_signer_returns_instance_after_init():
+ """get_mcp_jwt_signer() returns the initialized signer instance."""
+ from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
+ get_mcp_jwt_signer,
+ )
+
+ signer = _make_signer()
+ assert get_mcp_jwt_signer() is signer
+
+
+# ---------------------------------------------------------------------------
+# FR-5: Verify + re-sign — uses upstream jwt_claims when available
+# ---------------------------------------------------------------------------
+
+
+def test_build_claims_uses_jwt_claims_sub_when_available():
+ """FR-5: When jwt_claims is populated, sub is taken from the upstream token."""
+ signer = _make_signer(
+ access_token_discovery_uri="https://okta.example.com/.well-known/openid-configuration"
+ )
+ user_dict = _make_user_api_key_dict(user_id="litellm-user-999")
+ # Simulate upstream Okta JWT claims already decoded by litellm JWT auth.
+ user_dict.jwt_claims = {"sub": "okta-user-abc123", "email": "alice@corp.com"}
+
+ claims = signer._build_claims(user_dict, {})
+
+ # sub must come from upstream jwt_claims, not litellm user_id
+ assert claims["sub"] == "okta-user-abc123"
+ assert claims["email"] == "alice@corp.com"
+
+
+def test_build_claims_falls_back_to_user_id_when_no_jwt_claims():
+ """FR-5 M2M mode: falls back to user_id when jwt_claims is absent."""
+ signer = _make_signer(
+ access_token_discovery_uri="https://okta.example.com/.well-known/openid-configuration"
+ )
+ user_dict = _make_user_api_key_dict(user_id="svc-account-42")
+ user_dict.jwt_claims = None
+
+ claims = signer._build_claims(user_dict, {})
+
+ assert claims["sub"] == "svc-account-42"
+
+
+# ---------------------------------------------------------------------------
+# FR-12: End-user identity mapping via end_user_claim_sources
+# ---------------------------------------------------------------------------
+
+
+def test_end_user_claim_sources_picks_first_non_empty():
+ """FR-12: Identity is resolved from the first non-empty source."""
+ signer = _make_signer(
+ end_user_claim_sources=["sso_id", "preferred_username", "sub"]
+ )
+ user_dict = _make_user_api_key_dict(user_id="ignored")
+ user_dict.jwt_claims = {
+ "sub": "fallback-sub",
+ "preferred_username": "alice",
+ # sso_id absent
+ }
+
+ claims = signer._build_claims(user_dict, {})
+
+ # sso_id absent, preferred_username present → use preferred_username
+ assert claims["sub"] == "alice"
+
+
+def test_end_user_claim_sources_header_resolution():
+ """FR-12: Identity can be resolved from a raw header when claim is absent."""
+ signer = _make_signer(end_user_claim_sources=["x-end-user-id", "sub"])
+ user_dict = _make_user_api_key_dict(user_id="litellm-user")
+ user_dict.jwt_claims = {"sub": "fallback-sub"}
+ data = {
+ "mcp_raw_headers": {"x-end-user-id": "header-user-789"},
+ }
+
+ claims = signer._build_claims(user_dict, data)
+
+ assert claims["sub"] == "header-user-789"
+
+
+def test_end_user_claim_sources_falls_through_all():
+ """FR-12: Falls back gracefully when no source matches."""
+ signer = _make_signer(end_user_claim_sources=["nonexistent_claim"])
+ user_dict = _make_user_api_key_dict(user_id="")
+ user_dict.user_id = None
+ user_dict.jwt_claims = {}
+ user_dict.token = "sk-abc123"
+
+ claims = signer._build_claims(user_dict, {})
+
+ # Falls back to apikey hash
+ assert claims["sub"].startswith("apikey:")
+
+
+# ---------------------------------------------------------------------------
+# FR-13: Claim operations (add_claims, set_claims, remove_claims)
+# ---------------------------------------------------------------------------
+
+
+def test_add_claims_adds_missing_claims():
+ """FR-13: add_claims adds claims that are not already present."""
+ signer = _make_signer(add_claims={"tenant_id": "acme", "env": "prod"})
+ user_dict = _make_user_api_key_dict()
+
+ claims = signer._build_claims(user_dict, {})
+
+ assert claims["tenant_id"] == "acme"
+ assert claims["env"] == "prod"
+
+
+def test_add_claims_does_not_override_existing():
+ """FR-13: add_claims does NOT override claims that already exist (e.g., iss)."""
+ signer = _make_signer(
+ issuer="https://litellm.example.com",
+ add_claims={"iss": "https://imposter.example.com"},
+ )
+ user_dict = _make_user_api_key_dict()
+
+ claims = signer._build_claims(user_dict, {})
+
+ # add_claims must not override the signer-built issuer
+ assert claims["iss"] == "https://litellm.example.com"
+
+
+def test_set_claims_overrides_existing():
+ """FR-13: set_claims overrides existing claims (including signer-built ones)."""
+ signer = _make_signer(
+ audience="mcp",
+ set_claims={"aud": "custom-audience", "custom_key": "custom_val"},
+ )
+ user_dict = _make_user_api_key_dict()
+
+ claims = signer._build_claims(user_dict, {})
+
+ assert claims["aud"] == "custom-audience"
+ assert claims["custom_key"] == "custom_val"
+
+
+def test_remove_claims_strips_specified_claims():
+ """FR-13: remove_claims strips listed claim names from the output JWT."""
+ signer = _make_signer(remove_claims=["email", "act"])
+ user_dict = _make_user_api_key_dict(user_email="alice@example.com")
+
+ claims = signer._build_claims(user_dict, {})
+
+ assert "email" not in claims
+ assert "act" not in claims
+
+
+def test_claim_operations_order():
+ """FR-13: Operations are applied add → set → remove. Set wins over add; remove wins over both."""
+ signer = _make_signer(
+ add_claims={"x": "from-add"},
+ set_claims={"x": "from-set", "y": "from-set"},
+ remove_claims=["y"],
+ )
+ user_dict = _make_user_api_key_dict()
+
+ claims = signer._build_claims(user_dict, {})
+
+ # add adds x, then set overrides x
+ assert claims["x"] == "from-set"
+ # set adds y, then remove strips it
+ assert "y" not in claims
+
+
+# ---------------------------------------------------------------------------
+# FR-14: Two-token model (channel token)
+# ---------------------------------------------------------------------------
+
+
+def test_build_claims_uses_channel_token_as_act():
+ """FR-14: When channel_token_claims is provided, act.sub comes from its sub."""
+ signer = _make_signer()
+ user_dict = _make_user_api_key_dict(team_id="team-ignored")
+ channel_claims = {"sub": "agent-service-001", "client_id": "agent-client"}
+
+ claims = signer._build_claims(user_dict, {}, channel_token_claims=channel_claims)
+
+ assert claims["act"]["sub"] == "agent-service-001"
+ assert claims["act"]["client_id"] == "agent-client"
+
+
+def test_build_claims_channel_token_fallback_to_client_id():
+ """FR-14: Falls back to client_id when sub is absent in channel token."""
+ signer = _make_signer()
+ user_dict = _make_user_api_key_dict()
+ channel_claims = {"client_id": "m2m-client-xyz"}
+
+ claims = signer._build_claims(user_dict, {}, channel_token_claims=channel_claims)
+
+ assert claims["act"]["sub"] == "m2m-client-xyz"
+
+
+@pytest.mark.asyncio
+async def test_hook_reads_channel_token_from_raw_headers():
+ """FR-14: async_pre_call_hook picks up X-Channel-Token from mcp_raw_headers."""
+ signer = _make_signer()
+
+ # Build a valid channel token to inject
+ channel_payload = {"sub": "channel-agent", "exp": int(time.time()) + 300}
+ channel_token = jwt.encode(channel_payload, signer._private_key, algorithm="RS256")
+
+ user_dict = _make_user_api_key_dict()
+ data = {
+ "mcp_tool_name": "search",
+ "mcp_raw_headers": {"x-channel-token": channel_token},
+ }
+
+ result = await signer.async_pre_call_hook(
+ user_api_key_dict=user_dict,
+ cache=MagicMock(),
+ data=data,
+ call_type="call_mcp_tool",
+ )
+
+ assert isinstance(result, dict)
+ token = result["extra_headers"]["Authorization"].removeprefix("Bearer ")
+ decoded = _decode_unverified(token)
+ # act should come from channel token (no sig verification since no discovery_uri)
+ assert decoded["act"]["sub"] == "channel-agent"
+
+
+# ---------------------------------------------------------------------------
+# FR-15: Required/optional claim validation
+# ---------------------------------------------------------------------------
+
+
+def test_required_claims_passes_when_all_present():
+ """FR-15: No error when all required_claims are present in jwt_claims."""
+ from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
+ _validate_required_claims,
+ )
+
+ _validate_required_claims(
+ jwt_claims={"sub": "alice", "email": "alice@example.com"},
+ required_claims=["sub", "email"],
+ ) # must not raise
+
+
+def test_required_claims_raises_when_missing():
+ """FR-15: ValueError when a required claim is missing from jwt_claims."""
+ from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
+ _validate_required_claims,
+ )
+
+ with pytest.raises(ValueError, match="missing required_claims"):
+ _validate_required_claims(
+ jwt_claims={"sub": "alice"},
+ required_claims=["sub", "groups"],
+ )
+
+
+def test_required_claims_raises_when_no_jwt_claims():
+ """FR-15: ValueError when required_claims set but no JWT claims present."""
+ from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
+ _validate_required_claims,
+ )
+
+ with pytest.raises(ValueError, match="no JWT claims"):
+ _validate_required_claims(
+ jwt_claims=None,
+ required_claims=["sub"],
+ )
+
+
+def test_required_claims_empty_list_always_passes():
+ """FR-15: Empty required_claims never raises (even with None jwt_claims)."""
+ from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
+ _validate_required_claims,
+ )
+
+ _validate_required_claims(jwt_claims=None, required_claims=[]) # must not raise
+
+
+@pytest.mark.asyncio
+async def test_hook_raises_when_required_claims_missing():
+ """FR-15: async_pre_call_hook raises ValueError when required_claims are absent."""
+ signer = _make_signer(required_claims=["groups"])
+ user_dict = _make_user_api_key_dict()
+ user_dict.jwt_claims = {"sub": "alice"} # no 'groups' claim
+
+ with pytest.raises(ValueError, match="missing required_claims"):
+ await signer.async_pre_call_hook(
+ user_api_key_dict=user_dict,
+ cache=MagicMock(),
+ data={"mcp_tool_name": "do_thing"},
+ call_type="call_mcp_tool",
+ )
+
+
+# ---------------------------------------------------------------------------
+# FR-9: Debug headers
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.asyncio
+async def test_hook_emits_debug_header_by_default():
+ """FR-9: x-litellm-mcp-debug header is emitted by default."""
+ signer = _make_signer()
+ user_dict = _make_user_api_key_dict()
+
+ result = await signer.async_pre_call_hook(
+ user_api_key_dict=user_dict,
+ cache=MagicMock(),
+ data={"mcp_tool_name": "test_tool"},
+ call_type="call_mcp_tool",
+ )
+
+ assert isinstance(result, dict)
+ assert "x-litellm-mcp-debug" in result["extra_headers"]
+ debug = json.loads(result["extra_headers"]["x-litellm-mcp-debug"])
+ assert debug["signer"] == "mcp_jwt_signer"
+ assert debug["kid"] == signer._kid
+ assert debug["issuer"] == signer.issuer
+ assert "sub" in debug
+ assert "mode" in debug
+
+
+@pytest.mark.asyncio
+async def test_hook_omits_debug_header_when_disabled():
+ """FR-9: x-litellm-mcp-debug is not emitted when debug_header=False."""
+ signer = _make_signer(debug_header=False)
+ user_dict = _make_user_api_key_dict()
+
+ result = await signer.async_pre_call_hook(
+ user_api_key_dict=user_dict,
+ cache=MagicMock(),
+ data={"mcp_tool_name": "test_tool"},
+ call_type="call_mcp_tool",
+ )
+
+ assert isinstance(result, dict)
+ assert "x-litellm-mcp-debug" not in result["extra_headers"]
+
+
+@pytest.mark.asyncio
+async def test_debug_header_reports_re_sign_mode_when_jwt_claims_present():
+ """FR-9: Debug header shows mode=re-sign when upstream jwt_claims are available."""
+ signer = _make_signer()
+ user_dict = _make_user_api_key_dict()
+ user_dict.jwt_claims = {"sub": "upstream-user"}
+
+ result = await signer.async_pre_call_hook(
+ user_api_key_dict=user_dict,
+ cache=MagicMock(),
+ data={"mcp_tool_name": "tool"},
+ call_type="call_mcp_tool",
+ )
+
+ assert isinstance(result, dict)
+ debug = json.loads(result["extra_headers"]["x-litellm-mcp-debug"])
+ assert debug["mode"] == "re-sign"
+
+
+# ---------------------------------------------------------------------------
+# FR-10: Configurable scope (allowed_tools)
+# ---------------------------------------------------------------------------
+
+
+def test_allowed_tools_restricts_scope():
+ """FR-10: allowed_tools overrides auto-generated scope to admin-defined list."""
+ signer = _make_signer(allowed_tools=["search_web", "get_weather"])
+ user_dict = _make_user_api_key_dict()
+
+ claims = signer._build_claims(user_dict, {"mcp_tool_name": "search_web"})
+
+ scopes = set(claims["scope"].split())
+ assert "mcp:tools/search_web:call" in scopes
+ assert "mcp:tools/search_web:list" in scopes
+ assert "mcp:tools/get_weather:call" in scopes
+ # tools/list should NOT be present during a specific tool call
+ assert "mcp:tools/list" not in scopes
+
+
+def test_allowed_tools_grants_list_when_no_tool_name():
+ """FR-10: allowed_tools grants mcp:tools/list when not calling a specific tool."""
+ signer = _make_signer(allowed_tools=["search_web"])
+ user_dict = _make_user_api_key_dict()
+
+ claims = signer._build_claims(user_dict, {})
+
+ assert "mcp:tools/list" in claims["scope"]
+
+
+def test_allowed_tools_empty_falls_back_to_auto_scope():
+ """FR-10: Empty allowed_tools uses original auto-generated scope."""
+ signer = _make_signer(allowed_tools=[])
+ user_dict = _make_user_api_key_dict()
+
+ claims = signer._build_claims(user_dict, {"mcp_tool_name": "some_tool"})
+
+ # Auto-scope: call + tool-specific
+ assert "mcp:tools/call" in claims["scope"]
+ assert "mcp:tools/some_tool:call" in claims["scope"]
+
+
+# ---------------------------------------------------------------------------
+# FR-5: access_token_discovery_uri with jwt_claims integration
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.asyncio
+async def test_hook_uses_jwt_claims_for_sub_when_discovery_uri_set():
+ """FR-5: In re-sign mode, hook uses upstream jwt_claims.sub rather than user_id."""
+ signer = _make_signer(
+ access_token_discovery_uri="https://login.example.com/.well-known/openid-configuration"
+ )
+ user_dict = _make_user_api_key_dict(user_id="litellm-internal-user")
+ user_dict.jwt_claims = {"sub": "upstream-alice@corp.com"}
+
+ result = await signer.async_pre_call_hook(
+ user_api_key_dict=user_dict,
+ cache=MagicMock(),
+ data={"mcp_tool_name": "tool"},
+ call_type="call_mcp_tool",
+ )
+
+ assert isinstance(result, dict)
+ token = result["extra_headers"]["Authorization"].removeprefix("Bearer ")
+ decoded = _decode_unverified(token)
+ # sub must come from upstream jwt_claims
+ assert decoded["sub"] == "upstream-alice@corp.com"