mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
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:
parent
9496f16f12
commit
a33539d107
2 changed files with 118 additions and 6 deletions
|
|
@ -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.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue