mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(mcp): address Greptile P1/P2 findings in org-level permission enforcement
- Cache org object_permission in user_api_key_cache to avoid a DB hit on every MCP request (was: raw find_unique on every call). - Expand org mcp_servers list via expand_permission_list() so name-based entries resolve to canonical IDs, consistent with key/team/end-user path. - Expand org mcp_tool_permissions via expand_tool_permissions() in both _get_allowed_mcp_servers_for_org and get_allowed_tools_for_server, closing the silent name-vs-ID mismatch that could let restricted tools through. - The second _get_org_object_permission call in get_allowed_tools_for_server now hits the cache (warm from the earlier server-list check), resolving the double DB round-trip without changing the call structure. Made-with: Cursor
This commit is contained in:
parent
b540a71e47
commit
8398952e81
1 changed files with 30 additions and 13 deletions
|
|
@ -668,11 +668,15 @@ class MCPRequestHandler:
|
|||
|
||||
# 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 = (
|
||||
org_obj_perm.mcp_tool_permissions.get(server_id)
|
||||
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
|
||||
)
|
||||
|
|
@ -854,12 +858,9 @@ class MCPRequestHandler:
|
|||
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).
|
||||
Get org object_permission, using user_api_key_cache to avoid DB hits on every request.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
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
|
||||
|
|
@ -868,14 +869,24 @@ class MCPRequestHandler:
|
|||
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}"
|
||||
|
||||
try:
|
||||
cached = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
org_row = await prisma_client.db.litellm_organizationtable.find_unique(
|
||||
where={"organization_id": user_api_key_auth.org_id},
|
||||
where={"organization_id": 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
|
||||
|
||||
obj_perm = org_row.object_permission
|
||||
await user_api_key_cache.async_set_cache(key=cache_key, value=obj_perm)
|
||||
return obj_perm
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get org object permission: {str(e)}")
|
||||
return None
|
||||
|
|
@ -898,19 +909,25 @@ class MCPRequestHandler:
|
|||
if object_permissions is None:
|
||||
return []
|
||||
|
||||
# Direct server IDs
|
||||
direct_mcp_servers = object_permissions.mcp_servers or []
|
||||
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 []
|
||||
)
|
||||
|
||||
# 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()
|
||||
global_mcp_server_manager.expand_tool_permissions(
|
||||
object_permissions.mcp_tool_permissions
|
||||
).keys()
|
||||
)
|
||||
|
||||
all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue