Merge pull request #33612 from BerriAI/litellm_lit4448_toolset_call_grants

fix(mcp): expand toolset grants in shared permission primitives so tools/call honors them
This commit is contained in:
tin-berri 2026-07-17 10:16:07 -07:00 committed by GitHub
commit ea48ded1b1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 326 additions and 57 deletions

View file

@ -1220,11 +1220,29 @@ class MCPRequestHandler:
global_mcp_server_manager,
)
key_tools = (
key_direct_tools = (
global_mcp_server_manager.expand_tool_permissions(key_obj_perm.mcp_tool_permissions).get(server_id)
if key_obj_perm
else None
)
# Tools granted through the key's toolsets restrict this server exactly
# as direct tool permissions do; union with any direct grants so the
# tool-level check sees the key's full effective tool scope
key_toolset_ids = (key_obj_perm.mcp_toolsets or []) if key_obj_perm else []
key_toolset_tools = (
(await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=key_toolset_ids)).get(
server_id
)
if key_toolset_ids
else None
)
key_tools = (
list(set(key_direct_tools or []) | set(key_toolset_tools or []))
if key_direct_tools is not None or key_toolset_tools is not None
else None
)
team_tools = (
global_mcp_server_manager.expand_tool_permissions(team_obj_perm.mcp_tool_permissions).get(server_id)
if team_obj_perm
@ -1430,8 +1448,18 @@ class MCPRequestHandler:
global_mcp_server_manager.expand_tool_permissions(key_object_permission.mcp_tool_permissions).keys()
)
# servers referenced by the key's toolset grants are part of the key's
# scope on every path (list, call, REST), subject to the same team/org
# ceilings as any other key-level grant
toolset_ids = key_object_permission.mcp_toolsets or []
toolset_servers = (
list((await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=toolset_ids)).keys())
if toolset_ids
else []
)
# Combine all lists
all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers
all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers + toolset_servers
return list(set(all_servers))
except Exception as e:
verbose_logger.warning(f"Failed to get allowed MCP servers for key: {str(e)}")

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

