mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge pull request #26338 from BerriAI/litellm_feat-mcp-server-alias-permissions
feat(mcp): resolve team/key MCP permissions by name or alias
This commit is contained in:
commit
9f60b751e1
5 changed files with 485 additions and 108 deletions
|
|
@ -516,15 +516,26 @@ class MCPRequestHandler:
|
|||
user_api_key_auth
|
||||
)
|
||||
|
||||
# Extract tool permissions for this server
|
||||
# Extract tool permissions for this server. Dict keys may be
|
||||
# server_ids OR names/aliases; normalize to server_id-keyed form
|
||||
# before lookup so a name-based key does not silently drop its
|
||||
# tool restrictions when server_id is the resolved uuid.
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
key_tools = (
|
||||
key_obj_perm.mcp_tool_permissions.get(server_id)
|
||||
if key_obj_perm and key_obj_perm.mcp_tool_permissions
|
||||
global_mcp_server_manager.expand_tool_permissions(
|
||||
key_obj_perm.mcp_tool_permissions
|
||||
).get(server_id)
|
||||
if key_obj_perm
|
||||
else None
|
||||
)
|
||||
team_tools = (
|
||||
team_obj_perm.mcp_tool_permissions.get(server_id)
|
||||
if team_obj_perm and team_obj_perm.mcp_tool_permissions
|
||||
global_mcp_server_manager.expand_tool_permissions(
|
||||
team_obj_perm.mcp_tool_permissions
|
||||
).get(server_id)
|
||||
if team_obj_perm
|
||||
else None
|
||||
)
|
||||
|
||||
|
|
@ -643,8 +654,14 @@ class MCPRequestHandler:
|
|||
if key_object_permission is None:
|
||||
return []
|
||||
|
||||
# Get direct MCP servers
|
||||
direct_mcp_servers = key_object_permission.mcp_servers or []
|
||||
# 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,
|
||||
)
|
||||
|
||||
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(
|
||||
key_object_permission.mcp_servers or []
|
||||
)
|
||||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers = (
|
||||
|
|
@ -655,7 +672,9 @@ class MCPRequestHandler:
|
|||
|
||||
# servers referenced in tool permissions should also be accessible
|
||||
tool_perm_servers = list(
|
||||
(key_object_permission.mcp_tool_permissions or {}).keys()
|
||||
global_mcp_server_manager.expand_tool_permissions(
|
||||
key_object_permission.mcp_tool_permissions
|
||||
).keys()
|
||||
)
|
||||
|
||||
# Combine all lists
|
||||
|
|
@ -685,8 +704,14 @@ class MCPRequestHandler:
|
|||
if object_permissions is None:
|
||||
return []
|
||||
|
||||
# Get direct MCP servers
|
||||
direct_mcp_servers = object_permissions.mcp_servers or []
|
||||
# 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,
|
||||
)
|
||||
|
||||
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(
|
||||
object_permissions.mcp_servers or []
|
||||
)
|
||||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers = (
|
||||
|
|
@ -697,7 +722,9 @@ class MCPRequestHandler:
|
|||
|
||||
# servers referenced in tool permissions should also be accessible
|
||||
tool_perm_servers = list(
|
||||
(object_permissions.mcp_tool_permissions or {}).keys()
|
||||
global_mcp_server_manager.expand_tool_permissions(
|
||||
object_permissions.mcp_tool_permissions
|
||||
).keys()
|
||||
)
|
||||
|
||||
# Combine all lists
|
||||
|
|
@ -746,8 +773,14 @@ class MCPRequestHandler:
|
|||
if end_user_obj is None or end_user_obj.object_permission is None:
|
||||
return []
|
||||
|
||||
# Get direct MCP servers
|
||||
direct_mcp_servers = end_user_obj.object_permission.mcp_servers or []
|
||||
# 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,
|
||||
)
|
||||
|
||||
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(
|
||||
end_user_obj.object_permission.mcp_servers or []
|
||||
)
|
||||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers = (
|
||||
|
|
@ -758,7 +791,9 @@ class MCPRequestHandler:
|
|||
|
||||
# servers referenced in tool permissions should also be accessible
|
||||
tool_perm_servers = list(
|
||||
(end_user_obj.object_permission.mcp_tool_permissions or {}).keys()
|
||||
global_mcp_server_manager.expand_tool_permissions(
|
||||
end_user_obj.object_permission.mcp_tool_permissions
|
||||
).keys()
|
||||
)
|
||||
|
||||
# Combine all lists
|
||||
|
|
@ -836,12 +871,21 @@ class MCPRequestHandler:
|
|||
if isinstance(mcp_access_groups, str):
|
||||
mcp_access_groups = []
|
||||
|
||||
# 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,
|
||||
)
|
||||
|
||||
expanded_direct_servers = global_mcp_server_manager.expand_permission_list(
|
||||
list(direct_mcp_servers)
|
||||
)
|
||||
|
||||
access_group_servers = (
|
||||
await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
mcp_access_groups
|
||||
)
|
||||
)
|
||||
all_servers = list(direct_mcp_servers) + access_group_servers
|
||||
all_servers = expanded_direct_servers + access_group_servers
|
||||
return list(set(all_servers))
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
|
|
@ -878,12 +922,16 @@ class MCPRequestHandler:
|
|||
return None
|
||||
|
||||
mcp_tool_permissions = getattr(obj_perm, "mcp_tool_permissions", None)
|
||||
if not mcp_tool_permissions:
|
||||
if not mcp_tool_permissions or not isinstance(mcp_tool_permissions, dict):
|
||||
return None
|
||||
if isinstance(mcp_tool_permissions, dict):
|
||||
tools = mcp_tool_permissions.get(server_id)
|
||||
else:
|
||||
tools = None
|
||||
# Dict keys may be server_ids OR names/aliases; normalize before lookup.
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
tools = global_mcp_server_manager.expand_tool_permissions(
|
||||
mcp_tool_permissions
|
||||
).get(server_id)
|
||||
return list(tools) if tools else None
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
|
|
|
|||
|
|
@ -2765,6 +2765,72 @@ class MCPServerManager:
|
|||
servers.append(server)
|
||||
return servers
|
||||
|
||||
def expand_permission_list(self, identifiers: List[str]) -> List[str]:
|
||||
"""
|
||||
Expand a permission list of server_ids/names/aliases into concrete
|
||||
server_ids against the current region's config + DB registry union.
|
||||
|
||||
Entries that match a server_id pass through unchanged. Entries that
|
||||
match an alias/server_name/name are replaced with every matching
|
||||
server_id (duplicate names grant access to all matches). Entries
|
||||
that resolve to nothing pass through as-is and a debug log is
|
||||
emitted so admins can diagnose stale/typo permission entries — the
|
||||
downstream access-check denies them when compared against the
|
||||
concrete request server_id.
|
||||
"""
|
||||
if not identifiers:
|
||||
return []
|
||||
registry = self.get_registry()
|
||||
expanded: Set[str] = set()
|
||||
for identifier in identifiers:
|
||||
if identifier in registry:
|
||||
expanded.add(identifier)
|
||||
continue
|
||||
matches: List[str] = [
|
||||
server_id
|
||||
for server_id, server in registry.items()
|
||||
if server.alias == identifier
|
||||
or server.server_name == identifier
|
||||
or server.name == identifier
|
||||
]
|
||||
if matches:
|
||||
expanded.update(matches)
|
||||
else:
|
||||
# %r quotes and escapes control chars so an admin-controlled
|
||||
# identifier with newlines cannot forge log lines.
|
||||
verbose_logger.debug(
|
||||
"MCP permission entry %r does not resolve to any known "
|
||||
"server (config + DB union). Passing through — the "
|
||||
"downstream access check will deny it if it's stale.",
|
||||
identifier,
|
||||
)
|
||||
expanded.add(identifier)
|
||||
return list(expanded)
|
||||
|
||||
def expand_tool_permissions(
|
||||
self,
|
||||
tool_permissions: Optional[Dict[str, List[str]]],
|
||||
) -> Dict[str, List[str]]:
|
||||
"""
|
||||
Rewrite an ``mcp_tool_permissions`` dict keyed by id/name/alias so
|
||||
every key is a concrete server_id where possible. Tool lists from
|
||||
keys that point at the same server are unioned, matching the
|
||||
"duplicate names grant access to all matches" semantics of
|
||||
``expand_permission_list``.
|
||||
|
||||
Required so name-based keys don't silently drop their tool
|
||||
restrictions when the lookup uses the resolved server_id. Unresolved
|
||||
keys pass through via ``expand_permission_list`` so stale id-keyed
|
||||
restrictions still apply when the same string is used for lookup.
|
||||
"""
|
||||
if not tool_permissions:
|
||||
return {}
|
||||
result: Dict[str, List[str]] = {}
|
||||
for key, tools in tool_permissions.items():
|
||||
for server_id in self.expand_permission_list([key]):
|
||||
result.setdefault(server_id, []).extend(tools or [])
|
||||
return result
|
||||
|
||||
def get_mcp_server_by_name(
|
||||
self, server_name: str, client_ip: Optional[str] = None
|
||||
) -> Optional[MCPServer]:
|
||||
|
|
|
|||
|
|
@ -315,10 +315,12 @@ if MCP_AVAILABLE:
|
|||
and user_api_key_auth.object_permission
|
||||
and user_api_key_auth.object_permission.mcp_tool_permissions
|
||||
):
|
||||
# Dict keys may be server_ids OR names/aliases; normalize so lookup
|
||||
# by concrete server_id resolves name-keyed restrictions too.
|
||||
allowed_tools_for_server = (
|
||||
user_api_key_auth.object_permission.mcp_tool_permissions.get(
|
||||
server.server_id
|
||||
)
|
||||
global_mcp_server_manager.expand_tool_permissions(
|
||||
user_api_key_auth.object_permission.mcp_tool_permissions
|
||||
).get(server.server_id)
|
||||
)
|
||||
if (
|
||||
allowed_tools_for_server is not None
|
||||
|
|
|
|||
|
|
@ -1391,55 +1391,73 @@ 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.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
# 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=[],
|
||||
)
|
||||
# Register placeholder ids in the manager so expand_permission_list resolves them.
|
||||
for sid in ("direct-server1", "direct-server2"):
|
||||
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:
|
||||
# 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=[],
|
||||
)
|
||||
|
||||
# Create mock user auth
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
team_id="team-789",
|
||||
)
|
||||
# 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:
|
||||
# Mock the helper methods
|
||||
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",
|
||||
]
|
||||
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",
|
||||
]
|
||||
|
||||
# Call the method
|
||||
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 the result contains both direct and access group servers
|
||||
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 _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"])
|
||||
# Verify access groups were resolved
|
||||
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)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1569,35 +1587,53 @@ async def test_get_allowed_mcp_servers_for_key_returns_empty_when_db_returns_non
|
|||
async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission():
|
||||
"""Ensure in-memory object_permission is used without hitting the DB."""
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
perms = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm-in-memory",
|
||||
mcp_servers=["direct-server"],
|
||||
mcp_access_groups=["grp-alpha"],
|
||||
)
|
||||
user_api_key_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
object_permission=perms,
|
||||
# Register "direct-server" in the manager so permission expansion resolves it.
|
||||
# Without this, expand_permission_list drops unknown ids as stale — which is
|
||||
# the correct production behavior but unrelated to what this test asserts.
|
||||
global_mcp_server_manager.registry["direct-server"] = MCPServer(
|
||||
server_id="direct-server",
|
||||
name="direct-server",
|
||||
server_name="direct-server",
|
||||
url="https://direct-server.example.com",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
try:
|
||||
perms = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm-in-memory",
|
||||
mcp_servers=["direct-server"],
|
||||
mcp_access_groups=["grp-alpha"],
|
||||
)
|
||||
user_api_key_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
object_permission=perms,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_object_permission",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_perm:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_mcp_servers_from_access_groups"
|
||||
) as mock_access_groups:
|
||||
mock_access_groups.return_value = ["group-server"]
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_object_permission",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_perm:
|
||||
with patch.object(
|
||||
MCPRequestHandler, "_get_mcp_servers_from_access_groups"
|
||||
) as mock_access_groups:
|
||||
mock_access_groups.return_value = ["group-server"]
|
||||
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(
|
||||
user_api_key_auth
|
||||
)
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(
|
||||
user_api_key_auth
|
||||
)
|
||||
|
||||
assert set(result) == {"direct-server", "group-server"}
|
||||
mock_get_perm.assert_not_called()
|
||||
mock_access_groups.assert_called_once_with(["grp-alpha"])
|
||||
assert set(result) == {"direct-server", "group-server"}
|
||||
mock_get_perm.assert_not_called()
|
||||
mock_access_groups.assert_called_once_with(["grp-alpha"])
|
||||
finally:
|
||||
global_mcp_server_manager.registry.pop("direct-server", None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1753,28 +1789,46 @@ async def test_tool_permission_servers_included_in_allowed_servers():
|
|||
|
||||
Regression test for https://github.com/BerriAI/litellm/issues/21954
|
||||
"""
|
||||
perm = MagicMock()
|
||||
perm.mcp_servers = []
|
||||
perm.mcp_access_groups = []
|
||||
perm.mcp_tool_permissions = {"server_id_123": ["tool_a", "tool_b"]}
|
||||
|
||||
user_api_key_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
MCPRequestHandler, "_get_key_object_permission", return_value=perm
|
||||
),
|
||||
patch.object(
|
||||
MCPRequestHandler,
|
||||
"_get_mcp_servers_from_access_groups",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[],
|
||||
),
|
||||
):
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
# Register the server id so expand_permission_list resolves it rather than
|
||||
# dropping it as stale.
|
||||
global_mcp_server_manager.registry["server_id_123"] = MCPServer(
|
||||
server_id="server_id_123",
|
||||
name="server_id_123",
|
||||
server_name="server_id_123",
|
||||
url="https://server-id-123.example.com",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
try:
|
||||
perm = MagicMock()
|
||||
perm.mcp_servers = []
|
||||
perm.mcp_access_groups = []
|
||||
perm.mcp_tool_permissions = {"server_id_123": ["tool_a", "tool_b"]}
|
||||
|
||||
user_api_key_auth = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
)
|
||||
assert "server_id_123" in result
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
MCPRequestHandler, "_get_key_object_permission", return_value=perm
|
||||
),
|
||||
patch.object(
|
||||
MCPRequestHandler,
|
||||
"_get_mcp_servers_from_access_groups",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[],
|
||||
),
|
||||
):
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
assert "server_id_123" in result
|
||||
finally:
|
||||
global_mcp_server_manager.registry.pop("server_id_123", None)
|
||||
|
|
|
|||
|
|
@ -2610,5 +2610,212 @@ class TestMCPServerManagerUpstreamInstructionsCache:
|
|||
assert "srv_b" in by_name and by_name["srv_b"].instructions is None
|
||||
|
||||
|
||||
class TestMCPServerManagerExpandPermissionList:
|
||||
"""Tests for the alias/name-aware permission list expansion used by team MCP permissions."""
|
||||
|
||||
def _make_server(
|
||||
self,
|
||||
server_id: str,
|
||||
server_name: str,
|
||||
alias=None,
|
||||
name=None,
|
||||
) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id=server_id,
|
||||
name=name if name is not None else (alias or server_name),
|
||||
alias=alias,
|
||||
server_name=server_name,
|
||||
url=f"https://{server_id}.example.com",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
|
||||
def test_empty_list_returns_empty(self):
|
||||
manager = MCPServerManager()
|
||||
assert manager.expand_permission_list([]) == []
|
||||
|
||||
def test_expands_server_name(self):
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["id-usw1"] = self._make_server(
|
||||
"id-usw1", server_name="a"
|
||||
)
|
||||
|
||||
assert manager.expand_permission_list(["a"]) == ["id-usw1"]
|
||||
|
||||
def test_expands_alias(self):
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["id-1"] = self._make_server(
|
||||
"id-1", server_name="internal_name", alias="public_alias"
|
||||
)
|
||||
|
||||
assert manager.expand_permission_list(["public_alias"]) == ["id-1"]
|
||||
|
||||
def test_passes_through_unknown_entry(self):
|
||||
"""Unresolved entries pass through unchanged (with a debug log) —
|
||||
the downstream access check denies them when compared to the
|
||||
concrete request server_id."""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["id-1"] = self._make_server("id-1", server_name="b")
|
||||
|
||||
assert manager.expand_permission_list(["a"]) == ["a"]
|
||||
|
||||
def test_name_collision_expands_to_all_matches(self):
|
||||
"""Two servers sharing a server_name both resolve — the documented behavior."""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["id-config"] = self._make_server(
|
||||
"id-config", server_name="shared"
|
||||
)
|
||||
manager.registry["id-db"] = self._make_server("id-db", server_name="shared")
|
||||
|
||||
assert sorted(manager.expand_permission_list(["shared"])) == [
|
||||
"id-config",
|
||||
"id-db",
|
||||
]
|
||||
|
||||
def test_searches_config_and_registry_union(self):
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["cfg-id"] = self._make_server(
|
||||
"cfg-id", server_name="a"
|
||||
)
|
||||
manager.registry["reg-id"] = self._make_server("reg-id", server_name="b")
|
||||
|
||||
assert manager.expand_permission_list(["a"]) == ["cfg-id"]
|
||||
assert manager.expand_permission_list(["b"]) == ["reg-id"]
|
||||
|
||||
def test_id_match_takes_precedence_over_name_match(self):
|
||||
"""
|
||||
If a permission entry matches a server_id directly, don't also add
|
||||
servers whose server_name happens to equal that id.
|
||||
"""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["id-1"] = self._make_server(
|
||||
"id-1", server_name="other_name"
|
||||
)
|
||||
manager.config_mcp_servers["id-2"] = self._make_server(
|
||||
"id-2", server_name="id-1"
|
||||
)
|
||||
|
||||
assert manager.expand_permission_list(["id-1"]) == ["id-1"]
|
||||
|
||||
def test_mixed_ids_and_names_in_same_list(self):
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-1"] = self._make_server(
|
||||
"uuid-1", server_name="a"
|
||||
)
|
||||
manager.config_mcp_servers["uuid-2"] = self._make_server(
|
||||
"uuid-2", server_name="b"
|
||||
)
|
||||
|
||||
# ["uuid-1", "b"] -> uuid-1 passes through, "b" resolves to uuid-2
|
||||
assert sorted(manager.expand_permission_list(["uuid-1", "b"])) == [
|
||||
"uuid-1",
|
||||
"uuid-2",
|
||||
]
|
||||
|
||||
def test_deduplicates_overlapping_id_and_name_entries(self):
|
||||
"""If a list references the same server by both id and name, return it once."""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-1"] = self._make_server(
|
||||
"uuid-1", server_name="a"
|
||||
)
|
||||
|
||||
assert manager.expand_permission_list(["uuid-1", "a"]) == ["uuid-1"]
|
||||
|
||||
def test_simulates_cross_region_portability(self):
|
||||
"""
|
||||
Same permission entry "a" resolves to different concrete IDs per region —
|
||||
the cross-region portability the customer is asking for.
|
||||
"""
|
||||
usw1 = MCPServerManager()
|
||||
usw1.config_mcp_servers["hash-usw1"] = self._make_server(
|
||||
"hash-usw1", server_name="a"
|
||||
)
|
||||
|
||||
usc1 = MCPServerManager()
|
||||
usc1.config_mcp_servers["hash-usc1"] = self._make_server(
|
||||
"hash-usc1", server_name="a"
|
||||
)
|
||||
|
||||
assert usw1.expand_permission_list(["a"]) == ["hash-usw1"]
|
||||
assert usc1.expand_permission_list(["a"]) == ["hash-usc1"]
|
||||
|
||||
|
||||
class TestMCPServerManagerExpandToolPermissions:
|
||||
"""Tests for tool-permission dict rewriting — the privilege-escalation guard."""
|
||||
|
||||
def _make_server(self, server_id: str, server_name: str, alias=None) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id=server_id,
|
||||
name=alias or server_name,
|
||||
alias=alias,
|
||||
server_name=server_name,
|
||||
url=f"https://{server_id}.example.com",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
|
||||
def test_empty_or_none_returns_empty_dict(self):
|
||||
manager = MCPServerManager()
|
||||
assert manager.expand_tool_permissions(None) == {}
|
||||
assert manager.expand_tool_permissions({}) == {}
|
||||
|
||||
def test_rewrites_name_key_to_server_id(self):
|
||||
"""Privilege-escalation guard: a name-based key must resolve to the
|
||||
concrete server_id, otherwise `.get(server_id)` misses and the tool
|
||||
restriction is silently dropped (caller treats None as allow-all)."""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-a"] = self._make_server(
|
||||
"uuid-a", server_name="my-alias"
|
||||
)
|
||||
|
||||
result = manager.expand_tool_permissions({"my-alias": ["read_file"]})
|
||||
assert result == {"uuid-a": ["read_file"]}
|
||||
|
||||
def test_passes_through_existing_server_id_key(self):
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-a"] = self._make_server(
|
||||
"uuid-a", server_name="alpha"
|
||||
)
|
||||
|
||||
result = manager.expand_tool_permissions({"uuid-a": ["read_file"]})
|
||||
assert result == {"uuid-a": ["read_file"]}
|
||||
|
||||
def test_unresolved_key_passes_through_unchanged(self):
|
||||
"""A stale id-keyed restriction (server since deleted, or just a
|
||||
test-fixture placeholder) must still apply when something looks it
|
||||
up by that same string — dropping would silently remove the
|
||||
restriction."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
result = manager.expand_tool_permissions({"stale-uuid": ["read_file"]})
|
||||
assert result == {"stale-uuid": ["read_file"]}
|
||||
|
||||
def test_name_collision_unions_tool_lists(self):
|
||||
"""Two servers sharing a server_name both match; their tool lists get
|
||||
the restriction (matches the list-expansion collision semantics)."""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-1"] = self._make_server(
|
||||
"uuid-1", server_name="shared"
|
||||
)
|
||||
manager.registry["uuid-2"] = self._make_server("uuid-2", server_name="shared")
|
||||
|
||||
result = manager.expand_tool_permissions({"shared": ["read_file"]})
|
||||
assert sorted(result.keys()) == ["uuid-1", "uuid-2"]
|
||||
assert result["uuid-1"] == ["read_file"]
|
||||
assert result["uuid-2"] == ["read_file"]
|
||||
|
||||
def test_id_and_name_keys_pointing_at_same_server_union_tools(self):
|
||||
"""If the admin writes both {"uuid-a": [...], "alias-a": [...]} and
|
||||
both refer to the same server, the tool lists are unioned rather
|
||||
than one overwriting the other."""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-a"] = self._make_server(
|
||||
"uuid-a", server_name="alias-a"
|
||||
)
|
||||
|
||||
result = manager.expand_tool_permissions(
|
||||
{"uuid-a": ["read_file"], "alias-a": ["write_file"]}
|
||||
)
|
||||
assert sorted(result["uuid-a"]) == ["read_file", "write_file"]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue