diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 8aba8307b9d..703d9360a3d 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -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}, + )