mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge pull request #26960 from BerriAI/litellm_org_mcp_permissions
feat(mcp): enforce org-level MCP server and toolset permissions
This commit is contained in:
commit
f576eb3228
2 changed files with 443 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
|
||||
|
|
@ -435,6 +438,10 @@ class MCPRequestHandler:
|
|||
# Calculate key/team allowed servers using inheritance and intersection logic
|
||||
#########################################################
|
||||
allowed_mcp_servers: List[str] = []
|
||||
has_lower_level_mcp_restrictions = (
|
||||
len(allowed_mcp_servers_for_key) > 0
|
||||
or len(allowed_mcp_servers_for_team) > 0
|
||||
)
|
||||
if len(allowed_mcp_servers_for_team) > 0:
|
||||
if len(allowed_mcp_servers_for_key) > 0:
|
||||
# Key has its own MCP permissions - use intersection with team permissions
|
||||
|
|
@ -459,6 +466,7 @@ class MCPRequestHandler:
|
|||
|
||||
# If end_user has explicit MCP server permissions, apply intersection
|
||||
if len(allowed_mcp_servers_for_end_user) > 0:
|
||||
has_lower_level_mcp_restrictions = True
|
||||
verbose_logger.debug(
|
||||
f"End user {user_api_key_auth.end_user_id} has explicit MCP permissions: {allowed_mcp_servers_for_end_user}"
|
||||
)
|
||||
|
|
@ -490,6 +498,7 @@ class MCPRequestHandler:
|
|||
)
|
||||
)
|
||||
if len(allowed_mcp_servers_for_agent) > 0:
|
||||
has_lower_level_mcp_restrictions = True
|
||||
# Intersect: agent can only use servers allowed by BOTH key/team AND agent config
|
||||
allowed_mcp_servers = [
|
||||
s
|
||||
|
|
@ -500,6 +509,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 has_lower_level_mcp_restrictions:
|
||||
# Lower-level restrictions exist, so org can only cap them.
|
||||
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 +671,27 @@ 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:
|
||||
# _get_org_object_permission uses user_api_key_cache, so this is not a
|
||||
# fresh DB round-trip when get_allowed_mcp_servers was already called.
|
||||
org_obj_perm = await MCPRequestHandler._get_org_object_permission(
|
||||
user_api_key_auth
|
||||
)
|
||||
org_tools = (
|
||||
global_mcp_server_manager.expand_tool_permissions(
|
||||
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 +859,120 @@ class MCPRequestHandler:
|
|||
)
|
||||
return []
|
||||
|
||||
# Sentinel stored in cache when an org has no object_permission, so we
|
||||
# don't re-query the DB on every MCP request for that org.
|
||||
_ORG_NO_PERMISSION_SENTINEL = "__org_no_mcp_permission__"
|
||||
|
||||
@staticmethod
|
||||
async def _get_org_object_permission(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
):
|
||||
"""
|
||||
Get org object_permission, using user_api_key_cache to avoid DB hits on every request.
|
||||
|
||||
Caches both positive results and the absence of an object_permission so that orgs
|
||||
with no MCP permissions configured (the common default) do not trigger a DB query
|
||||
on every request.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
|
||||
|
||||
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
|
||||
|
||||
org_id = user_api_key_auth.org_id
|
||||
cache_key = f"org_object_permission:{org_id}"
|
||||
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
|
||||
try:
|
||||
cached = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
if cached is not None:
|
||||
# Sentinel means the DB confirmed no object_permission for this org
|
||||
if cached == MCPRequestHandler._ORG_NO_PERMISSION_SENTINEL:
|
||||
return None
|
||||
# Redis deserialises to a plain dict; reconstruct the Pydantic model
|
||||
# so callers can access .mcp_servers / .mcp_tool_permissions as attrs.
|
||||
if isinstance(cached, dict):
|
||||
return LiteLLM_ObjectPermissionTable(**cached)
|
||||
return cached
|
||||
|
||||
org_row = await prisma_client.db.litellm_organizationtable.find_unique(
|
||||
where={"organization_id": org_id},
|
||||
include={"object_permission": True},
|
||||
)
|
||||
|
||||
if org_row is None or org_row.object_permission is None:
|
||||
# Cache the negative result so subsequent calls skip the DB
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=MCPRequestHandler._ORG_NO_PERMISSION_SENTINEL,
|
||||
)
|
||||
return None
|
||||
|
||||
# Convert raw Prisma model → Pydantic before caching. Caching the
|
||||
# Pydantic .dict() ensures the value survives a Redis JSON round-trip
|
||||
# as a plain dict that we can reconstruct above (same pattern used by
|
||||
# get_end_user_object / get_team_object in auth_checks.py).
|
||||
obj_perm = LiteLLM_ObjectPermissionTable(**org_row.object_permission.dict())
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=cache_key, value=obj_perm.dict()
|
||||
)
|
||||
return obj_perm
|
||||
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 []
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
# Expand names/aliases to canonical server IDs (consistent with key/team/end-user path)
|
||||
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(
|
||||
object_permissions.mcp_servers or []
|
||||
)
|
||||
|
||||
access_group_servers = (
|
||||
await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
object_permissions.mcp_access_groups or []
|
||||
)
|
||||
)
|
||||
|
||||
tool_perm_servers = list(
|
||||
global_mcp_server_manager.expand_tool_permissions(
|
||||
object_permissions.mcp_tool_permissions
|
||||
).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,275 @@ 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",
|
||||
),
|
||||
(
|
||||
["s1"],
|
||||
["s2"],
|
||||
["s1", "s2", "org_s1"],
|
||||
[],
|
||||
"key_team_conflict_not_expanded_by_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