@ -2218,43 +2218,6 @@ if MCP_AVAILABLE:
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
return [t for t in tools if strip_known_server_prefix(t.name, server) in allowed_tool_names]
async def _merge_toolset_permissions(
user_api_key_auth: Optional[UserAPIKeyAuth],
) -> Optional[UserAPIKeyAuth]:
"""
Resolve mcp_toolsets on the key's object_permission into tool-level permissions
and merge them (union) into object_permission.mcp_tool_permissions.
Returns the (possibly mutated copy of) user_api_key_auth.
"""
if user_api_key_auth is None:
return None
op = user_api_key_auth.object_permission
if op is None:
return user_api_key_auth
toolset_ids = getattr(op, "mcp_toolsets", None) or []
if not toolset_ids:
return user_api_key_auth
toolset_perms = await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=toolset_ids)
if not toolset_perms:
return user_api_key_auth
# Merge toolset_perms into existing mcp_tool_permissions (union)
existing = dict(op.mcp_tool_permissions or {})
for server_id, tool_names in toolset_perms.items():
existing_tools = existing.get(server_id, [])
merged = list(set(existing_tools) | set(tool_names))
existing[server_id] = merged
# Build updated object_permission with merged tool permissions and server IDs.
# Union the toolset's server IDs into mcp_servers so downstream server-level
# filtering doesn't silently drop servers that the toolset references but that
# aren't already in the key's explicit mcp_servers list.
merged_servers = list(set(op.mcp_servers or []) | set(existing.keys()))
updated_op = op.model_copy(update={"mcp_servers": merged_servers, "mcp_tool_permissions": existing})
return user_api_key_auth.model_copy(update={"object_permission": updated_op})
async def _list_mcp_tools(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
@ -2282,10 +2245,6 @@ if MCP_AVAILABLE:
if not MCP_AVAILABLE:
return []
# Resolve toolset permissions and merge into the key's object_permission
# so that the existing filter_tools_by_key_team_permissions logic picks them up.
user_api_key_auth = await _merge_toolset_permissions(user_api_key_auth)
# Get tools from managed MCP servers with error handling
managed_tools = []
try:

View file

@ -301,6 +301,187 @@ class TestMCPRequestHandler:
assert result == [SpecialMCPServerNames.no_mcp_servers.value]
def _toolset_only_object_permission(self, toolset_ids):
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_ids
return key_object_permission
def _mock_manager_with_toolsets(self, toolset_perms):
mock_manager = MagicMock()
mock_manager.expand_permission_list = MagicMock(side_effect=lambda servers: servers)
mock_manager.expand_tool_permissions = MagicMock(side_effect=lambda perms: perms or {})
mock_manager.resolve_toolset_tool_permissions = AsyncMock(return_value=toolset_perms)
return mock_manager
async def test_get_allowed_mcp_servers_for_key_includes_toolset_servers(self):
"""A key granted only mcp_toolsets must reach the toolset's servers on
every path (list, call, REST); regression for the list-ok/call-403 bug"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
key_object_permission = self._toolset_only_object_permission(["toolset-1"])
mock_manager = self._mock_manager_with_toolsets({"server-a": ["lookup_status"]})
with (
patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_object_permission),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
mock_manager,
),
patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])),
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
assert result == ["server-a"]
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(toolset_ids=["toolset-1"])
async def test_get_allowed_mcp_servers_for_key_skips_toolset_resolution_when_none_granted(self):
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
key_object_permission = self._toolset_only_object_permission([])
key_object_permission.mcp_servers = ["server-direct"]
mock_manager = self._mock_manager_with_toolsets({})
with (
patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_object_permission),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
mock_manager,
),
patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])),
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
assert result == ["server-direct"]
mock_manager.resolve_toolset_tool_permissions.assert_not_awaited()
async def test_get_allowed_mcp_servers_toolset_only_key_end_to_end_inheritance(self):
"""The full get_allowed_mcp_servers flow (key/team inheritance, no team
restriction) surfaces toolset-granted servers for a toolset-only key"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
key_object_permission = self._toolset_only_object_permission(["toolset-1"])
mock_manager = self._mock_manager_with_toolsets({"server-a": ["lookup_status"]})
with (
patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_object_permission),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
mock_manager,
),
patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])),
patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team", AsyncMock(return_value=[])),
patch.object(MCPRequestHandler, "_get_key_access_group_mcp_server_extras", AsyncMock(return_value=[])),
):
result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
assert result == ["server-a"]
async def test_toolset_servers_stay_capped_by_team_ceiling(self):
"""Toolset grants expand the KEY's scope, which the team ceiling still
intersects; a toolset must never grant a server the team does not allow.
Pins that toolset expansion lives in the intersected key scope, not the
additive access-group path"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user", team_id="test-team")
key_object_permission = self._toolset_only_object_permission(["toolset-1"])
mock_manager = self._mock_manager_with_toolsets(
{"server-in-team": ["lookup_status"], "server-outside-team": ["other_tool"]}
)
with (
patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_object_permission),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
mock_manager,
),
patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])),
patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_team",
AsyncMock(return_value=["server-in-team", "server-unrelated"]),
),
patch.object(MCPRequestHandler, "_get_key_access_group_mcp_server_extras", AsyncMock(return_value=[])),
):
result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
assert result == ["server-in-team"]
async def test_get_allowed_tools_for_server_unions_toolset_and_direct_tools(self):
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
key_object_permission = self._toolset_only_object_permission(["toolset-1"])
key_object_permission.mcp_tool_permissions = {"server-a": ["direct_tool"]}
mock_manager = self._mock_manager_with_toolsets({"server-a": ["lookup_status"]})
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 MCPRequestHandler.get_allowed_tools_for_server(
server_id="server-a",
user_api_key_auth=user_api_key_auth,
)
assert result is not None
assert set(result) == {"direct_tool", "lookup_status"}
async def test_get_allowed_tools_for_server_toolset_only_key_restricts_to_toolset_tools(self):
"""A toolset grant must RESTRICT the server's tools, not fall through to
the allow-all default; otherwise merging servers alone would over-grant
every tool on a toolset-referenced server"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
key_object_permission = self._toolset_only_object_permission(["toolset-1"])
mock_manager = self._mock_manager_with_toolsets({"server-a": ["lookup_status"]})
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,
),
):
allowed = await MCPRequestHandler.get_allowed_tools_for_server(
server_id="server-a",
user_api_key_auth=user_api_key_auth,
)
is_granted_tool_allowed = await MCPRequestHandler.is_tool_allowed_for_server(
tool_name="lookup_status",
server_id="server-a",
user_api_key_auth=user_api_key_auth,
)
is_other_tool_allowed = await MCPRequestHandler.is_tool_allowed_for_server(
tool_name="delete_everything",
server_id="server-a",
user_api_key_auth=user_api_key_auth,
)
assert allowed == ["lookup_status"]
assert is_granted_tool_allowed is True
assert is_other_tool_allowed is False
async def test_get_allowed_tools_for_server_without_restrictions_stays_allow_all(self):
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
key_object_permission = self._toolset_only_object_permission([])
mock_manager = self._mock_manager_with_toolsets({})
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 MCPRequestHandler.get_allowed_tools_for_server(
server_id="server-a",
user_api_key_auth=user_api_key_auth,
)
assert result is None
async def test_permission_inheritance_edge_cases(self):
"""Test edge cases in permission inheritance"""

View file

@ -8496,3 +8496,34 @@ def test_build_mcp_server_table_carries_null_oauth2_flow():
table = manager._build_mcp_server_table(server)
assert table.oauth2_flow is None
@pytest.mark.asyncio
async def test_resolve_toolset_tool_permissions_single_db_fetch_across_checks():
"""The server-level and tool-level permission primitives each resolve the
key's toolsets during one request; the shared cache must dedupe the DB
fetch so the request costs a single toolset query however many checks run"""
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
manager = MCPServerManager()
toolset = MagicMock()
toolset.tools = [{"server_id": "server-a", "tool_name": "lookup_status"}]
list_toolsets_mock = AsyncMock(return_value=[toolset])
with (
patch(
"litellm.proxy._experimental.mcp_server.toolset_db.list_mcp_toolsets",
list_toolsets_mock,
),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
):
first = await manager.resolve_toolset_tool_permissions(toolset_ids=["ts-1"])
second = await manager.resolve_toolset_tool_permissions(toolset_ids=["ts-1"])
assert first == {"server-a": ["lookup_status"]}
assert second == first
list_toolsets_mock.assert_awaited_once()

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