diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b6760e58852..101d1ca04ae 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index a6ef7de07ae..9cbd137ac1a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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): diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index fe96d9c260a..a21e585a2b9 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -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): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index b3b0e8adcf6..110c7b9e63a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 5dec580c771..91dc5e9b78e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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 — diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index 26c8c774812..8d543d19bbb 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -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",