From b540a71e47e88c0080993eb60c1dad4966e13fe5 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 1 May 2026 10:10:38 +0530 Subject: [PATCH] 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 --- .../mcp_server/auth/user_api_key_auth_mcp.py | 122 +++++++- .../auth/test_user_api_key_auth_mcp.py | 265 ++++++++++++++++++ 2 files changed, 384 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 756b2ed91d7..8ffaada3db2 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index dd352d0999a..c4bcad64fb4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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"]