mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
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:
commit
ea48ded1b1
6 changed files with 326 additions and 57 deletions
|
|
@ -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)}")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue