feat(mcp): validate mcp_toolsets in key-vs-team permission check

This commit is contained in:
Ishaan Jaffer 2026-03-21 18:59:37 -07:00
parent 7ac2e590be
commit fc1558f6a9

View file

@ -208,10 +208,10 @@ async def _resolve_team_allowed_mcp_servers(
)
direct_servers: List[str] = team_object_permission.mcp_servers or []
access_group_servers: List[
str
] = await MCPRequestHandler._get_mcp_servers_from_access_groups(
team_object_permission.mcp_access_groups or []
access_group_servers: List[str] = (
await MCPRequestHandler._get_mcp_servers_from_access_groups(
team_object_permission.mcp_access_groups or []
)
)
raw_tool_perms = team_object_permission.mcp_tool_permissions or {}
if isinstance(raw_tool_perms, str):
@ -286,6 +286,19 @@ def _extract_requested_mcp_access_groups(
return set()
def _extract_requested_mcp_toolsets(
object_permission: Optional[dict],
) -> Set[str]:
"""Extract MCP toolset IDs from a key's object_permission dict."""
if not object_permission or not isinstance(object_permission, dict):
return set()
toolsets = object_permission.get("mcp_toolsets")
if isinstance(toolsets, list):
return set(toolsets)
return set()
async def validate_key_mcp_servers_against_team(
object_permission: Optional[dict],
team_obj: Optional["LiteLLM_TeamTableCachedObj"],
@ -305,8 +318,10 @@ async def validate_key_mcp_servers_against_team(
requested_servers = _extract_requested_mcp_server_ids(object_permission)
requested_access_groups = _extract_requested_mcp_access_groups(object_permission)
requested_toolsets = _extract_requested_mcp_toolsets(object_permission)
# Nothing to validate
if not requested_servers and not requested_access_groups:
if not requested_servers and not requested_access_groups and not requested_toolsets:
return
allow_all_keys_servers = _get_allow_all_keys_server_ids()
@ -364,3 +379,31 @@ async def validate_key_mcp_servers_against_team(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": detail},
)
# Validate requested toolsets (must be subset of team's toolsets)
if requested_toolsets:
team_toolsets: Set[str] = set()
if (
team_obj is not None
and team_obj.object_permission is not None
and team_obj.object_permission.mcp_toolsets
):
team_toolsets = set(team_obj.object_permission.mcp_toolsets)
disallowed_toolsets = requested_toolsets - team_toolsets
if disallowed_toolsets:
if team_obj is not None:
detail = (
f"Key requests MCP toolsets not allowed by team '{team_obj.team_id}': "
f"{sorted(disallowed_toolsets)}. "
f"Team allows: {sorted(team_toolsets)}."
)
else:
detail = (
f"Key is not in a team. MCP toolsets cannot be assigned to "
f"keys outside of a team. Disallowed toolsets: {sorted(disallowed_toolsets)}."
)
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": detail},
)