feat: add all-mcp-servers sentinel for team MCP allocation

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Krrish Dholakia 2026-07-01 01:35:37 +00:00
parent 2860dad514
commit 0d43292707
6 changed files with 167 additions and 2 deletions

View file

@ -3836,6 +3836,8 @@ class MCPServerManager:
if not identifiers:
return []
registry = self.get_registry()
if SpecialMCPServerNames.all_mcp_servers.value in identifiers:
return list(registry.keys())
expanded: Set[str] = set()
for identifier in identifiers:
if identifier in registry:

View file

@ -2885,6 +2885,7 @@ class SpecialModelNames(enum.Enum):
class SpecialMCPServerNames(enum.Enum):
no_mcp_servers = "no-mcp-servers"
all_mcp_servers = "all-mcp-servers"
class SpecialProxyStrings(enum.Enum):

View file

@ -266,8 +266,11 @@ def _rewrite_object_permission_mcp_servers(
normalized_servers: List[str] = []
for identifier in mcp_servers:
if identifier == SpecialMCPServerNames.no_mcp_servers.value:
normalized_servers.append(SpecialMCPServerNames.no_mcp_servers.value)
if identifier in (
SpecialMCPServerNames.no_mcp_servers.value,
SpecialMCPServerNames.all_mcp_servers.value,
):
normalized_servers.append(identifier)
continue
normalized_servers.extend(sorted(identifier_to_server_ids.get(identifier, [])))
object_permission["mcp_servers"] = _dedupe_preserving_order(normalized_servers)
@ -334,6 +337,13 @@ async def _resolve_team_allowed_mcp_servers(
)
direct_servers: List[str] = team_object_permission.mcp_servers or []
if SpecialMCPServerNames.all_mcp_servers.value in direct_servers:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
return global_mcp_server_manager.get_all_mcp_server_ids()
access_group_servers: List[str] = await MCPRequestHandler._get_mcp_servers_from_access_groups(
team_object_permission.mcp_access_groups or []
)
@ -400,6 +410,7 @@ def _extract_requested_mcp_server_ids(
if isinstance(mcp_servers, list):
server_ids.update(mcp_servers)
server_ids.discard(SpecialMCPServerNames.no_mcp_servers.value)
server_ids.discard(SpecialMCPServerNames.all_mcp_servers.value)
mcp_tool_permissions = object_permission.get("mcp_tool_permissions")
if isinstance(mcp_tool_permissions, dict):

View file

@ -315,6 +315,68 @@ class TestMCPRequestHandler:
assert result == [SpecialMCPServerNames.no_mcp_servers.value]
@pytest.mark.parametrize(
"team_servers,key_servers,expected_servers,scenario",
[
# Team has all-mcp-servers, key has no permissions -> inherit all
(
["server1", "server2", "server3"],
[],
["server1", "server2", "server3"],
"team_all_key_empty",
),
# Team has all-mcp-servers, key has specific servers -> key's servers
# (intersection of all and key = key)
(
["server1", "server2", "server3"],
["server1"],
["server1"],
"team_all_key_subset",
),
# Both have full sets -> all servers
(
["server1", "server2", "server3"],
["server1", "server2", "server3"],
["server1", "server2", "server3"],
"both_all",
),
],
)
async def test_all_mcp_servers_sentinel_on_team(
self, team_servers, key_servers, expected_servers, scenario
):
"""When _get_allowed_mcp_servers_for_team returns all IDs (because the
team has all-mcp-servers), the key/team intersection works correctly"""
auth = UserAPIKeyAuth(
api_key="test-key", user_id="test-user", team_id="test-team"
)
with (
patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_key",
new_callable=AsyncMock,
return_value=key_servers,
),
patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_team",
new_callable=AsyncMock,
return_value=team_servers,
),
patch.object(
MCPRequestHandler,
"_get_key_access_group_mcp_server_extras",
new_callable=AsyncMock,
return_value=[],
),
patch(
"litellm.proxy.proxy_server.general_settings",
{},
),
):
result = await MCPRequestHandler.get_allowed_mcp_servers(auth)
assert sorted(result) == sorted(expected_servers)
async def test_permission_inheritance_edge_cases(self):
"""Test edge cases in permission inheritance"""

View file

@ -3928,6 +3928,26 @@ class TestMCPServerManagerExpandPermissionList:
assert manager.expand_permission_list(["uuid-1", "a"]) == ["uuid-1"]
def test_all_mcp_servers_sentinel_returns_all_registry_keys(self):
"""The all-mcp-servers sentinel expands to every registered server ID"""
manager = MCPServerManager()
manager.config_mcp_servers["cfg-1"] = self._make_server("cfg-1", server_name="a")
manager.registry["reg-1"] = self._make_server("reg-1", server_name="b")
manager.registry["reg-2"] = self._make_server("reg-2", server_name="c")
result = manager.expand_permission_list(["all-mcp-servers"])
assert sorted(result) == ["cfg-1", "reg-1", "reg-2"]
def test_all_mcp_servers_sentinel_ignores_other_entries(self):
"""When all-mcp-servers is present alongside other IDs, the sentinel
dominates and returns all server IDs (the extra entries are a subset)"""
manager = MCPServerManager()
manager.config_mcp_servers["cfg-1"] = self._make_server("cfg-1", server_name="a")
manager.registry["reg-1"] = self._make_server("reg-1", server_name="b")
result = manager.expand_permission_list(["all-mcp-servers", "cfg-1"])
assert sorted(result) == ["cfg-1", "reg-1"]
def test_simulates_cross_region_portability(self):
"""
Same permission entry "a" resolves to different concrete IDs per region —

View file

@ -118,12 +118,23 @@ def test_extract_requested_mcp_server_ids_excludes_no_mcp_servers_sentinel():
assert _extract_requested_mcp_server_ids(obj_perm) == {"server-1"}
def test_extract_requested_mcp_server_ids_excludes_all_mcp_servers_sentinel():
obj_perm = {"mcp_servers": ["all-mcp-servers", "server-1"]}
assert _extract_requested_mcp_server_ids(obj_perm) == {"server-1"}
def test_rewrite_object_permission_mcp_servers_preserves_sentinel():
obj_perm = {"mcp_servers": ["no-mcp-servers", "alias-1"]}
_rewrite_object_permission_mcp_servers(obj_perm, {"alias-1": {"server-1"}})
assert obj_perm["mcp_servers"] == ["no-mcp-servers", "server-1"]
def test_rewrite_object_permission_mcp_servers_preserves_all_sentinel():
obj_perm = {"mcp_servers": ["all-mcp-servers", "alias-1"]}
_rewrite_object_permission_mcp_servers(obj_perm, {"alias-1": {"server-1"}})
assert obj_perm["mcp_servers"] == ["all-mcp-servers", "server-1"]
@pytest.mark.asyncio
async def test_validate_no_mcp_servers_sentinel_passes_and_preserved():
"""A key scoped to no-mcp-servers passes team validation untouched, keeping the
@ -204,6 +215,7 @@ def _make_mock_mcp_manager(*existing_ids: str, servers=None):
server_objs.setdefault(server_id, _make_mock_mcp_server(server_id))
mock_mgr.get_registry.return_value = server_objs
mock_mgr.get_mcp_server_by_id.side_effect = lambda sid: server_objs.get(sid)
mock_mgr.get_all_mcp_server_ids.return_value = set(server_objs.keys())
return mock_mgr
@ -483,6 +495,63 @@ async def test_validate_team_no_mcp_config_blocks_all(
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
@patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
new=_make_mock_mcp_manager("server-1", "server-2", "server-3"),
)
@patch(
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
return_value=set(),
)
@patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_key_against_team_with_all_mcp_servers_sentinel(
mock_access_groups, mock_allow_all
):
"""When a team has the all-mcp-servers sentinel, any key server should
pass validation because the team grants access to everything"""
team_obj = _make_team_obj(mcp_servers=["all-mcp-servers"])
obj_perm = {"mcp_servers": ["server-1", "server-2"]}
result = await validate_key_mcp_servers_against_team(
object_permission=obj_perm,
team_obj=team_obj,
)
assert result is not None
assert "server-1" in result["mcp_servers"]
assert "server-2" in result["mcp_servers"]
@pytest.mark.asyncio
@patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
)
@patch(
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
return_value=set(),
)
@patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
)
async def test_resolve_team_allowed_mcp_servers_with_all_sentinel(
mock_access_groups, mock_allow_all, mock_mgr
):
"""_resolve_team_allowed_mcp_servers returns all registry IDs when the
team has the all-mcp-servers sentinel"""
mock_mgr.get_all_mcp_server_ids.return_value = {"s1", "s2", "s3"}
team_perm = MagicMock(spec=LiteLLM_ObjectPermissionTable)
team_perm.mcp_servers = ["all-mcp-servers"]
team_perm.mcp_access_groups = []
team_perm.mcp_tool_permissions = {}
result = await _resolve_team_allowed_mcp_servers(team_perm)
assert result == {"s1", "s2", "s3"}
@pytest.mark.asyncio
@patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",