Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
19 commits
Select commit Hold shift + click to select a range
93db9e0
fix(mcp): JWT on tools/list, REST server_id resolution, tool_server_m…
Sameerlite May 19, 2026
5499662
fix(mcp): restrict list JWTs to mcp:tools/list and default REST argum…
cursoragent May 19, 2026
4c9f578
fix(mcp): validate tool/server in call_tool; skip JWT signer when not…
cursoragent May 19, 2026
16c6474
fix(mcp): align tests and mypy with user_api_key_auth on tools/list
Sameerlite May 19, 2026
c10ff56
fix(test): accept user_api_key_auth in get_tools_from_mcp_servers mock
Sameerlite May 19, 2026
002c3ad
fix(mcp): fail fast for unknown tools when server mapping exists
Sameerlite May 19, 2026
e7c8e55
fix mypy
Sameerlite May 20, 2026
bfe3f10
Fix mypy
Sameerlite May 20, 2026
97d17c3
Merge branch 'litellm_internal_staging' into litellm_mcp-rest-jwt-fixes
claude May 20, 2026
d2438b5
fix(mcp): preserve tools/call scope on missing tool name; pass user_a…
cursoragent May 20, 2026
4653773
fix(mcp): match alias/server_name in _resolve_mcp_server_for_tool_call
cursoragent May 20, 2026
ed10815
fix(mcp): reuse proxy_logging DualCache in inject_mcp_jwt_headers_for…
claude May 20, 2026
5b98ec0
fix(mcp): return 403 ip_filtering for IP-restricted servers in tools/…
cursoragent May 20, 2026
1015d8f
fix(test): accept user_api_key_auth kwarg in list_tools mocks
claude May 20, 2026
8da32a5
fix(mcp): skip JWT injection when per-user mcp_auth_header is set
claude May 20, 2026
d429e84
fix(mcp): skip JWT injection when extra_headers already has Authoriza…
cursoragent May 20, 2026
c0cbf6a
test(mcp): cover JWT signer + tool-call resolution branches
claude May 20, 2026
bc2e0d2
fix(mcp): retry tool-server lookup with prefixed name in REST mismatc…
cursoragent May 20, 2026
80ec93d
fix(mcp): always reject unknown tools in server-name fallback
May 20, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
209 changes: 153 additions & 56 deletions litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -1226,6 +1226,7 @@ async def _fetch_server_tools(server_id: str) -> List[MCPTool]:
tools = await self._get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
user_api_key_auth=user_api_key_auth,
)
return tools
except Exception as e:
Expand Down Expand Up @@ -1406,6 +1407,7 @@ async def _get_tools_from_server(
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,
Comment thread
cursor[bot] marked this conversation as resolved.
) -> List[MCPTool]:
"""
Helper method to get tools from a single MCP server with prefixed names.
Expand All @@ -1432,6 +1434,46 @@ async def _get_tools_from_server(
extra_headers = {}
extra_headers.update(server.static_headers)

# MCPJWTSigner: inject signed JWT for tools/list (list path skips pre_call_hook).
# Skip entirely when the signer is not configured (avoid an unnecessary
# dict copy on every list call), when the server has its own static
# Authorization header, when a per-user mcp_auth_header has already
# been resolved, or when the caller already supplied an Authorization
# entry in extra_headers (e.g. a per-user OAuth token resolved
# upstream) — admin-configured static auth and per-user OAuth must
# take precedence so the signer doesn't silently overwrite e.g. an
# upstream API key or a user's OAuth token (MCPClient._get_auth_headers
# applies extra_headers after writing Authorization from auth_value, so
# an injected JWT would otherwise clobber the per-user token).
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 (
get_mcp_jwt_signer,
inject_mcp_jwt_headers_for_upstream,
)

static_headers = server.static_headers or {}
has_static_authorization = any(
isinstance(k, str) and k.lower() == "authorization"
for k in static_headers.keys()
)
Comment thread
Sameerlite marked this conversation as resolved.
has_extra_authorization = bool(extra_headers) and any(
isinstance(k, str) and k.lower() == "authorization"
for k in (extra_headers or {}).keys()
)

if (
get_mcp_jwt_signer() is not None
and not has_static_authorization
and not mcp_auth_header
and not has_extra_authorization
):
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,
)
Comment thread
cursor[bot] marked this conversation as resolved.

stdio_env = self._build_stdio_env(server, raw_headers)

client = await self._create_mcp_client(
Expand Down Expand Up @@ -2791,6 +2833,112 @@ async def _call_tool_via_client(client, params):

return cast(CallToolResult, result)

def _resolve_mcp_server_for_tool_call(
self,
server_name: str,
name: str,
) -> MCPServer:
"""Resolve MCP server for call_tool (prefixed name, registry, fallback)."""
prefixed_tool_name = add_server_prefix_to_name(name, server_name)
mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name)
resolved_by_server_name_only = False
normalized_server_name = normalize_server_name(server_name)

def _candidate_matches_server_name(candidate: MCPServer) -> bool:
for identifier in (
candidate.alias,
candidate.server_name,
candidate.name,
):
if identifier and normalize_server_name(identifier) == (
normalized_server_name
):
return True
return False

if mcp_server is None:
for candidate in self.get_registry().values():
if _candidate_matches_server_name(candidate):
mcp_server = candidate
resolved_by_server_name_only = True
break
Comment thread
cursor[bot] marked this conversation as resolved.
if mcp_server is None:
fallback = self._get_mcp_server_from_tool_name(name)
if fallback is not None and (
not server_name or _candidate_matches_server_name(fallback)
):
mcp_server = fallback
if mcp_server is None:
raise ValueError(f"Tool {name} not found")

if resolved_by_server_name_only:
tool_known = (
name in self.tool_name_to_mcp_server_name_mapping
or prefixed_tool_name in self.tool_name_to_mcp_server_name_mapping
)
if not tool_known:
raise ValueError(f"Tool {name} not found")

return mcp_server

async def _resolve_oauth2_headers_for_tool_call(
self,
mcp_server: MCPServer,
oauth2_headers: Optional[Dict[str, str]],
user_api_key_auth: Optional[UserAPIKeyAuth],
) -> Optional[Dict[str, str]]:
"""Look up per-user OAuth headers when the client did not supply a token."""
if (
not mcp_server.needs_user_oauth_token
or oauth2_headers
or user_api_key_auth is None
):
return oauth2_headers

user_id = getattr(user_api_key_auth, "user_id", None)
if not user_id:
return oauth2_headers

try:
from litellm.proxy._experimental.mcp_server.server import ( # noqa: PLC0415
_get_user_oauth_extra_headers_from_db,
)

stored_headers = await _get_user_oauth_extra_headers_from_db(
server=mcp_server,
user_api_key_auth=user_api_key_auth,
)
if stored_headers:
return stored_headers
except Exception as _lookup_exc:
verbose_logger.debug(
"call_tool: per-user token lookup failed for " "user=%s server=%s: %s",
user_id,
mcp_server.server_id,
_lookup_exc,
)
return oauth2_headers

async def _gather_openapi_tool_tasks(
self,
tasks: List[Any],
proxy_logging_obj: Optional[ProxyLogging],
) -> CallToolResult:
"""Await OpenAPI tool tasks and return the tool call result."""
try:
mcp_responses = await asyncio.gather(*tasks)
result_index = 1 if proxy_logging_obj else 0
return cast(CallToolResult, mcp_responses[result_index])
except (
BlockedPiiEntityError,
GuardrailRaisedException,
HTTPException,
) as e:
verbose_logger.error(
f"Guardrail blocked MCP tool call during result check: {str(e)}"
)
raise e

async def call_tool(
self,
server_name: str,
Expand Down Expand Up @@ -2821,12 +2969,7 @@ async def call_tool(
CallToolResult from the MCP server
"""
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)
if mcp_server is None:
raise ValueError(f"Tool {name} not found")
mcp_server = self._resolve_mcp_server_for_tool_call(server_name, name)

#########################################################
# Pre MCP Tool Call Hook
Expand Down Expand Up @@ -2860,36 +3003,9 @@ async def call_tool(
)
tasks.append(during_hook_task)

# For per-user OAuth servers: if the client didn't supply a token in
# oauth2_headers, look up the stored token from Redis / DB. This is the
# call_tool equivalent of _get_user_oauth_extra_headers_from_db used in
# list_tools.
if (
mcp_server.needs_user_oauth_token
and not oauth2_headers
and user_api_key_auth is not None
):
user_id = getattr(user_api_key_auth, "user_id", None)
if user_id:
try:
from litellm.proxy._experimental.mcp_server.server import ( # noqa: PLC0415
_get_user_oauth_extra_headers_from_db,
)

stored_headers = await _get_user_oauth_extra_headers_from_db(
server=mcp_server,
user_api_key_auth=user_api_key_auth,
)
if stored_headers:
oauth2_headers = stored_headers
except Exception as _lookup_exc:
verbose_logger.debug(
"call_tool: per-user token lookup failed for "
"user=%s server=%s: %s",
user_id,
mcp_server.server_id,
_lookup_exc,
)
oauth2_headers = await self._resolve_oauth2_headers_for_tool_call(
mcp_server, oauth2_headers, user_api_key_auth
)

# For OpenAPI servers, call the tool handler directly instead of via MCP client
if mcp_server.spec_path:
Expand Down Expand Up @@ -2925,26 +3041,7 @@ async def call_tool(
hook_extra_headers=hook_result.get("extra_headers"),
)

# For OpenAPI tools, await outside the client context
try:
mcp_responses = await asyncio.gather(*tasks)

# If proxy_logging_obj is None, the tool call result is at index 0
# If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task)
result_index = 1 if proxy_logging_obj else 0
result = mcp_responses[result_index]

return cast(CallToolResult, result)
except (
BlockedPiiEntityError,
GuardrailRaisedException,
HTTPException,
) as e:
# Re-raise guardrail exceptions to properly fail the MCP call
verbose_logger.error(
f"Guardrail blocked MCP tool call during result check: {str(e)}"
)
raise e
return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj)

#########################################################
# End of Methods that call the upstream MCP servers
Expand Down
Loading
Loading