fix(mcp): JWT on tools/list, REST server_id resolution, tool_server_mismatch

Sign outbound MCP JWTs for list_mcp_tools and inject headers on the tools/list
path. Resolve server_id on /mcp-rest/tools/call and return 403 tool_server_mismatch
when the tool does not belong to the requested server. Default missing arguments to {}.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-05-19 13:34:05 +05:30
parent cff3e0b75e
commit 93db9e0dad
No known key found for this signature in database
5 changed files with 203 additions and 22 deletions

View file

@ -1406,6 +1406,7 @@ class MCPServerManager:
extra_headers: Optional[Dict[str, str]] = None,
add_prefix: bool = True,
raw_headers: Optional[Dict[str, str]] = None,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[MCPTool]:
"""
Helper method to get tools from a single MCP server with prefixed names.
@ -1432,6 +1433,19 @@ class MCPServerManager:
extra_headers = {}
extra_headers.update(server.static_headers)
# MCPJWTSigner: inject signed JWT for tools/list (list path skips pre_call_hook)
if user_api_key_auth is not None and not server.spec_path:
from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import (
inject_mcp_jwt_headers_for_upstream,
)
extra_headers = await inject_mcp_jwt_headers_for_upstream(
user_api_key_dict=user_api_key_auth,
extra_headers=extra_headers,
raw_headers=raw_headers,
for_list_tools=True,
)
stdio_env = self._build_stdio_env(server, raw_headers)
client = await self._create_mcp_client(
@ -2822,9 +2836,19 @@ class MCPServerManager:
"""
start_time = datetime.datetime.now()
# Get the MCP server
prefixed_tool_name = add_server_prefix_to_name(name, server_name)
mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name)
# Resolve server (REST may pass server_name + prefixed or unprefixed tool name)
mcp_server: Optional[MCPServer] = None
for candidate in self.get_registry().values():
if normalize_server_name(candidate.name) == normalize_server_name(
server_name
):
mcp_server = candidate
break
if mcp_server is None:
prefixed_tool_name = add_server_prefix_to_name(name, server_name)
mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name)
if mcp_server is None:
mcp_server = self._get_mcp_server_from_tool_name(name)
if mcp_server is None:
raise ValueError(f"Tool {name} not found")

View file

@ -1,6 +1,17 @@
import importlib
from datetime import datetime
from typing import Any, Awaitable, Callable, Dict, List, Literal, Optional, Set, Union
from typing import (
Any,
Awaitable,
Callable,
Dict,
List,
Literal,
Optional,
Set,
Tuple,
Union,
)
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
@ -231,11 +242,32 @@ if MCP_AVAILABLE:
)
return mcp_auth_header, mcp_server_auth_headers, raw_headers
def _resolve_mcp_server_id_for_rest(
server_id: str,
allowed_server_ids: Union[Set[str], List[str]],
client_ip: Optional[str] = None,
) -> str:
"""
Map REST ``server_id`` (UUID, server_name, or alias) to canonical server_id.
tools/list already did this; tools/call must match so clients can pass
server names like ``order_status_mcp`` instead of only UUIDs.
"""
allowed = set(allowed_server_ids)
if server_id in allowed:
return server_id
by_name = global_mcp_server_manager.get_mcp_server_by_name(
server_id, client_ip=client_ip
)
if by_name is not None and by_name.server_id in allowed:
return by_name.server_id
return server_id
async def _resolve_allowed_mcp_servers_with_ip_filter(
request: Request,
user_api_key_dict: UserAPIKeyAuth,
server_id: str,
) -> List[MCPServer]:
) -> Tuple[List[MCPServer], str]:
"""
Resolve allowed MCP servers for a tool call with IP filtering.
@ -245,10 +277,10 @@ if MCP_AVAILABLE:
server_id: The server ID to validate access for
Returns:
List of allowed MCPServer objects
Tuple of (allowed MCPServer objects, canonical server_id)
Raises:
HTTPException: If the server_id is not allowed
HTTPException: If the server_id is not allowed or not found
"""
# Get all auth contexts
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
@ -268,8 +300,24 @@ if MCP_AVAILABLE:
)
)
# Check if the specified server_id is allowed
if server_id not in allowed_server_ids_set:
canonical_server_id = _resolve_mcp_server_id_for_rest(
server_id, allowed_server_ids_set, _rest_client_ip
)
if canonical_server_id not in allowed_server_ids_set:
_server = global_mcp_server_manager.get_mcp_server_by_id(
server_id
) or global_mcp_server_manager.get_mcp_server_by_name(
server_id, client_ip=_rest_client_ip
)
if _server is None:
raise HTTPException(
status_code=404,
detail={
"error": "server_not_found",
"message": f"MCP server '{server_id}' was not found",
},
)
raise HTTPException(
status_code=403,
detail={
@ -285,7 +333,7 @@ if MCP_AVAILABLE:
if server is not None:
allowed_mcp_servers.append(server)
return allowed_mcp_servers
return allowed_mcp_servers, canonical_server_id
async def _get_tools_for_single_server(
server,
@ -301,6 +349,7 @@ if MCP_AVAILABLE:
extra_headers=extra_headers,
add_prefix=False,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
# Filter tools based on allowed_tools configuration
@ -786,14 +835,18 @@ if MCP_AVAILABLE:
data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"]
# Resolve allowed MCP servers with IP filtering
allowed_mcp_servers = await _resolve_allowed_mcp_servers_with_ip_filter(
(
allowed_mcp_servers,
canonical_server_id,
) = await _resolve_allowed_mcp_servers_with_ip_filter(
request, user_api_key_dict, server_id
)
# Look up per-user OAuth headers for this server (mirrors list_tool_rest_api).
user_oauth_extra_headers: Optional[Dict[str, str]] = None
target_server = next(
(s for s in allowed_mcp_servers if s.server_id == server_id), None
(s for s in allowed_mcp_servers if s.server_id == canonical_server_id),
None,
)
if target_server is not None:
user_oauth_extra_headers = await _get_user_oauth_extra_headers(
@ -812,6 +865,7 @@ if MCP_AVAILABLE:
oauth2_headers=user_oauth_extra_headers or data.get("oauth2_headers"),
raw_headers=data.get("raw_headers"),
litellm_logging_obj=data.get("litellm_logging_obj"),
requested_server_id=canonical_server_id,
)
return result
except BlockedPiiEntityError as e:

View file

@ -1368,6 +1368,7 @@ if MCP_AVAILABLE:
extra_headers=extra_headers,
add_prefix=True, # Always add server prefix
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
filtered_tools = filter_tools_by_allowed_tools(tools, server)
@ -2074,6 +2075,7 @@ if MCP_AVAILABLE:
"""
# Track resolved MCP server for both permission checks and dispatch
mcp_server: Optional[MCPServer] = None
requested_server_id: Optional[str] = kwargs.get("requested_server_id")
# If the client called with a display-name override (e.g. "Get Pet"),
# translate it back to the original prefixed name before any routing.
@ -2082,6 +2084,13 @@ if MCP_AVAILABLE:
# Remove prefix from tool name for logging and processing
original_tool_name, server_name = split_server_prefix_from_name(name)
requested_server: Optional[MCPServer] = None
if requested_server_id:
requested_server = next(
(s for s in allowed_mcp_servers if s.server_id == requested_server_id),
None,
)
# Resolve the actual MCP server up-front so the permission check uses
# the canonical server.name even when the tool name is prefixed with a
# short ID (LITELLM_USE_SHORT_MCP_TOOL_PREFIX) that doesn't match the
@ -2090,6 +2099,26 @@ if MCP_AVAILABLE:
if mcp_server is not None:
server_name = mcp_server.name
# REST /mcp-rest/tools/call passes server_id — tool must belong to that server
if requested_server is not None:
if (
mcp_server is not None
and mcp_server.server_id != requested_server.server_id
):
raise HTTPException(
status_code=403,
detail={
"error": "tool_server_mismatch",
"message": (
f"Tool '{name}' belongs to MCP server '{mcp_server.name}' "
f"but request specified server_id for '{requested_server.name}'."
),
},
)
if mcp_server is None:
mcp_server = requested_server
server_name = requested_server.name
# Only enforce server-level permissions when we can resolve a server
if server_name:
if not MCPRequestHandler.is_tool_allowed(

View file

@ -92,6 +92,8 @@ from litellm.types.utils import CallTypesLiteral
# Module-level singleton for the JWKS discovery endpoint to access.
_mcp_jwt_signer_instance: Optional["MCPJWTSigner"] = None
_MCP_JWT_CALL_TYPES = frozenset({"call_mcp_tool", "list_mcp_tools"})
# Simple in-memory JWKS cache: keyed by JWKS URI → (keys_list, fetched_at).
_jwks_cache: Dict[str, tuple] = {}
_JWKS_CACHE_TTL = 3600 # 1 hour
@ -779,16 +781,20 @@ class MCPJWTSigner(CustomGuardrail):
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.
Signs outbound MCP tool calls and tools/list requests.
"""
if call_type != "call_mcp_tool":
if call_type not in _MCP_JWT_CALL_TYPES:
return data
hook_data = dict(data)
if call_type == "list_mcp_tools":
hook_data["mcp_tool_name"] = ""
# ------------------------------------------------------------------
# FR-5: Verify incoming token before re-signing
# ------------------------------------------------------------------
jwt_claims: Optional[Dict[str, Any]] = None
raw_token: Optional[str] = data.get("incoming_bearer_token")
raw_token: Optional[str] = hook_data.get("incoming_bearer_token")
if self.access_token_discovery_uri and raw_token:
# Three-dot pattern → JWT; otherwise opaque.
@ -837,7 +843,7 @@ class MCPJWTSigner(CustomGuardrail):
# ------------------------------------------------------------------
# Build outbound access token
# ------------------------------------------------------------------
claims = self._build_claims(user_api_key_dict, data, jwt_claims)
claims = self._build_claims(user_api_key_dict, hook_data, jwt_claims)
signed_token = jwt.encode(
claims,
@ -848,7 +854,7 @@ class MCPJWTSigner(CustomGuardrail):
# 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 {}
existing_headers: Dict[str, str] = hook_data.get("extra_headers") or {}
new_headers: Dict[str, str] = {
**existing_headers,
"Authorization": f"Bearer {signed_token}",
@ -875,17 +881,61 @@ class MCPJWTSigner(CustomGuardrail):
claims, self._kid
)
data["extra_headers"] = new_headers
hook_data["extra_headers"] = new_headers
verbose_proxy_logger.debug(
"MCPJWTSigner: signed JWT sub=%s act=%s tool=%s exp=%d "
"verified=%s channel=%s",
"verified=%s channel=%s call_type=%s",
claims.get("sub"),
claims.get("act", {}).get("sub"),
data.get("mcp_tool_name"),
hook_data.get("mcp_tool_name"),
claims["exp"],
jwt_claims is not None,
bool(self.channel_token_audience),
call_type,
)
return data
return hook_data
async def inject_mcp_jwt_headers_for_upstream(
user_api_key_dict: Optional[UserAPIKeyAuth],
extra_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
*,
for_list_tools: bool = False,
mcp_tool_name: str = "",
) -> Dict[str, str]:
"""
Sign outbound MCP headers when MCPJWTSigner is configured.
Used by tools/list paths that do not go through proxy pre_call_hook.
"""
merged = dict(extra_headers or {})
signer = get_mcp_jwt_signer()
if signer is None or user_api_key_dict is None:
return merged
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 ") :]
hook_data: Dict[str, Any] = {
"mcp_tool_name": "" if for_list_tools else mcp_tool_name,
"incoming_bearer_token": incoming_bearer_token,
"extra_headers": merged,
}
call_type: CallTypesLiteral = (
"list_mcp_tools" if for_list_tools else "call_mcp_tool"
)
result = await signer.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data=hook_data,
call_type=call_type,
)
if isinstance(result, dict) and result.get("extra_headers"):
merged.update(result["extra_headers"])
return merged

View file

@ -338,7 +338,7 @@ async def test_hook_skips_non_mcp_call_types():
user_dict = _make_user_api_key_dict()
data = {"messages": [{"role": "user", "content": "hello"}]}
for call_type in ("completion", "acompletion", "embedding", "list_mcp_tools"):
for call_type in ("completion", "acompletion", "embedding"):
original_data = {**data}
result = await signer.async_pre_call_hook(
user_api_key_dict=user_dict,
@ -351,6 +351,30 @@ async def test_hook_skips_non_mcp_call_types():
), f"extra_headers should not be set for {call_type}"
@pytest.mark.asyncio
async def test_hook_signs_list_mcp_tools():
"""async_pre_call_hook() signs JWT for list_mcp_tools with list scope."""
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": "should_be_cleared"}
result = await signer.async_pre_call_hook(
user_api_key_dict=user_dict,
cache=MagicMock(),
data=data,
call_type="list_mcp_tools",
)
assert isinstance(result, dict)
assert "extra_headers" in result
assert result["extra_headers"]["Authorization"].startswith("Bearer ")
token = result["extra_headers"]["Authorization"].removeprefix("Bearer ")
decoded = _decode_unverified(token)
assert "mcp:tools/list" in decoded["scope"]
@pytest.mark.asyncio
async def test_signed_token_is_verifiable():
"""The JWT injected by the hook can be verified against the JWKS public key."""