mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
cff3e0b75e
commit
93db9e0dad
5 changed files with 203 additions and 22 deletions
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue