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:
Sameer Kankute 2026-05-02 11:25:55 +05:30 • committed by GitHub
commit f576eb3228
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 443 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
@ -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,

View file

@ -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"]