mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
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:
parent
efa33bfe50
commit
b540a71e47
2 changed files with 384 additions and 3 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue