feat(guardrails): MCPJWTSigner - built-in guardrail for zero trust MCP auth (#23897)

* Allow pre_mcp_call guardrail hooks to mutate outbound MCP headers

* Enhance MCPServerManager to support hook-modified arguments and extra headers. Update tests to validate argument mutation and header injection behavior, including warnings for OpenAPI-backed servers when headers are present.

* Refactor MCPServerManager to raise HTTPException for extra headers in OpenAPI-backed servers. Update tests to reflect this change, ensuring proper exception handling instead of logging warnings.

* Allow pre_mcp_call guardrail hooks to mutate outbound MCP headers

* Enhance MCPServerManager to support hook-modified arguments and extra headers. Update tests to validate argument mutation and header injection behavior, including warnings for OpenAPI-backed servers when headers are present.

* Refactor MCPServerManager to raise HTTPException for extra headers in OpenAPI-backed servers. Update tests to reflect this change, ensuring proper exception handling instead of logging warnings.

* feat(guardrails): add MCPJWTSigner built-in guardrail for zero trust MCP auth

Signs outbound MCP tool calls with a LiteLLM-issued RS256 JWT so MCP servers
can trust a single signing authority instead of every upstream IdP.

Enable in config.yaml:
  guardrails:
    - guardrail_name: mcp-jwt-signer
      litellm_params:
        guardrail: mcp_jwt_signer
        mode: pre_mcp_call
        default_on: true

JWT carries sub (user_id), act.sub (team_id, RFC 8693), tool-level scope, iss,
aud, iat/exp/nbf. RSA-2048 keypair auto-generated at startup unless
MCP_JWT_SIGNING_KEY env var is set.

Adds /.well-known/jwks.json endpoint and jwks_uri to /.well-known/openid-configuration
so MCP servers can verify LiteLLM-issued tokens via OIDC discovery.

* Update MCPServerManager to raise HTTPException with status code 400 for extra headers in OpenAPI-backed servers. Adjust tests to verify the correct status code and exception message.

* fix: address P1 issues in MCPJWTSigner

- OpenAPI servers: warn + skip header injection instead of 500
- JWKS Cache-Control: 5min for auto-generated keys, 1h for persistent
- sub claim: fallback to apikey:{token_hash} for anonymous callers
- ttl_seconds: validate > 0 at init time

* docs: add MCP zero trust auth guide with architecture diagram

* docs: add FastMCP JWT verification guide to zero trust doc

* fix: address remaining Greptile review issues (round 2)

- mcp_server_manager: warn when hook Authorization overwrites existing header
- __init__: remove _mcp_jwt_signer_instance from __all__ (private internal)
- discoverable_endpoints: copy dict instead of mutating in-place on OIDC augmentation
- test docstring: reflect warn-and-continue behavior for OpenAPI servers
- test: update scope assertions for least-privilege (no mcp:tools/list on tool-call JWTs)

* fix: address Greptile round 3 feedback

- initialize_guardrail: validate mode='pre_mcp_call' at init time — misconfigured
  mode silently bypasses JWT injection, which is a zero-trust bypass
- _build_claims: remove duplicate inline 'import re' (module-level import already present)
- _types.py: add TODO comment explaining jwt_claims is forward-compat plumbing
  for a follow-up PR that will forward upstream IdP claims into outbound MCP JWTs

* feat(mcp_jwt_signer): add verify+re-sign, claim ops, two-token model, configurable scopes

Addresses all missing pieces from the scoping doc review:

FR-5 (Verify + re-sign): MCPJWTSigner now accepts access_token_discovery_uri
and token_introspection_endpoint.  When set, the incoming Bearer token is
extracted from raw_headers (threaded through pre_call_tool_check), verified
against the IdP's JWKS (JWT) or introspected (opaque), and only re-signed if
valid.  Falls back to user_api_key_dict.jwt_claims for LiteLLM JWT-auth mode.

FR-12 (Configurable end-user identity mapping): end_user_claim_sources
ordered list drives sub resolution — sources: token:<claim>, litellm:user_id,
litellm:email, litellm:end_user_id, litellm:team_id.

FR-13 (Claim operations): add_claims (insert-if-absent), set_claims (always
override), remove_claims (delete) applied in that order.

FR-14 (Two-token model): channel_token_audience + channel_token_ttl issue a
second JWT injected as x-mcp-channel-token: Bearer <token>.

FR-15 (Incoming claim validation): required_claims raises HTTP 403 when any
listed claim is absent; optional_claims passes listed claims from verified
token into the outbound JWT.

FR-9 (Debug headers): debug_headers: true emits x-litellm-mcp-debug with kid,
sub, iss, exp, scope.

FR-10 (Configurable scopes): allowed_scopes replaces auto-generation.  Also
fixed: tool-call JWTs no longer grant mcp:tools/list (overpermission).

P1 fixes:
- proxy/utils.py: _convert_mcp_hook_response_to_kwargs merges rather than
  replaces extra_headers, preserving headers from prior guardrails.
- mcp_server_manager.py: warns when hook injects Authorization alongside a
  server-configured authentication_token (previously silent).
- mcp_server_manager.py: pre_call_tool_check now accepts raw_headers and
  extracts incoming_bearer_token so FR-5 verification has the raw token.
- proxy/utils.py: remove stray inline import inspect inside loop (pre-existing
  lint error, now cleaned up).

Tests: 43 passing (28 new tests covering all FR flags + P1 fixes).

* feat(mcp_jwt_signer): add verify+re-sign, claim ops, two-token model, configurable scopes (core)

Remaining files from the FR implementation:

mcp_jwt_signer.py — full rewrite with all new params:
  FR-5:  access_token_discovery_uri, token_introspection_endpoint,
         verify_issuer, verify_audience + _verify_incoming_jwt(),
         _introspect_opaque_token()
  FR-12: end_user_claim_sources ordered resolution chain
  FR-13: add_claims, set_claims, remove_claims
  FR-14: channel_token_audience, channel_token_ttl → x-mcp-channel-token
  FR-15: required_claims (raises 403), optional_claims (passthrough)
  FR-9:  debug_headers → x-litellm-mcp-debug
  FR-10: allowed_scopes; tool-call JWTs no longer over-grant tools/list

mcp_server_manager.py:
  - pre_call_tool_check gains raw_headers param to extract incoming_bearer_token
  - Silent Authorization override warning fixed: now fires when server has
    authentication_token AND hook injects Authorization

tests/test_mcp_jwt_signer.py:
  28 new tests covering all FR flags + P1 fixes (43 total, all passing)

* fix(mcp_jwt_signer): address pre-landing review issues

- Remove stale TODO comment on UserAPIKeyAuth.jwt_claims — the field is
  already populated and consumed by MCPJWTSigner in the same PR
- Fix _get_oidc_discovery to only cache the OIDC discovery doc when
  jwks_uri is present; a malformed/empty doc now retries on the next
  request instead of being permanently cached until proxy restart
- Add FR-5 test coverage for _fetch_jwks (cache hit/miss),
  _get_oidc_discovery (cache/no-cache on bad doc), _verify_incoming_jwt
  (valid token, expired token), _introspect_opaque_token (active,
  inactive, no endpoint), and the end-to-end 401 hook path — 53 tests
  total, all passing

* docs(mcp_zero_trust): rewrite as use-case guide covering all new JWT signer features

Add scenario-driven sections for each new config area:
- Verify+re-sign with Okta/Azure AD (access_token_discovery_uri,
  end_user_claim_sources, token_introspection_endpoint)
- Enforcing caller attributes with required_claims / optional_claims
- Adding metadata via add_claims / set_claims / remove_claims
- Two-token model for AWS Bedrock AgentCore Gateway
  (channel_token_audience / channel_token_ttl)
- Controlling scopes with allowed_scopes
- Debugging JWT rejections with debug_headers

Update JWT claims table to reflect configurable sub (end_user_claim_sources)

* fix(mcp_jwt_signer): wire all config.yaml params through initialize_guardrail

The factory was only passing issuer/audience/ttl_seconds to MCPJWTSigner.
All FR-5/9/10/12/13/14/15 params (access_token_discovery_uri,
end_user_claim_sources, add/set/remove_claims, channel_token_audience,
required/optional_claims, debug_headers, allowed_scopes, etc.) were
silently dropped, making every advertised advanced feature non-functional
when loaded from config.yaml.

Add regression test that asserts every param is wired through correctly.

* docs(mcp_zero_trust): add hero image

* docs(mcp_zero_trust): apply Linear-style edits

- Lead with the problem (unsigned direct calls bypass access controls)
- Shorter statement section headers instead of question-form headers
- Move diagram/OIDC discovery block after the reader is bought in
- Add 'read further only if you need to' callout after basic setup
- Two-token section now opens from the user problem not product jargon
- Add concrete 403 error response example in required_claims section
- Debug section opens from the symptom (MCP server returning 401)
- Lowercase claims reference header for consistency

* fix(mcp_jwt_signer): fix algorithm confusion attack + add OIDC discovery 24h TTL

- Remove alg from unverified JWT header; use signing_jwk.algorithm_name from JWKS key instead.
  Reading alg from attacker-controlled headers enables alg:none / HS256 confusion attacks.
- Add _oidc_discovery_fetched_at timestamp and _OIDC_DISCOVERY_TTL = 86400 (24h).
  Without a TTL the cached discovery doc never refreshes, so IdP key rotation is invisible.

---------

Co-authored-by: Noah Nistler <60981020+noahnistler@users.noreply.github.com>
This commit is contained in:
Ishaan Jaff 2026-03-17 17:30:16 -07:00 • committed by GitHub
parent 2f7dcbaeb9
commit d9a6036162
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 3224 additions and 9 deletions

View file

@ -0,0 +1,294 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# MCP Zero Trust Auth (JWT Signer)
![Zero Trust MCP Gateway](/img/mcp_zero_trust_gateway.png)
MCP servers have no built-in way to verify that a request actually came through LiteLLM. Without this guardrail, any client that can reach your MCP server directly can call tools — bypassing your access controls entirely.
`MCPJWTSigner` fixes this. It signs every outbound tool call with a short-lived RS256 JWT. Your MCP server verifies the signature against LiteLLM's public key. Requests that didn't go through LiteLLM have no valid signature and are rejected.
---
## Basic setup
Add the guardrail to your config and point your MCP server at LiteLLM's JWKS endpoint. Every tool call gets a signed JWT automatically — no changes needed on the client side.
```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" # defaults to request base URL
audience: "mcp" # default: "mcp"
ttl_seconds: 300 # default: 300
```
**Bring your own signing key** — recommended for production. Auto-generated keys are lost on restart.
```bash
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"
```
**Build a verified MCP server with [FastMCP](https://gofastmcp.com):**
```python title="weather_server.py"
from fastmcp import FastMCP, Context
from fastmcp.server.auth.providers.jwt import JWTVerifier
auth = JWTVerifier(
jwks_uri="https://my-litellm.example.com/.well-known/jwks.json",
issuer="https://my-litellm.example.com",
audience="mcp",
algorithm="RS256",
)
mcp = FastMCP("weather-server", auth=auth)
@mcp.tool()
async def get_weather(city: str, ctx: Context) -> str:
caller = ctx.client_id # JWT `sub` — the verified user identity
return f"Weather in {city}: sunny, 72°F (requested by {caller})"
if __name__ == "__main__":
mcp.run(transport="http", host="0.0.0.0", port=8000)
```
FastMCP fetches the JWKS automatically and re-fetches when the signing key changes.
LiteLLM publishes OIDC discovery so MCP servers find the key without any manual configuration:
```
GET /.well-known/openid-configuration → { "jwks_uri": "https://<litellm>/.well-known/jwks.json" }
GET /.well-known/jwks.json → { "keys": [{ "kty": "RSA", "alg": "RS256", ... }] }
```
> **Read further only if you need to:** thread a corporate IdP identity into the JWT, enforce specific claims on callers, add custom metadata, use AWS Bedrock AgentCore Gateway, or debug JWT rejections.
---
## Thread IdP identity into MCP JWTs
By default the outbound JWT `sub` is LiteLLM's internal `user_id`. If your users authenticate with Okta, Azure AD, or another IdP, the MCP server sees a LiteLLM-internal ID — not the user's email or employee ID.
With verify+re-sign, LiteLLM validates the incoming IdP token first, then builds the outbound JWT using the real identity claims from that token. The MCP server gets the user's actual identity without ever having to trust the original IdP directly.
```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"
# Validate the incoming Bearer token against the IdP
access_token_discovery_uri: "https://login.microsoftonline.com/{tenant}/v2.0/.well-known/openid-configuration"
verify_issuer: "https://login.microsoftonline.com/{tenant}/v2.0"
verify_audience: "api://my-app"
# Which claim to use for `sub` in the outbound JWT — first non-empty value wins
end_user_claim_sources:
- "token:sub" # from the verified incoming JWT
- "token:email" # fallback to email
- "litellm:user_id" # last resort: LiteLLM's internal user_id
```
If the incoming token is **opaque** (not a JWT — some IdPs issue these), add an introspection endpoint. LiteLLM will POST the token to it (RFC 7662) and use the returned claims:
```yaml
token_introspection_endpoint: "https://idp.example.com/oauth2/introspect"
```
**Supported `end_user_claim_sources` values:**
| Source | Resolves to |
|--------|-------------|
| `token:<claim>` | Any claim from the verified incoming JWT (e.g. `token:sub`, `token:email`, `token:oid`) |
| `litellm:user_id` | LiteLLM's internal user ID |
| `litellm:email` | User email from LiteLLM auth context |
| `litellm:end_user_id` | End-user ID if set separately |
| `litellm:team_id` | Team ID from LiteLLM auth context |
---
## Block callers missing required attributes
Some MCP servers expose sensitive operations that should only be reachable by verified employees — not service accounts, not external API keys. You can enforce this at the LiteLLM layer so the MCP server never receives the request at all.
`required_claims` rejects with `403` if the incoming token is missing any listed claim. `optional_claims` forwards claims that are useful but not mandatory.
```yaml title="config.yaml"
guardrails:
- guardrail_name: mcp-jwt-signer
litellm_params:
guardrail: mcp_jwt_signer
mode: pre_mcp_call
default_on: true
access_token_discovery_uri: "https://idp.example.com/.well-known/openid-configuration"
# Service accounts without `employee_id` are blocked before the tool runs
required_claims:
- "sub"
- "employee_id"
# Forward these into the outbound JWT when present — skipped silently if absent
optional_claims:
- "groups"
- "department"
```
**What the client sees when blocked:**
```json
HTTP 403
{ "error": "MCPJWTSigner: incoming token is missing required claims: ['employee_id']. Configure the IdP to include these claims." }
```
---
## Add custom metadata to every JWT
Your MCP server may need context that LiteLLM doesn't carry natively — which deployment sent the request, a tenant ID, an environment tag. Use claim operations to inject, override, or strip claims from the outbound JWT.
```yaml title="config.yaml"
guardrails:
- guardrail_name: mcp-jwt-signer
litellm_params:
guardrail: mcp_jwt_signer
mode: pre_mcp_call
default_on: true
# add: insert only when the key is not already in the JWT
add_claims:
deployment_id: "prod-us-east-1"
tenant_id: "acme-corp"
# set: always override — even if the claim came from the incoming token
set_claims:
env: "production"
# remove: strip claims the MCP server shouldn't see
remove_claims:
- "nbf" # some validators reject nbf; remove it if yours does
```
Operations run in order — `add_claims` → `set_claims` → `remove_claims`. `set_claims` always wins over `add_claims`; `remove_claims` beats both.
---
## AWS Bedrock AgentCore Gateway
Bedrock AgentCore Gateway uses two separate JWTs: one to authenticate the transport connection and another to authorize tool calls. They need different `aud` values and TTLs — a single JWT won't work for both.
LiteLLM can issue both in one hook and inject them into separate headers:
```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"
audience: "mcp-resource" # for the MCP resource layer
ttl_seconds: 300
# Second JWT for the transport channel — same sub/act/scope, different aud + TTL
channel_token_audience: "bedrock-agentcore-gateway"
channel_token_ttl: 60 # transport tokens should be short-lived
```
LiteLLM injects two headers on every tool call:
- `Authorization: Bearer <resource-token>` — audience `mcp-resource`, TTL 300s
- `x-mcp-channel-token: Bearer <channel-token>` — audience `bedrock-agentcore-gateway`, TTL 60s
Both tokens are signed with the same LiteLLM key, so your MCP server only needs to trust one JWKS endpoint.
---
## Control which scopes go into the JWT
By default LiteLLM generates least-privilege scopes per request:
- Tool call → `mcp:tools/call mcp:tools/{name}:call`
- List tools → `mcp:tools/call mcp:tools/list`
If your MCP server does its own scope enforcement and needs a specific format, set `allowed_scopes` to replace auto-generation entirely:
```yaml title="config.yaml"
guardrails:
- guardrail_name: mcp-jwt-signer
litellm_params:
guardrail: mcp_jwt_signer
mode: pre_mcp_call
default_on: true
allowed_scopes:
- "mcp:tools/call"
- "mcp:tools/list"
- "mcp:admin"
```
Every JWT carries exactly those scopes regardless of which tool is being called.
---
## Debug JWT rejections
Your MCP server is returning 401 and you're not sure what's in the JWT. Enable `debug_headers` and LiteLLM adds a `x-litellm-mcp-debug` response header with the key claims that were signed:
```yaml title="config.yaml"
guardrails:
- guardrail_name: mcp-jwt-signer
litellm_params:
guardrail: mcp_jwt_signer
mode: pre_mcp_call
default_on: true
debug_headers: true
```
Response header:
```
x-litellm-mcp-debug: v=1; kid=a3f1b2c4d5e6f708; sub=alice@corp.com; iss=https://my-litellm.example.com; exp=1712345678; scope=mcp:tools/call mcp:tools/get_weather:call
```
Check that `kid` matches what the MCP server fetched from JWKS, `iss`/`aud` match your server's expected values, and `exp` hasn't passed. Disable in production — the header leaks claim metadata.
---
## JWT claims reference
| Claim | Value |
|-------|-------|
| `iss` | `issuer` config value (or request base URL) |
| `aud` | `audience` config value (default: `"mcp"`) |
| `sub` | Resolved via `end_user_claim_sources` (default: `user_id` → api-key hash → `"litellm-proxy"`) |
| `act.sub` | `team_id` → `org_id` → `"litellm-proxy"` (RFC 8693 delegation) |
| `email` | `user_email` from LiteLLM auth context (when available) |
| `scope` | Auto-generated per tool call, or `allowed_scopes` when set |
| `iat`, `exp`, `nbf` | Standard timing claims (RFC 7519) |
---
## Limitations
- **OpenAPI-backed MCP servers** (`spec_path` set) do not support JWT injection. LiteLLM logs a warning and skips the header. Use SSE/HTTP transport servers to get full JWT injection.
- The keypair is **in-memory by default** and rotated on each restart unless `MCP_JWT_SIGNING_KEY` is set. FastMCP's `JWTVerifier` handles key rotation transparently via JWKS key ID matching.
---
## 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

Binary file not shown.

After

Width:  |  Height:  |  Size: 294 KiB

View file

@ -636,6 +636,7 @@ const sidebars = {
"mcp_control",
"mcp_cost",
"mcp_guardrail",
"mcp_zero_trust",
"mcp_troubleshoot",
]
},

View file

@ -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

View file

@ -1908,7 +1908,15 @@ 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
"""
## check if the tool is allowed or banned for the given server
if not self.check_allowed_or_banned_tools(name, server):
raise HTTPException(
@ -1932,6 +1940,14 @@ class MCPServerManager:
server=server,
)
# Extract incoming Bearer token from raw request headers so
# guardrails like MCPJWTSigner can verify + re-sign it (FR-5).
normalized_raw = {k.lower(): v for k, v in (raw_headers or {}).items()}
incoming_bearer_token: Optional[str] = None
auth_hdr = normalized_raw.get("authorization", "")
if auth_hdr.lower().startswith("bearer "):
incoming_bearer_token = auth_hdr[len("bearer "):]
pre_hook_kwargs = {
"name": name,
"arguments": arguments,
@ -1957,6 +1973,7 @@ class MCPServerManager:
if user_api_key_auth
else None
),
"incoming_bearer_token": incoming_bearer_token,
}
# Create MCP request object for processing
@ -1969,6 +1986,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 +2002,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 +2015,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 +2069,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 +2084,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 +2142,31 @@ 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 hook_extra_headers:
if "Authorization" in extra_headers:
verbose_logger.warning(
"MCPServerManager: hook_extra_headers 'Authorization' will overwrite "
"the existing Authorization header from static_headers. "
"The hook JWT will take precedence."
)
elif server_auth_header is not None:
# server_auth_header is passed separately to _create_mcp_client as
# auth_value. Both will reach the upstream server — warn so admins
# know two Authorization credentials are being sent.
verbose_logger.warning(
"MCPServerManager: hook_extra_headers injects 'Authorization' while "
"server '%s' already has a configured authentication_token. "
"Both credentials will be sent; the hook header is in extra_headers "
"and the server token is in auth_value — the upstream server decides "
"which one wins. Consider unsetting authentication_token if you want "
"the hook JWT to be the sole credential.",
mcp_server.server_name or mcp_server.name,
)
extra_headers.update(hook_extra_headers)
stdio_env = self._build_stdio_env(mcp_server, raw_headers)
client = await self._create_mcp_client(
@ -2201,15 +2252,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 +2282,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 +2310,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

View file

@ -2471,6 +2471,9 @@ class UserAPIKeyAuth(
Any
] = None # Expanded created_by user when expand=user is used
end_user_object_permission: Optional[LiteLLM_ObjectPermissionTable] = None
# Decoded upstream IdP claims (groups, roles, etc.) propagated by JWT auth machinery
# and forwarded into outbound tokens by guardrails such as MCPJWTSigner.
jwt_claims: Optional[Dict] = None
model_config = ConfigDict(arbitrary_types_allowed=True)

View file

@ -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

View file

@ -0,0 +1,84 @@
"""MCP JWT Signer guardrail — built-in LiteLLM guardrail for zero trust MCP auth."""
from typing import TYPE_CHECKING
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 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."
)
optional_params = getattr(litellm_params, "optional_params", None)
def _get(key): # type: ignore[no-untyped-def]
if optional_params is not None:
v = getattr(optional_params, key, None)
if v is not None:
return v
return getattr(litellm_params, key, None)
signer = MCPJWTSigner(
guardrail_name=guardrail_name,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
# Core signing
issuer=_get("issuer"),
audience=_get("audience"),
ttl_seconds=_get("ttl_seconds"),
# FR-5: verify + re-sign
access_token_discovery_uri=_get("access_token_discovery_uri"),
token_introspection_endpoint=_get("token_introspection_endpoint"),
verify_issuer=_get("verify_issuer"),
verify_audience=_get("verify_audience"),
# FR-12: end-user identity mapping
end_user_claim_sources=_get("end_user_claim_sources"),
# FR-13: claim operations
add_claims=_get("add_claims"),
set_claims=_get("set_claims"),
remove_claims=_get("remove_claims"),
# FR-14: two-token model
channel_token_audience=_get("channel_token_audience"),
channel_token_ttl=_get("channel_token_ttl"),
# FR-15: incoming claim validation
required_claims=_get("required_claims"),
optional_claims=_get("optional_claims"),
# FR-9: debug headers
debug_headers=_get("debug_headers") or False,
# FR-10: configurable scopes
allowed_scopes=_get("allowed_scopes"),
)
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",
]

View file

@ -0,0 +1,889 @@
"""
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.
Usage in config.yaml:
guardrails:
- guardrail_name: "mcp-jwt-signer"
litellm_params:
guardrail: mcp_jwt_signer
mode: "pre_mcp_call"
default_on: true
# Core signing config
issuer: "https://my-litellm.example.com" # optional
audience: "mcp" # optional
ttl_seconds: 300 # optional
# FR-5: Verify + re-sign — validate incoming Bearer token before signing
access_token_discovery_uri: "https://idp.example.com/.well-known/openid-configuration"
token_introspection_endpoint: "https://idp.example.com/introspect" # opaque tokens
verify_issuer: "https://idp.example.com" # expected iss in incoming JWT
verify_audience: "api://my-app" # expected aud in incoming JWT
# FR-12: End-user identity mapping — ordered resolution chain
# Supported: token:<claim>, litellm:user_id, litellm:email,
# litellm:end_user_id, litellm:team_id
end_user_claim_sources:
- "token:sub"
- "token:email"
- "litellm:user_id"
# FR-13: Claim operations
add_claims: # add if key not already present in the JWT
deployment_id: "prod-001"
set_claims: # always set (overrides computed value)
env: "production"
remove_claims: # remove from final JWT
- "nbf"
# FR-14: Two-token model — issue a second JWT for the MCP transport channel
channel_token_audience: "bedrock-gateway"
channel_token_ttl: 60
# FR-15: Incoming claim validation — enforce required IdP claims
required_claims:
- "sub"
- "email"
optional_claims: # pass through from jwt_claims into outbound JWT
- "groups"
- "roles"
# FR-9: Debug headers
debug_headers: false # emit x-litellm-mcp-debug header when true
# FR-10: Configurable scopes — explicit list replaces auto-generation
allowed_scopes:
- "mcp:tools/call"
- "mcp:tools/list"
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 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
# Simple in-memory JWKS cache: keyed by JWKS URI → (keys_list, fetched_at).
_jwks_cache: Dict[str, tuple] = {}
_JWKS_CACHE_TTL = 3600 # 1 hour
def get_mcp_jwt_signer() -> Optional["MCPJWTSigner"]:
"""Return the active MCPJWTSigner singleton, or None if not initialized."""
return _mcp_jwt_signer_instance
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]
async def _fetch_jwks(jwks_uri: str) -> List[Dict[str, Any]]:
"""
Fetch and cache a JWKS from the given URI.
Results are cached for _JWKS_CACHE_TTL seconds to avoid hammering the IdP.
"""
now = time.time()
cached = _jwks_cache.get(jwks_uri)
if cached is not None:
keys, fetched_at = cached
if now - fetched_at < _JWKS_CACHE_TTL:
return keys # type: ignore[return-value]
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
resp = await client.get(jwks_uri, headers={"Accept": "application/json"})
resp.raise_for_status()
keys = resp.json().get("keys", [])
_jwks_cache[jwks_uri] = (keys, now)
return keys # type: ignore[return-value]
async def _fetch_oidc_discovery(discovery_uri: str) -> Dict[str, Any]:
"""Fetch an OIDC discovery document and return its parsed JSON."""
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
resp = await client.get(discovery_uri, headers={"Accept": "application/json"})
resp.raise_for_status()
return resp.json() # type: ignore[return-value]
class MCPJWTSigner(CustomGuardrail):
"""
Built-in LiteLLM guardrail that signs outbound MCP requests with a
LiteLLM-issued RS256 JWT, enabling zero trust authentication.
MCP servers verify tokens using liteLLM's OIDC discovery endpoint and
JWKS endpoint rather than trusting each upstream IdP directly.
The signed JWT carries:
- iss: LiteLLM issuer identifier
- aud: MCP audience (configurable)
- sub: End-user identity (resolved via end_user_claim_sources, RFC 8693)
- act: Actor/agent identity (team_id or org_id, RFC 8693 delegation)
- scope: Tool-level access scopes (configurable via allowed_scopes)
- iat, exp, nbf: Standard timing claims
Feature set:
FR-5: Verify + re-sign (access_token_discovery_uri, token_introspection_endpoint)
FR-9: Debug headers (debug_headers)
FR-10: Configurable scopes (allowed_scopes)
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 (channel_token_audience, channel_token_ttl)
FR-15: Incoming claim validation (required_claims, optional_claims)
"""
ALGORITHM = "RS256"
DEFAULT_TTL = 300
DEFAULT_AUDIENCE = "mcp"
SIGNING_KEY_ENV = "MCP_JWT_SIGNING_KEY"
def __init__(
self,
# Core signing config
issuer: Optional[str] = None,
audience: Optional[str] = None,
ttl_seconds: Optional[int] = None,
# FR-5: Verify + re-sign
access_token_discovery_uri: Optional[str] = None,
token_introspection_endpoint: Optional[str] = None,
verify_issuer: Optional[str] = None,
verify_audience: 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_audience: Optional[str] = None,
channel_token_ttl: Optional[int] = None,
# FR-15: Incoming claim validation
required_claims: Optional[List[str]] = None,
optional_claims: Optional[List[str]] = None,
# FR-9: Debug headers
debug_headers: bool = False,
# FR-10: Configurable scopes
allowed_scopes: Optional[List[str]] = None,
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
# --- Signing 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 config ---
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: Verify + re-sign ---
self.access_token_discovery_uri: Optional[str] = access_token_discovery_uri
self.token_introspection_endpoint: Optional[str] = token_introspection_endpoint
self.verify_issuer: Optional[str] = verify_issuer
self.verify_audience: Optional[str] = verify_audience
# Cached OIDC discovery document (fetched lazily, TTL = 24 h)
self._oidc_discovery_doc: Optional[Dict[str, Any]] = None
self._oidc_discovery_fetched_at: float = 0.0
# --- FR-12: End-user identity mapping ---
# Default chain: try incoming JWT sub, fall back to litellm user_id
self.end_user_claim_sources: List[str] = end_user_claim_sources or [
"token:sub",
"litellm:user_id",
]
# --- FR-13: Claim operations ---
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: Two-token model ---
self.channel_token_audience: Optional[str] = channel_token_audience
self.channel_token_ttl: int = (
channel_token_ttl if channel_token_ttl is not None else self.ttl_seconds
)
# --- FR-15: Incoming claim validation ---
self.required_claims: List[str] = required_claims or []
self.optional_claims: List[str] = optional_claims or []
# --- FR-9: Debug headers ---
self.debug_headers: bool = debug_headers
# --- FR-10: Configurable scopes ---
self.allowed_scopes: Optional[List[str]] = allowed_scopes
# Register singleton for JWKS/OIDC discovery endpoints.
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 "
"verify=%s channel_token=%s debug=%s",
self.issuer,
self.audience,
self.ttl_seconds,
self._kid,
bool(self.access_token_discovery_uri),
bool(self.channel_token_audience),
self.debug_headers,
)
# ------------------------------------------------------------------
# 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).
1 hour for persistent keys; 5 minutes for auto-generated keys so MCP
servers re-fetch quickly after a proxy restart.
"""
return 3600 if self._persistent_key else 300
def get_jwks(self) -> Dict[str, Any]:
"""
Return the JWKS 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-5: Verify + re-sign helpers
# ------------------------------------------------------------------
# 24-hour TTL for the OIDC discovery doc — long enough to avoid hammering
# the IdP, short enough to pick up jwks_uri changes after key rotation.
_OIDC_DISCOVERY_TTL = 86400
async def _get_oidc_discovery(self) -> Dict[str, Any]:
"""Fetch and cache the OIDC discovery document with a 24-hour TTL.
Only caches when the doc contains a 'jwks_uri' so that a transient or
malformed response doesn't permanently disable JWT verification.
"""
now = time.time()
cache_expired = (now - self._oidc_discovery_fetched_at) >= self._OIDC_DISCOVERY_TTL
if (self._oidc_discovery_doc is None or cache_expired) and self.access_token_discovery_uri:
doc = await _fetch_oidc_discovery(self.access_token_discovery_uri)
if "jwks_uri" in doc:
self._oidc_discovery_doc = doc
self._oidc_discovery_fetched_at = now
else:
return doc
return self._oidc_discovery_doc or {}
async def _verify_incoming_jwt(self, raw_token: str) -> Dict[str, Any]:
"""
Verify an incoming Bearer JWT against the configured IdP's JWKS.
Returns the verified payload claims dict.
Raises jwt.PyJWTError (or subclass) if verification fails.
"""
discovery = await self._get_oidc_discovery()
jwks_uri = discovery.get("jwks_uri")
if not jwks_uri:
raise ValueError(
"MCPJWTSigner: access_token_discovery_uri discovery document "
f"at {self.access_token_discovery_uri!r} has no 'jwks_uri'."
)
jwks_keys = await _fetch_jwks(jwks_uri)
# Only read `kid` from the unverified header — never `alg`.
# Reading `alg` from an attacker-controlled header enables algorithm
# confusion attacks (e.g. alg:none, HS256 with the public key as secret).
# The algorithm is determined from the JWKS key entry instead.
unverified_header = jwt.get_unverified_header(raw_token)
kid = unverified_header.get("kid")
# Build a JWKS object and pick the matching key.
# PyJWT's PyJWKSet handles key-type parsing and kid matching correctly.
from jwt import PyJWKSet
try:
jwks_set = PyJWKSet.from_dict({"keys": jwks_keys})
except Exception as exc:
raise jwt.exceptions.PyJWKSetError( # type: ignore[attr-defined]
f"Failed to parse JWKS from {jwks_uri!r}: {exc}"
) from exc
signing_jwk = None
for jwk_obj in jwks_set.keys:
if not kid or jwk_obj.key_id == kid:
signing_jwk = jwk_obj
break
if signing_jwk is None:
raise jwt.exceptions.PyJWKSetError( # type: ignore[attr-defined]
f"No JWKS key matching kid={kid!r} at {jwks_uri!r}"
)
# Use the algorithm declared by the JWKS key entry, not the token header.
# PyJWT populates algorithm_name from the key's `alg` field; when absent
# it infers from the key type (RSAPublicKey → RS256).
alg = getattr(signing_jwk, "algorithm_name", None) or "RS256"
decode_options: Dict[str, Any] = {"verify_exp": True}
decode_kwargs: Dict[str, Any] = {
"algorithms": [alg],
"options": decode_options,
}
if self.verify_audience:
decode_kwargs["audience"] = self.verify_audience
else:
decode_options["verify_aud"] = False
if self.verify_issuer:
decode_kwargs["issuer"] = self.verify_issuer
payload: Dict[str, Any] = jwt.decode(
raw_token, signing_jwk.key, **decode_kwargs
)
return payload
async def _introspect_opaque_token(self, token: str) -> Dict[str, Any]:
"""
Perform RFC 7662 token introspection for opaque (non-JWT) tokens.
Returns the introspection response dict. Raises on HTTP error or
inactive token.
"""
if not self.token_introspection_endpoint:
raise ValueError(
"MCPJWTSigner: token_introspection_endpoint is required for "
"opaque token verification but is not configured."
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
resp = await client.post(
self.token_introspection_endpoint,
data={"token": token},
headers={"Accept": "application/json"},
)
resp.raise_for_status()
result: Dict[str, Any] = resp.json()
if not result.get("active", False):
raise jwt.exceptions.ExpiredSignatureError( # type: ignore[attr-defined]
"MCPJWTSigner: incoming token is inactive (introspection returned active=false)"
)
return result
# ------------------------------------------------------------------
# FR-15: Incoming claim validation
# ------------------------------------------------------------------
def _validate_required_claims(
self,
jwt_claims: Optional[Dict[str, Any]],
) -> None:
"""
Raise HTTP 403 if any required_claims are absent from the verified
incoming token claims.
"""
if not self.required_claims:
return
from fastapi import HTTPException
missing = [c for c in self.required_claims if not (jwt_claims or {}).get(c)]
if missing:
raise HTTPException(
status_code=403,
detail={
"error": (
f"MCPJWTSigner: incoming token is missing required claims: "
f"{missing}. Configure the IdP to include these claims."
)
},
)
# ------------------------------------------------------------------
# FR-12: End-user identity mapping
# ------------------------------------------------------------------
def _resolve_end_user_identity(
self,
user_api_key_dict: UserAPIKeyAuth,
jwt_claims: Optional[Dict[str, Any]],
) -> str:
"""
Resolve the outbound JWT 'sub' using the ordered end_user_claim_sources list.
Supported source prefixes:
token:<claim> — from verified incoming JWT / introspection claims
litellm:user_id — from UserAPIKeyAuth.user_id
litellm:email — from UserAPIKeyAuth.user_email
litellm:end_user_id — from UserAPIKeyAuth.end_user_id
litellm:team_id — from UserAPIKeyAuth.team_id
Falls back to a stable hash of the API token for service-account callers.
"""
for source in self.end_user_claim_sources:
value: Optional[str] = None
if source.startswith("token:"):
claim_name = source[len("token:"):]
raw = (jwt_claims or {}).get(claim_name)
value = str(raw) if raw else None
elif source == "litellm:user_id":
uid = getattr(user_api_key_dict, "user_id", None)
value = str(uid) if uid else None
elif source == "litellm:email":
email = getattr(user_api_key_dict, "user_email", None)
value = str(email) if email else None
elif source == "litellm:end_user_id":
eid = getattr(user_api_key_dict, "end_user_id", None)
value = str(eid) if eid else None
elif source == "litellm:team_id":
tid = getattr(user_api_key_dict, "team_id", None)
value = str(tid) if tid else None
else:
verbose_proxy_logger.warning(
"MCPJWTSigner: unknown end_user_claim_source %r — skipping", source
)
continue
if value:
return value
# Final fallback for service accounts with no user identity
token = getattr(user_api_key_dict, "token", None) or getattr(
user_api_key_dict, "api_key", None
)
if token:
return "apikey:" + hashlib.sha256(str(token).encode()).hexdigest()[:16]
return "litellm-proxy"
# ------------------------------------------------------------------
# FR-10: Scope building
# ------------------------------------------------------------------
def _build_scope(self, raw_tool_name: str) -> str:
"""
Build the JWT scope string.
When allowed_scopes is configured: join them verbatim.
Otherwise auto-generate minimal, least-privilege scopes:
- Tool call → mcp:tools/call mcp:tools/<name>:call
- No tool → mcp:tools/call mcp:tools/list
NOTE: tools/list is intentionally NOT granted on tool-call JWTs to
prevent callers from enumerating tools they didn't ask to use.
"""
if self.allowed_scopes is not None:
return " ".join(self.allowed_scopes)
tool_name = (
re.sub(r"[^a-zA-Z0-9_\-]", "_", raw_tool_name) if raw_tool_name else ""
)
if tool_name:
scopes = ["mcp:tools/call", f"mcp:tools/{tool_name}:call"]
else:
scopes = ["mcp:tools/call", "mcp:tools/list"]
return " ".join(scopes)
# ------------------------------------------------------------------
# FR-13: Claim operations
# ------------------------------------------------------------------
def _apply_claim_operations(self, claims: Dict[str, Any]) -> Dict[str, Any]:
"""Apply add_claims, set_claims, and remove_claims to the claim dict."""
# add_claims: insert only when key is absent
for k, v in self.add_claims.items():
if k not in claims:
claims[k] = v
# set_claims: always override (highest priority)
claims = {**claims, **self.set_claims}
# remove_claims: delete listed keys
for k in self.remove_claims:
claims.pop(k, None)
return claims
# ------------------------------------------------------------------
# FR-15: optional_claims passthrough
# ------------------------------------------------------------------
def _passthrough_optional_claims(
self,
claims: Dict[str, Any],
jwt_claims: Optional[Dict[str, Any]],
) -> Dict[str, Any]:
"""Forward optional_claims from verified incoming token into the outbound JWT."""
if not self.optional_claims or not jwt_claims:
return claims
for claim in self.optional_claims:
if claim in jwt_claims and claim not in claims:
claims[claim] = jwt_claims[claim]
return claims
# ------------------------------------------------------------------
# Core JWT builder
# ------------------------------------------------------------------
def _build_claims(
self,
user_api_key_dict: UserAPIKeyAuth,
data: dict,
jwt_claims: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""
Build JWT claims for the outbound MCP access token.
Args:
user_api_key_dict: LiteLLM auth context for the current request.
data: Pre-call hook data dict (contains mcp_tool_name etc.).
jwt_claims: Verified incoming IdP claims (FR-5), or LiteLLM-decoded
jwt_claims if available. None for pure API-key requests.
"""
now = int(time.time())
claims: Dict[str, Any] = {
"iss": self.issuer,
"aud": self.audience,
"iat": now,
"exp": now + self.ttl_seconds,
"nbf": now,
}
# sub — resolved via ordered claim sources (FR-12)
claims["sub"] = self._resolve_end_user_identity(user_api_key_dict, jwt_claims)
# email passthrough when available from LiteLLM context
user_email = getattr(user_api_key_dict, "user_email", None)
if user_email:
claims["email"] = user_email
# act — RFC 8693 delegation claim (team/org context)
team_id = getattr(user_api_key_dict, "team_id", None)
org_id = getattr(user_api_key_dict, "org_id", None)
act_sub = team_id or org_id or "litellm-proxy"
claims["act"] = {"sub": act_sub}
# end_user_id when set separately from user_id
end_user_id = getattr(user_api_key_dict, "end_user_id", None)
if end_user_id:
claims["end_user_id"] = end_user_id
# scope (FR-10)
raw_tool_name: str = data.get("mcp_tool_name", "")
claims["scope"] = self._build_scope(raw_tool_name)
# optional_claims passthrough (FR-15)
claims = self._passthrough_optional_claims(claims, jwt_claims)
# Claim operations — applied last so admin overrides take effect (FR-13)
claims = self._apply_claim_operations(claims)
return claims
def _build_channel_token_claims(
self,
base_claims: Dict[str, Any],
) -> Dict[str, Any]:
"""
Build claims for the channel token (FR-14 two-token model).
Inherits sub/act/scope from the access token but uses a separate
audience and TTL so the transport layer and resource layer receive
purpose-bound credentials.
"""
now = int(time.time())
return {
**base_claims,
"aud": self.channel_token_audience,
"iat": now,
"exp": now + self.channel_token_ttl,
"nbf": now,
}
# ------------------------------------------------------------------
# FR-9: Debug header
# ------------------------------------------------------------------
@staticmethod
def _build_debug_header(claims: Dict[str, Any], kid: str) -> str:
"""
Build the x-litellm-mcp-debug header value.
Format: v=1; kid=<kid>; sub=<sub>; iss=<iss>; exp=<exp>; scope=<scope>
Scope is truncated to 80 chars for header safety.
"""
sub = claims.get("sub", "")
iss = claims.get("iss", "")
exp = claims.get("exp", 0)
scope = claims.get("scope", "")
if len(scope) > 80:
scope = scope[:77] + "..."
return f"v=1; kid={kid}; sub={sub}; iss={iss}; exp={exp}; scope={scope}"
# ------------------------------------------------------------------
# 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]]:
"""
Verifies the incoming token (when configured), validates required claims,
then signs an outbound JWT and injects it as the Authorization header.
All non-MCP call types pass through unchanged.
"""
if call_type != "call_mcp_tool":
return data
# ------------------------------------------------------------------
# FR-5: Verify incoming token before re-signing
# ------------------------------------------------------------------
jwt_claims: Optional[Dict[str, Any]] = None
raw_token: Optional[str] = data.get("incoming_bearer_token")
if self.access_token_discovery_uri and raw_token:
# Three-dot pattern → JWT; otherwise opaque.
is_jwt = raw_token.count(".") == 2
try:
if is_jwt:
jwt_claims = await self._verify_incoming_jwt(raw_token)
elif self.token_introspection_endpoint:
jwt_claims = await self._introspect_opaque_token(raw_token)
else:
verbose_proxy_logger.warning(
"MCPJWTSigner: access_token_discovery_uri is set but the "
"incoming token appears to be opaque and no "
"token_introspection_endpoint is configured. "
"Proceeding without incoming token verification."
)
except Exception as exc:
verbose_proxy_logger.error(
"MCPJWTSigner: incoming token verification failed: %s", exc
)
from fastapi import HTTPException
raise HTTPException(
status_code=401,
detail={
"error": (
f"MCPJWTSigner: incoming token verification failed: {exc}"
)
},
)
elif not raw_token and self.access_token_discovery_uri:
verbose_proxy_logger.debug(
"MCPJWTSigner: access_token_discovery_uri configured but no Bearer "
"token found in request (API-key auth request — skipping verification)."
)
# Fall back to LiteLLM-decoded JWT claims (available when proxy uses JWT auth).
if jwt_claims is None:
jwt_claims = getattr(user_api_key_dict, "jwt_claims", None)
# ------------------------------------------------------------------
# FR-15: Validate required claims
# ------------------------------------------------------------------
self._validate_required_claims(jwt_claims)
# ------------------------------------------------------------------
# Build outbound access token
# ------------------------------------------------------------------
claims = self._build_claims(user_api_key_dict, data, jwt_claims)
signed_token = jwt.encode(
claims,
self._private_key,
algorithm=self.ALGORITHM,
headers={"kid": self._kid},
)
# Merge into existing extra_headers — a prior guardrail in the chain may
# have already injected tracing headers or correlation IDs.
existing_headers: Dict[str, str] = data.get("extra_headers") or {}
new_headers: Dict[str, str] = {
**existing_headers,
"Authorization": f"Bearer {signed_token}",
}
# ------------------------------------------------------------------
# FR-14: Two-token model — channel token
# ------------------------------------------------------------------
if self.channel_token_audience:
channel_claims = self._build_channel_token_claims(claims)
channel_token = jwt.encode(
channel_claims,
self._private_key,
algorithm=self.ALGORITHM,
headers={"kid": self._kid},
)
new_headers["x-mcp-channel-token"] = f"Bearer {channel_token}"
# ------------------------------------------------------------------
# FR-9: Debug header
# ------------------------------------------------------------------
if self.debug_headers:
new_headers["x-litellm-mcp-debug"] = self._build_debug_header(
claims, self._kid
)
data["extra_headers"] = new_headers
verbose_proxy_logger.debug(
"MCPJWTSigner: signed JWT sub=%s act=%s tool=%s exp=%d "
"verified=%s channel=%s",
claims.get("sub"),
claims.get("act", {}).get("sub"),
data.get("mcp_tool_name"),
claims["exp"],
jwt_claims is not None,
bool(self.channel_token_audience),
)
return data

View file

@ -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,10 @@ 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 Bearer token from the original HTTP request — allows guardrails
# (e.g. MCPJWTSigner) to independently verify the caller's identity
# before re-signing an outbound token (FR-5 verify+re-sign).
"incoming_bearer_token": kwargs.get("incoming_bearer_token"),
}
return synthetic_data
@ -824,17 +826,27 @@ 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"):
# Merge rather than replace — a prior guardrail in the chain may have
# already injected headers (e.g. tracing IDs). Later guardrails win on
# key collisions so that the most-specific guardrail (e.g. JWT signer)
# takes precedence over earlier ones.
existing = modified_kwargs.get("extra_headers") or {}
modified_kwargs["extra_headers"] = {**existing, **response_data["extra_headers"]}
return modified_kwargs
async def process_pre_call_hook_response(self, response, data, call_type):

View file

@ -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):

View file

@ -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"]

File diff suppressed because it is too large Load diff