mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
2860dad514
commit
0d43292707
6 changed files with 167 additions and 2 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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 —
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue