From a33539d107cc8d4214292ebf7a3e16ff677b9456 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 15 Sep 2026 12:38:39 +0000 Subject: [PATCH] fix(key_management): count team unified access group MCP servers when validating key MCP grants Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../object_permission_utils.py | 33 +++++-- .../test_object_permission_utils.py | 91 +++++++++++++++++++ 2 files changed, 118 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index daab38d3662..a1f1ba6c672 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -547,17 +547,35 @@ async def _get_team_allowed_mcp_servers( """ Get the full set of MCP server IDs a team allows. - If team has no object_permission or no MCP config, returns empty set - (meaning only allow_all_keys servers are permitted). + Combines servers granted via the team's object_permission with servers + granted via the team's unified access groups (access_group_ids). If the + team grants neither, returns empty set (meaning only allow_all_keys + servers are permitted). """ if team_obj is None: return set() + from litellm.proxy.auth.auth_checks import ( + _get_mcp_server_ids_from_access_groups, # pyright: ignore[reportPrivateUsage] # same shared resolver the runtime MCP auth path calls + ) + + access_group_servers: Final = await _get_mcp_server_ids_from_access_groups( + access_group_ids=team_obj.access_group_ids or [], + prisma_client=prisma_client, + ) + resolved_access_group_servers: Final = await _resolve_mcp_server_identifiers_to_ids( + identifiers=set(access_group_servers), + prisma_client=prisma_client, + ) + unified_servers: Final = _flatten_resolved_mcp_server_ids(resolved_access_group_servers) | { + server for server in access_group_servers if not resolved_access_group_servers.get(server) + } + team_object_permission: Final = team_obj.object_permission if team_object_permission is None: - return set() + return unified_servers - return await _resolve_team_allowed_mcp_servers( + return unified_servers | await _resolve_team_allowed_mcp_servers( team_object_permission=team_object_permission, prisma_client=prisma_client, ) @@ -632,14 +650,17 @@ async def validate_key_mcp_servers_against_team( Rules: - If key is in a team: key's mcp_servers must be a subset of - (team's allowed servers + allow_all_keys servers) + (team's allowed servers + allow_all_keys servers), where the team's + allowed servers include servers granted via the team's unified + access groups - If key is NOT in a team and the caller is a proxy admin: any server or access group may be assigned. A proxy admin can already reach every MCP server, and runtime access is granted directly from the key's own object_permission, so the key is scoped to exactly what the admin selected - If key is NOT in a team and the caller is not a proxy admin: key's mcp_servers must only contain allow_all_keys servers - - If team has no MCP config: key can only use allow_all_keys servers + - If team has no MCP config (no object_permission and no unified + access groups): key can only use allow_all_keys servers Raises HTTPException(403) if validation fails. """ 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 078315c2bf8..d68a8592f30 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 @@ -217,10 +217,12 @@ def _make_team_obj( mcp_servers=None, mcp_access_groups=None, mcp_tool_permissions=None, + access_group_ids=None, ): """Create a mock team object with the given MCP permissions.""" mock_team = MagicMock() mock_team.team_id = team_id + mock_team.access_group_ids = access_group_ids or [] if ( mcp_servers is not None @@ -541,6 +543,95 @@ 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"), +) +@patch( + "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + return_value=set(), +) +@patch( + "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + new_callable=AsyncMock, + return_value=["server-1"], +) +async def test_validate_key_servers_granted_via_team_unified_access_group_pass( + mock_unified_access_groups, mock_allow_all +): + """A team whose only MCP grant comes from a unified access group still + allows keys in that team to request those servers.""" + team_obj = _make_team_obj(access_group_ids=["ag-1"]) + await validate_key_mcp_servers_against_team( + object_permission={"mcp_servers": ["server-1"]}, + team_obj=team_obj, + ) + mock_unified_access_groups.assert_awaited_once() + assert mock_unified_access_groups.await_args.kwargs["access_group_ids"] == ["ag-1"] + + +@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"), +) +@patch( + "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + return_value=set(), +) +@patch( + "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + new_callable=AsyncMock, + return_value=["server-1"], +) +async def test_validate_key_servers_outside_team_unified_access_group_rejected( + mock_unified_access_groups, mock_allow_all +): + """A server not granted by the team's unified access group is rejected, + and the error lists the access-group-granted servers as the team scope.""" + team_obj = _make_team_obj(access_group_ids=["ag-1"]) + with pytest.raises(HTTPException) as exc_info: + await validate_key_mcp_servers_against_team( + object_permission={"mcp_servers": ["server-2"]}, + team_obj=team_obj, + ) + assert exc_info.value.status_code == 403 + assert "server-2" in str(exc_info.value.detail) + assert "Team allows: ['server-1']" in str(exc_info.value.detail) + + +@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"), +) +@patch( + "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + return_value=set(), +) +@patch( + "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + new_callable=AsyncMock, + return_value=["server-2"], +) +@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_team_allowed_servers_union_object_permission_and_unified_access_group( + mock_access_groups, mock_unified_access_groups, mock_allow_all +): + """Team scope is the union of object_permission servers and unified + access group servers.""" + team_obj = _make_team_obj(mcp_servers=["server-1"], access_group_ids=["ag-1"]) + await validate_key_mcp_servers_against_team( + object_permission={"mcp_servers": ["server-1", "server-2"]}, + team_obj=team_obj, + ) + + @pytest.mark.asyncio @patch( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",