fix(mcp): resolve team.access_group_ids → MCP servers (#28997)

* fix(mcp): resolve team.access_group_ids → MCP servers

A virtual key whose team has an MCP-granting access group attached
via /v1/access_group now sees that server through /v1/mcp/server (and
can call tools on it) instead of getting an empty list. The runtime
already resolves the key's unified access_group_ids; this adds the
symmetric resolution on the team side, mirroring the model-side
pattern in can_team_access_model — the group being on the team is
itself the gate, so no assigned_team_ids re-check is needed.

Resolves #27657

* chore(mcp): address greptile review on team access-group resolver

Forward already-imported prisma_client / user_api_key_cache /
proxy_logging_obj to _get_mcp_server_ids_from_access_groups so it
skips its lazy re-import path. Update test docstring + assertions to
reflect that the resolver is invoked with [] (and short-circuits
without DB access) rather than skipped entirely.
This commit is contained in:
ryan-crabbe-berri
2026-05-27 12:36:50 -07:00
committed by GitHub
parent 73e9071311
commit ea015332d8
4 changed files with 365 additions and 129 deletions
@@ -992,42 +992,78 @@ class MCPRequestHandler:
"""
Get allowed MCP servers for a team.
Note: object_permission is automatically loaded by get_team_object() in main auth flow.
Unions two sources:
- Legacy team.object_permission (mcp_servers, mcp_access_groups,
mcp_tool_permissions).
- Unified team.access_group_ids → access_group.access_mcp_server_ids.
Mirrors the model-side pattern in can_team_access_model — the group
is already attached to the team, so the team relationship is itself
the gate (no assigned_team_ids check needed here).
"""
try:
# Get team object permission (already loaded in main auth flow)
object_permissions = await MCPRequestHandler._get_team_object_permission(
user_api_key_auth
)
if object_permissions is None:
return []
# Permission entries may be server_ids OR names/aliases — expand to ids.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy.auth.auth_checks import (
_get_mcp_server_ids_from_access_groups,
get_team_object,
)
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if (
user_api_key_auth is None
or not user_api_key_auth.team_id
or prisma_client is None
):
return []
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
team_id=user_api_key_auth.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if team_obj is None:
return []
team_access_group_servers = await _get_mcp_server_ids_from_access_groups(
access_group_ids=team_obj.access_group_ids or [],
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
object_permissions = team_obj.object_permission
if object_permissions is None:
return list(set(team_access_group_servers))
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(
object_permissions.mcp_servers or []
)
# Get MCP servers from access groups
access_group_servers = (
legacy_access_group_servers = (
await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
)
)
# servers referenced in tool permissions should also be accessible
tool_perm_servers = list(
global_mcp_server_manager.expand_tool_permissions(
object_permissions.mcp_tool_permissions
).keys()
)
# Combine all lists
all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers
all_servers = (
direct_mcp_servers
+ legacy_access_group_servers
+ tool_perm_servers
+ team_access_group_servers
)
return list(set(all_servers))
except Exception as e:
verbose_logger.warning(
@@ -2444,13 +2444,14 @@ async def test_get_team_object_permission_with_core_auth_auto_loading():
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_for_team_uses_helper():
"""
Test that _get_allowed_mcp_servers_for_team properly uses _get_team_object_permission
helper which handles both loaded and unloaded object_permission cases.
Test that _get_allowed_mcp_servers_for_team resolves both legacy
object_permission fields (mcp_servers, mcp_access_groups) and the unified
team.access_group_ids → access_mcp_server_ids path.
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@@ -2464,53 +2465,51 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper():
transport=MCPTransport.http,
)
try:
# Create mock object permission with servers and access groups
mock_object_permission = LiteLLM_ObjectPermissionTable(
object_permission_id="perm-789",
mcp_servers=["direct-server1", "direct-server2"],
mcp_access_groups=["dev-group"],
vector_stores=[],
)
mock_team = LiteLLM_TeamTable(
team_id="team-789",
access_group_ids=[],
object_permission_id="perm-789",
)
mock_team.object_permission = mock_object_permission
# Create mock user auth
mock_user_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
team_id="team-789",
)
# Mock the helper methods
with patch.object(
MCPRequestHandler, "_get_team_object_permission"
) as mock_get_team_perm:
with patch.object(
MCPRequestHandler, "_get_mcp_servers_from_access_groups"
) as mock_get_access_group_servers:
# Configure mocks
mock_get_team_perm.return_value = mock_object_permission
mock_get_access_group_servers.return_value = [
"group-server1",
"group-server2",
]
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=mock_team,
),
patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=["group-server1", "group-server2"],
) as mock_get_access_group_servers,
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
mock_user_auth
)
# Call the method
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
mock_user_auth
)
assert set(result) == {
"direct-server1",
"direct-server2",
"group-server1",
"group-server2",
}
# Assert the result contains both direct and access group servers
assert set(result) == {
"direct-server1",
"direct-server2",
"group-server1",
"group-server2",
}
# Verify _get_team_object_permission was called (the helper we fixed)
mock_get_team_perm.assert_called_once_with(mock_user_auth)
# Verify access groups were resolved
mock_get_access_group_servers.assert_called_once_with(["dev-group"])
mock_get_access_group_servers.assert_called_once_with(["dev-group"])
finally:
for sid in ("direct-server1", "direct-server2"):
global_mcp_server_manager.registry.pop(sid, None)
@@ -2520,32 +2519,36 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper():
async def test_get_allowed_mcp_servers_for_team_with_no_object_permission():
"""
Test that _get_allowed_mcp_servers_for_team returns empty list when
team has no object_permission.
the team has no object_permission and no access_group_ids.
"""
# Create mock user auth
from litellm.proxy._types import LiteLLM_TeamTable
mock_team = LiteLLM_TeamTable(
team_id="team-no-perm",
access_group_ids=[],
object_permission_id=None,
)
mock_user_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
team_id="team-no-perm",
)
# Mock the helper to return None (no object permission)
with patch.object(
MCPRequestHandler, "_get_team_object_permission"
) as mock_get_team_perm:
mock_get_team_perm.return_value = None
# Call the method
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=mock_team,
),
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
mock_user_auth
)
# Assert empty list is returned
assert result == []
# Verify the helper was called
mock_get_team_perm.assert_called_once_with(mock_user_auth)
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_for_team_without_user_auth_returns_empty():
@@ -3456,3 +3459,185 @@ async def test_get_allowed_mcp_servers_no_union_when_no_authorized_extras():
# key ∩ team = {} (no overlap), extras = [] → final = []
result = await MCPRequestHandler.get_allowed_mcp_servers(auth)
assert result == []
# ---------------------------------------------------------------------------
# Issue #27657: team unified access_group_ids resolve to MCP servers
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_team_access_group_ids_resolve_to_mcp_servers():
"""A virtual key with empty access_group_ids inherits MCP servers from
its team's access_group_ids (mirror of the model-side resolution).
Reproduction of https://github.com/BerriAI/litellm/issues/27657:
the runtime used to ignore team.access_group_ids when computing the
MCP scope, so virtual keys saw empty server lists even when their
team had an MCP-granting access group attached.
"""
from litellm.proxy._types import LiteLLM_TeamTable
mock_team = LiteLLM_TeamTable(
team_id="team-a",
access_group_ids=["mcp-premium"],
object_permission_id=None,
)
auth = UserAPIKeyAuth(
token="test-token-hash",
api_key="sk-test",
team_id="team-a",
access_group_ids=[],
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=mock_team,
),
patch(
"litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups",
new_callable=AsyncMock,
return_value=["srv-stripe"],
) as mock_resolver,
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth)
assert result == ["srv-stripe"]
mock_resolver.assert_called_once()
assert mock_resolver.call_args.kwargs["access_group_ids"] == ["mcp-premium"]
@pytest.mark.asyncio
async def test_team_access_group_ids_union_with_object_permission():
"""When both legacy object_permission and unified team.access_group_ids
grant MCP servers, the final list is their union."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
for sid in ("srv-direct",):
global_mcp_server_manager.registry[sid] = MCPServer(
server_id=sid,
name=sid,
server_name=sid,
url=f"https://{sid}.example.com",
transport=MCPTransport.http,
)
try:
mock_object_permission = LiteLLM_ObjectPermissionTable(
object_permission_id="perm-1",
mcp_servers=["srv-direct"],
mcp_access_groups=[],
vector_stores=[],
)
mock_team = LiteLLM_TeamTable(
team_id="team-a",
access_group_ids=["mcp-premium"],
object_permission_id="perm-1",
)
mock_team.object_permission = mock_object_permission
auth = UserAPIKeyAuth(
token="test-token-hash",
api_key="sk-test",
team_id="team-a",
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=mock_team,
),
patch(
"litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups",
new_callable=AsyncMock,
return_value=["srv-stripe"],
),
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth)
assert set(result) == {"srv-direct", "srv-stripe"}
finally:
global_mcp_server_manager.registry.pop("srv-direct", None)
@pytest.mark.asyncio
async def test_team_access_group_ids_empty_returns_no_extras():
"""Empty team.access_group_ids → resolver called with [], short-circuits
without DB access, no extras added."""
from litellm.proxy._types import LiteLLM_TeamTable
mock_team = LiteLLM_TeamTable(
team_id="team-a",
access_group_ids=[],
object_permission_id=None,
)
auth = UserAPIKeyAuth(
token="test-token-hash",
api_key="sk-test",
team_id="team-a",
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=mock_team,
),
patch(
"litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups",
new_callable=AsyncMock,
return_value=[],
) as mock_resolver,
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth)
assert result == []
mock_resolver.assert_called_once()
assert mock_resolver.call_args.kwargs["access_group_ids"] == []
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_includes_team_access_group_extras_end_to_end():
"""End-to-end: virtual key has nothing of its own, team has an MCP
access group → key sees the granted server through get_allowed_mcp_servers."""
auth = UserAPIKeyAuth(
token="test-token",
api_key="sk-test",
team_id="team-a",
access_group_ids=[],
)
with (
patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_key",
new_callable=AsyncMock,
return_value=[],
),
patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_team",
new_callable=AsyncMock,
return_value=["srv-stripe"],
),
patch.object(
MCPRequestHandler,
"_get_key_access_group_mcp_server_extras",
new_callable=AsyncMock,
return_value=[],
),
):
result = await MCPRequestHandler.get_allowed_mcp_servers(auth)
assert result == ["srv-stripe"]
@@ -462,6 +462,9 @@ async def test_e2e_jwt_team_mcp_key_intersection(monkeypatch):
monkeypatch.setattr(
"litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object
)
monkeypatch.setattr(
"litellm.proxy.auth.auth_checks.get_team_object", mock_get_team_object
)
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_ids_jwt_field="groups")
@@ -495,28 +498,25 @@ async def test_e2e_jwt_team_mcp_key_intersection(monkeypatch):
object_permission=key_object_permission, # Key has its own permissions
)
# Mock the helper methods to return our test data
with patch.object(
MCPRequestHandler, "_get_team_object_permission"
) as mock_team_perm:
mock_team_perm.return_value = team_object_permission
with (
patch.object(
MCPRequestHandler,
"_get_key_object_permission",
return_value=key_object_permission,
),
patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
),
):
allowed_servers = await MCPRequestHandler.get_allowed_mcp_servers(
user_api_key_auth
)
with patch.object(
MCPRequestHandler, "_get_key_object_permission"
) as mock_key_perm:
mock_key_perm.return_value = key_object_permission
with patch.object(
MCPRequestHandler, "_get_mcp_servers_from_access_groups"
) as mock_access_groups:
mock_access_groups.return_value = []
allowed_servers = await MCPRequestHandler.get_allowed_mcp_servers(
user_api_key_auth
)
# Should be intersection: only server-2 is in both
expected = ["server-2"]
assert sorted(allowed_servers) == sorted(
expected
), f"Expected intersection {expected}, got {allowed_servers}"
# Should be intersection: only server-2 is in both
expected = ["server-2"]
assert sorted(allowed_servers) == sorted(
expected
), f"Expected intersection {expected}, got {allowed_servers}"
@@ -41,37 +41,44 @@ async def test_simple_jwt_mcp_permissions_enforced():
object_permission_id="perm-123",
mcp_servers=team_mcp_servers,
)
team_obj = LiteLLM_TeamTable(
team_id="my-team",
access_group_ids=[],
object_permission_id="perm-123",
)
team_obj.object_permission = team_object_permission
# 3. Mock the team permission lookup
with patch.object(
MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock
) as mock_team_perm:
mock_team_perm.return_value = team_object_permission
# 3. Mock the team object lookup (object_permission attached) and prisma_client
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=team_obj,
) as mock_get_team,
patch.object(
MCPRequestHandler,
"_get_key_object_permission",
new_callable=AsyncMock,
return_value=None,
),
patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
),
):
# 4. Call get_allowed_mcp_servers - this is what MCP routes use
allowed = await MCPRequestHandler.get_allowed_mcp_servers(user_auth)
# Mock key permissions (empty - user has no key-level MCP permissions)
with patch.object(
MCPRequestHandler, "_get_key_object_permission", new_callable=AsyncMock
) as mock_key_perm:
mock_key_perm.return_value = None
# 5. Verify only team's MCP servers are returned
assert sorted(allowed) == sorted(
team_mcp_servers
), f"Expected {team_mcp_servers}, got {allowed}"
# Mock access groups (empty)
with patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
) as mock_access_groups:
mock_access_groups.return_value = []
# 4. Call get_allowed_mcp_servers - this is what MCP routes use
allowed = await MCPRequestHandler.get_allowed_mcp_servers(user_auth)
# 5. Verify only team's MCP servers are returned
assert sorted(allowed) == sorted(
team_mcp_servers
), f"Expected {team_mcp_servers}, got {allowed}"
# Verify team permission was looked up
mock_team_perm.assert_called_once_with(user_auth)
# Verify team was looked up
mock_get_team.assert_called()
@pytest.mark.asyncio
@@ -120,25 +127,33 @@ async def test_simple_jwt_team_id_required_for_mcp_permissions():
object_permission_id="perm-1",
mcp_servers=team_mcp_servers,
)
team_obj = LiteLLM_TeamTable(
team_id="team-abc",
access_group_ids=[],
object_permission_id="perm-1",
)
team_obj.object_permission = team_perm
with patch.object(
MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock
) as mock_perm:
mock_perm.return_value = team_perm
with patch.object(
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=team_obj,
) as mock_get_team,
patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
) as mock_groups:
mock_groups.return_value = []
return_value=[],
),
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
user_with_team
)
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
user_with_team
)
assert sorted(result) == sorted(team_mcp_servers)
mock_perm.assert_called_once() # Permission WAS checked
assert sorted(result) == sorted(team_mcp_servers)
mock_get_team.assert_called() # Team WAS looked up
# Case 2: team_id is None -> team permissions NOT checked
user_without_team = UserAPIKeyAuth(