feat(mcp): enforce org-level MCP server and toolset permissions

Apply organization object_permission as a ceiling on allowed MCP servers
and tool permissions, consistent with vector store org checks.

Includes unit tests for org ceiling, intersection, and tool filtering.

Made-with: Cursor
This commit is contained in:
Sameer Kankute 2026-05-01 10:10:38 +05:30
parent efa33bfe50
commit b540a71e47
No known key found for this signature in database
2 changed files with 384 additions and 3 deletions

View file

@ -409,9 +409,12 @@ class MCPRequestHandler:
Permission hierarchy (all rules are intersections):
1. Get allowed servers from key permissions
2. Get allowed servers from team permissions
3. Get allowed servers from end_user permissions
4. Final result = intersection of key/team AND end_user (if end_user has permissions set)
2. Get allowed servers from team permissions (key inherits from team, or intersection)
3. Get allowed servers from end_user permissions (intersected if set)
4. Get allowed servers from agent permissions (intersected if set)
5. Get allowed servers from org permissions org acts as a ceiling: if the org
has an explicit MCP server list, the combined key/team/end_user/agent result is
capped to that list. If the org has no list, no extra restriction is applied.
Returns:
List[str]: List of allowed MCP servers by server id
@ -500,6 +503,30 @@ class MCPRequestHandler:
f"Applied agent intersection filter. Final allowed servers: {allowed_mcp_servers}"
)
#########################################################
# Apply org-level ceiling if org_id is set
#########################################################
if user_api_key_auth and user_api_key_auth.org_id:
allowed_mcp_servers_for_org = (
await MCPRequestHandler._get_allowed_mcp_servers_for_org(
user_api_key_auth
)
)
if len(allowed_mcp_servers_for_org) > 0:
if len(allowed_mcp_servers) > 0:
# Both have explicit lists → intersection
allowed_mcp_servers = [
s
for s in allowed_mcp_servers
if s in allowed_mcp_servers_for_org
]
else:
# No lower-level restrictions → org list becomes the ceiling
allowed_mcp_servers = allowed_mcp_servers_for_org
verbose_logger.debug(
f"Applied org ceiling filter. Final allowed servers: {allowed_mcp_servers}"
)
return list(set(allowed_mcp_servers))
except Exception as e:
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}")
@ -638,6 +665,23 @@ class MCPRequestHandler:
allowed_tools = list(set(allowed_tools) & set(agent_tools))
else:
allowed_tools = agent_tools
# Apply org-level tool ceiling if org_id is set
if user_api_key_auth.org_id:
org_obj_perm = await MCPRequestHandler._get_org_object_permission(
user_api_key_auth
)
org_tools = (
org_obj_perm.mcp_tool_permissions.get(server_id)
if org_obj_perm and org_obj_perm.mcp_tool_permissions
else None
)
if org_tools is not None:
if allowed_tools is not None:
allowed_tools = list(set(allowed_tools) & set(org_tools))
else:
allowed_tools = list(org_tools)
return allowed_tools
except Exception as e:
@ -805,6 +849,78 @@ class MCPRequestHandler:
)
return []
@staticmethod
async def _get_org_object_permission(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
):
"""
Get org object_permission by fetching the org row with object_permission included.
Note: get_org_object() in auth_checks.py does not include the object_permission
relation, so we do a targeted DB lookup here (same pattern as _get_agent_object_permission).
"""
from litellm.proxy.proxy_server import prisma_client
if not user_api_key_auth or not user_api_key_auth.org_id:
return None
if prisma_client is None:
verbose_logger.debug("prisma_client is None")
return None
try:
org_row = await prisma_client.db.litellm_organizationtable.find_unique(
where={"organization_id": user_api_key_auth.org_id},
include={"object_permission": True},
)
if org_row is None or org_row.object_permission is None:
return None
return org_row.object_permission
except Exception as e:
verbose_logger.warning(f"Failed to get org object permission: {str(e)}")
return None
@staticmethod
async def _get_allowed_mcp_servers_for_org(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[str]:
"""
Get allowed MCP servers for an organization.
Returns the MCP servers from the org's object_permission.
An empty result means the org places no restriction (allow-all from this level).
"""
try:
object_permissions = await MCPRequestHandler._get_org_object_permission(
user_api_key_auth
)
if object_permissions is None:
return []
# Direct server IDs
direct_mcp_servers = object_permissions.mcp_servers or []
# Servers from access groups
access_group_servers = (
await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
)
)
# Servers referenced only in tool permissions should also be accessible
tool_perm_servers = list(
(object_permissions.mcp_tool_permissions or {}).keys()
)
all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers
return list(set(all_servers))
except Exception as e:
verbose_logger.warning(
f"Failed to get allowed MCP servers for org: {str(e)}"
)
return []
@staticmethod
async def _get_allowed_mcp_servers_for_end_user(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,

View file

@ -2139,3 +2139,268 @@ async def test_tool_permission_servers_included_in_allowed_servers():
assert "server_id_123" in result
finally:
global_mcp_server_manager.registry.pop("server_id_123", None)
# ---------------------------------------------------------------------------
# Org-level MCP permission tests
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
class TestOrgMCPPermissions:
"""Tests for org-level MCP server permission enforcement."""
def _make_auth(self, org_id=None, team_id=None) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
team_id=team_id,
org_id=org_id,
)
@pytest.mark.parametrize(
"key_servers,team_servers,org_servers,expected,scenario",
[
(
["s1", "s2"],
[],
None,
["s1", "s2"],
"no_org_id",
),
(
["s1", "s2"],
[],
[],
["s1", "s2"],
"org_empty_no_restriction",
),
(
[],
[],
["org_s1", "org_s2"],
["org_s1", "org_s2"],
"org_only_ceiling",
),
(
["s1", "s2"],
[],
["s1", "org_only"],
["s1"],
"org_intersection",
),
(
["s1", "s2"],
[],
["org_s1"],
[],
"no_overlap_denied",
),
(
["s1", "s2"],
["s1", "s2", "s3"],
["s1"],
["s1"],
"team_then_org",
),
],
)
async def test_get_allowed_mcp_servers_with_org(
self,
key_servers,
team_servers,
org_servers,
expected,
scenario,
):
org_id = "org-123" if org_servers is not None else None
auth = self._make_auth(org_id=org_id)
with (
patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_key",
new_callable=AsyncMock,
return_value=key_servers,
),
patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_team",
new_callable=AsyncMock,
return_value=team_servers,
),
patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_org",
new_callable=AsyncMock,
return_value=org_servers if org_servers is not None else [],
),
):
result = await MCPRequestHandler.get_allowed_mcp_servers(auth)
assert sorted(result) == sorted(expected), f"scenario={scenario}"
async def test_get_org_object_permission_no_org_id(self):
auth = self._make_auth(org_id=None)
result = await MCPRequestHandler._get_org_object_permission(auth)
assert result is None
async def test_get_org_object_permission_no_prisma(self):
auth = self._make_auth(org_id="org-123")
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_org_object_permission",
new_callable=AsyncMock,
return_value=None,
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_org(auth)
assert result == []
async def test_get_allowed_mcp_servers_for_org_direct_servers(self):
auth = self._make_auth(org_id="org-123")
mock_perm = MagicMock()
mock_perm.mcp_servers = ["org_server_1", "org_server_2"]
mock_perm.mcp_access_groups = []
mock_perm.mcp_tool_permissions = {}
with (
patch.object(
MCPRequestHandler,
"_get_org_object_permission",
new_callable=AsyncMock,
return_value=mock_perm,
),
patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
),
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_org(auth)
assert sorted(result) == ["org_server_1", "org_server_2"]
async def test_get_allowed_mcp_servers_for_org_access_groups(self):
auth = self._make_auth(org_id="org-123")
mock_perm = MagicMock()
mock_perm.mcp_servers = []
mock_perm.mcp_access_groups = ["group-a"]
mock_perm.mcp_tool_permissions = {}
with (
patch.object(
MCPRequestHandler,
"_get_org_object_permission",
new_callable=AsyncMock,
return_value=mock_perm,
),
patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=["group_server_1"],
),
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_org(auth)
assert "group_server_1" in result
async def test_get_allowed_mcp_servers_for_org_tool_permissions_only(self):
auth = self._make_auth(org_id="org-123")
mock_perm = MagicMock()
mock_perm.mcp_servers = []
mock_perm.mcp_access_groups = []
mock_perm.mcp_tool_permissions = {"tool_only_server": ["tool_x"]}
with (
patch.object(
MCPRequestHandler,
"_get_org_object_permission",
new_callable=AsyncMock,
return_value=mock_perm,
),
patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
),
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_org(auth)
assert "tool_only_server" in result
async def test_get_allowed_mcp_servers_for_org_no_object_permission(self):
auth = self._make_auth(org_id="org-123")
with patch.object(
MCPRequestHandler,
"_get_org_object_permission",
new_callable=AsyncMock,
return_value=None,
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_org(auth)
assert result == []
async def test_get_allowed_tools_for_server_org_ceiling(self):
auth = self._make_auth(org_id="org-123")
key_perm = MagicMock()
key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b", "tool_c"]}
org_perm = MagicMock()
org_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]}
with (
patch.object(
MCPRequestHandler, "_get_key_object_permission", return_value=key_perm
),
patch.object(
MCPRequestHandler,
"_get_team_object_permission",
new_callable=AsyncMock,
return_value=None,
),
patch.object(
MCPRequestHandler,
"_get_org_object_permission",
new_callable=AsyncMock,
return_value=org_perm,
),
):
result = await MCPRequestHandler.get_allowed_tools_for_server(
server_id="server_1",
user_api_key_auth=auth,
)
assert sorted(result) == ["tool_a", "tool_b"]
async def test_get_allowed_tools_for_server_org_no_restriction(self):
auth = self._make_auth(org_id="org-123")
key_perm = MagicMock()
key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]}
org_perm = MagicMock()
org_perm.mcp_tool_permissions = {}
with (
patch.object(
MCPRequestHandler, "_get_key_object_permission", return_value=key_perm
),
patch.object(
MCPRequestHandler,
"_get_team_object_permission",
new_callable=AsyncMock,
return_value=None,
),
patch.object(
MCPRequestHandler,
"_get_org_object_permission",
new_callable=AsyncMock,
return_value=org_perm,
),
):
result = await MCPRequestHandler.get_allowed_tools_for_server(
server_id="server_1",
user_api_key_auth=auth,
)
assert sorted(result) == ["tool_a", "tool_b"]