fix(mcp): route REST tools list filtering through the shared toolset-aware primitive

This commit is contained in:
Tin Chi Lo 2026-07-16 17:45:37 -07:00
parent c5cfe284cb
commit b5d38b84e0
2 changed files with 84 additions and 14 deletions

View file

@ -515,20 +515,19 @@ if MCP_AVAILABLE:
# enforced even when no allowlist is set (matches the SSE/HTTP path).
tools = filter_tools_by_allowed_tools(tools, server)
# Filter tools based on user_api_key_auth.object_permission.mcp_tool_permissions
# This provides per-key/team/org control over which tools can be accessed
if (
user_api_key_auth
and user_api_key_auth.object_permission
and user_api_key_auth.object_permission.mcp_tool_permissions
):
# Dict keys may be server_ids OR names/aliases; normalize so lookup
# by concrete server_id resolves name-keyed restrictions too.
allowed_tools_for_server = global_mcp_server_manager.expand_tool_permissions(
user_api_key_auth.object_permission.mcp_tool_permissions
).get(server.server_id)
if allowed_tools_for_server is not None and len(allowed_tools_for_server) > 0:
# Filter tools to only include those in the allowed list
# Filter by the key's effective tool permissions through the same
# primitive the MCP protocol path uses (direct grants, toolset grants,
# and team/agent/org ceilings), so REST listing cannot drift from it
if user_api_key_auth:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
allowed_tools_for_server = await MCPRequestHandler.get_allowed_tools_for_server(
server_id=server.server_id,
user_api_key_auth=user_api_key_auth,
)
if allowed_tools_for_server is not None:
tools = [tool for tool in tools if _tool_name_matches(tool.name, allowed_tools_for_server)]
return _create_tool_response_objects(tools, server)

View file

@ -2654,3 +2654,74 @@ class TestToolResponseMcpInfoEnrichment:
"server_id": "server-uuid",
"alias": None,
}
class TestRestListToolsetFiltering:
@pytest.mark.asyncio
async def test_rest_list_filters_toolset_only_key_to_toolset_tools(self, monkeypatch):
"""A toolset-only key reaching a toolset server via REST list must see
only the toolset's tools; the raw catalog leaked every tool on the
server when the filter read object_permission directly instead of the
shared toolset-aware primitive"""
from unittest.mock import patch
from mcp.types import Tool as MCPTool
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._experimental.mcp_server.server import MCPServer
from litellm.types.mcp import MCPTransport
stub_server = MCPServer(
server_id="server-a",
name="stubtools",
transport=MCPTransport.http,
)
stub_server.alias = "stubtools"
stub_server.server_name = "stubtools"
stub_server.allowed_tools = None
stub_server.disallowed_tools = None
stub_server.mcp_info = {"server_name": "stubtools"}
upstream_tools = [
MCPTool(name="lookup_status", inputSchema={"type": "object"}),
MCPTool(name="delete_everything", inputSchema={"type": "object"}),
]
key_object_permission = MagicMock()
key_object_permission.mcp_servers = []
key_object_permission.mcp_access_groups = []
key_object_permission.mcp_tool_permissions = None
key_object_permission.mcp_toolsets = ["toolset-1"]
user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
mock_manager = MagicMock()
mock_manager.expand_tool_permissions = MagicMock(side_effect=lambda perms: perms or {})
mock_manager.resolve_toolset_tool_permissions = AsyncMock(
return_value={"server-a": ["lookup_status"]}
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"_get_tools_from_server",
AsyncMock(return_value=upstream_tools),
)
with (
patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_object_permission),
patch.object(MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=None)),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
mock_manager,
),
):
result = await rest_endpoints._get_tools_for_single_server(
server=stub_server,
server_auth_header=None,
raw_headers=None,
user_api_key_auth=user_auth,
)
assert [tool.name for tool in result] == ["lookup_status"]