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>
This commit is contained in:
Devin AI 2026-09-15 12:38:39 +00:00
parent 9496f16f12
commit a33539d107
2 changed files with 118 additions and 6 deletions

View file

@ -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.
"""

View file

@ -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",