From 767fbd4377d2ed8441b1f125afeaf465b305b844 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Christopher=20M=C3=B6ller?= Date: Wed, 20 Oct 2021 16:57:44 +0200 Subject: [PATCH] feat(noderenderer): render child nodes --- example/src/Basic/index.tsx | 22 ++- src/components/ConnectionLine/index.tsx | 3 +- src/components/Nodes/wrapNode.tsx | 2 +- src/container/EdgeRenderer/utils.ts | 4 +- src/container/NodeRenderer/index.tsx | 187 ++++++++++++++++-------- src/store/index.ts | 14 +- src/types/index.ts | 1 + src/utils/graph.ts | 16 +- 8 files changed, 180 insertions(+), 69 deletions(-) diff --git a/example/src/Basic/index.tsx b/example/src/Basic/index.tsx index e13d4a7a..e59f8f6e 100644 --- a/example/src/Basic/index.tsx +++ b/example/src/Basic/index.tsx @@ -23,7 +23,27 @@ const initialNodes: Node[] = [ { id: '1', type: 'input', data: { label: 'Node 1' }, position: { x: 250, y: 5 }, className: 'light' }, { id: '2', data: { label: 'Node 2' }, position: { x: 100, y: 100 }, className: 'light' }, { id: '3', data: { label: 'Node 3' }, position: { x: 400, y: 100 }, className: 'light' }, - { id: '4', data: { label: 'Node 4' }, position: { x: 400, y: 200 }, className: 'light' }, + { + id: '4', + data: { label: 'Node 4' }, + position: { x: 400, y: 200 }, + className: 'light', + style: { backgroundColor: 'rgba(255, 0, 0, .2)' }, + childNodes: [ + { id: '4a', data: { label: 'Node 4a' }, position: { x: 400, y: 400 }, className: 'light' }, + { + id: '4b', + data: { label: 'Node 4b' }, + position: { x: 500, y: 500 }, + className: 'light', + style: { backgroundColor: 'rgba(255, 0, 0, .2)' }, + childNodes: [ + { id: '4b1', data: { label: 'Node 4b1' }, position: { x: 450, y: 450 }, className: 'light' }, + { id: '4b2', data: { label: 'Node 4b2' }, position: { x: 550, y: 550 }, className: 'light' }, + ], + }, + ], + }, ]; const initialEdges: Edge[] = [ diff --git a/src/components/ConnectionLine/index.tsx b/src/components/ConnectionLine/index.tsx index 926ed692..aed5a5c2 100644 --- a/src/components/ConnectionLine/index.tsx +++ b/src/components/ConnectionLine/index.tsx @@ -14,6 +14,7 @@ import { HandleType, ReactFlowState, } from '../../types'; +import { flattenNodes } from '../../utils/graph'; interface ConnectionLineProps { connectionNodeId: ElementId; @@ -28,7 +29,7 @@ interface ConnectionLineProps { CustomConnectionLineComponent?: ConnectionLineComponent; } -const nodesSelector = (s: ReactFlowState) => s.nodes; +const nodesSelector = (s: ReactFlowState) => flattenNodes(s.nodes); export default ({ connectionNodeId, diff --git a/src/components/Nodes/wrapNode.tsx b/src/components/Nodes/wrapNode.tsx index b61d0550..6da0f95c 100644 --- a/src/components/Nodes/wrapNode.tsx +++ b/src/components/Nodes/wrapNode.tsx @@ -68,7 +68,7 @@ export default (NodeComponent: ComponentType) => { pointerEvents: isSelectable || isDraggable || onClick || onMouseEnter || onMouseMove || onMouseLeave ? 'all' : 'none', // prevents jumping of nodes on start - opacity: isInitialized ? 1 : 0, + // opacity: isInitialized ? 1 : 0, ...style, }), [ diff --git a/src/container/EdgeRenderer/utils.ts b/src/container/EdgeRenderer/utils.ts index d3e4ac20..19e3b9ac 100644 --- a/src/container/EdgeRenderer/utils.ts +++ b/src/container/EdgeRenderer/utils.ts @@ -2,7 +2,7 @@ import { ComponentType } from 'react'; import { BezierEdge, StepEdge, SmoothStepEdge, StraightEdge } from '../../components/Edges'; import wrapEdge from '../../components/Edges/wrapEdge'; -import { rectToBox } from '../../utils/graph'; +import { rectToBox, flattenNodes } from '../../utils/graph'; import { EdgeTypesType, @@ -171,7 +171,7 @@ type SourceTargetNode = { }; export const getSourceTargetNodes = (edge: Edge, nodes: Node[]): SourceTargetNode => { - return nodes.reduce( + return flattenNodes(nodes).reduce( (res, node) => { if (node.id === edge.source) { res.sourceNode = node; diff --git a/src/container/NodeRenderer/index.tsx b/src/container/NodeRenderer/index.tsx index ad12d642..f362f63a 100644 --- a/src/container/NodeRenderer/index.tsx +++ b/src/container/NodeRenderer/index.tsx @@ -1,9 +1,9 @@ -import React, { memo, useMemo, ComponentType, MouseEvent, useCallback } from 'react'; +import React, { memo, useMemo, ComponentType, MouseEvent, useCallback, Fragment } from 'react'; import shallow from 'zustand/shallow'; import { useStore } from '../../store'; -import { Node, NodeTypesType, ReactFlowState, WrapNodeProps } from '../../types'; -import { getNodesInside } from '../../utils/graph'; +import { Node, NodeTypesType, ReactFlowState, WrapNodeProps, SnapGrid } from '../../types'; +import { getNodesInside, getRectOfNodes } from '../../utils/graph'; interface NodeRendererProps { nodeTypes: NodeTypesType; selectNodesOnDrag: boolean; @@ -29,6 +29,122 @@ const selector = (s: ReactFlowState) => ({ snapToGrid: s.snapToGrid, }); +interface NodesProps extends NodeRendererProps { + nodes?: Node[]; + isDraggable?: boolean; + resizeObserver: ResizeObserver | null; + scale: number; + snapToGrid: boolean; + snapGrid: SnapGrid; + nodesDraggable: boolean; + nodesConnectable: boolean; + elementsSelectable: boolean; +} + +const Nodes = memo( + ({ + nodes = [], + isDraggable, + resizeObserver, + scale, + snapToGrid, + snapGrid, + nodesDraggable, + nodesConnectable, + elementsSelectable, + ...props + }: NodesProps) => { + return ( + <> + {nodes.map((node) => { + const nodeType = node.type || 'default'; + + if (!props.nodeTypes[nodeType]) { + console.warn(`Node type "${nodeType}" not found. Using fallback type "default".`); + } + + const NodeComponent = (props.nodeTypes[nodeType] || props.nodeTypes.default) as ComponentType; + const isNodeDraggable = + typeof isDraggable !== 'undefined' + ? isDraggable + : !!(node.draggable || (nodesDraggable && typeof node.draggable === 'undefined')); + const isSelectable = !!(node.selectable || (elementsSelectable && typeof node.selectable === 'undefined')); + const isConnectable = !!(node.connectable || (nodesConnectable && typeof node.connectable === 'undefined')); + const isInitialized = + node.width !== null && + node.height !== null && + typeof node.width !== 'undefined' && + typeof node.height !== 'undefined'; + + if (node.childNodes) { + const childRect = getRectOfNodes(node.childNodes); + node.position = node.isDragging + ? node.position + : { x: Math.round(childRect.x) - 10, y: Math.round(childRect.y) - 10 }; + node.style = { + ...node.style, + width: Math.round(childRect.width) + 20, + height: Math.round(childRect.height) + 20, + boxSizing: 'border-box', + }; + } + + return ( + + + {node.childNodes && ( + + )} + + ); + })} + + ); + } +); + const NodeRenderer = (props: NodeRendererProps) => { const { transform, @@ -75,60 +191,17 @@ const NodeRenderer = (props: NodeRendererProps) => { return (
- {nodes.map((node) => { - const nodeType = node.type || 'default'; - const NodeComponent = (props.nodeTypes[nodeType] || props.nodeTypes.default) as ComponentType; - - if (!props.nodeTypes[nodeType]) { - console.warn(`Node type "${nodeType}" not found. Using fallback type "default".`); - } - - const isDraggable = !!(node.draggable || (nodesDraggable && typeof node.draggable === 'undefined')); - const isSelectable = !!(node.selectable || (elementsSelectable && typeof node.selectable === 'undefined')); - const isConnectable = !!(node.connectable || (nodesConnectable && typeof node.connectable === 'undefined')); - const isInitialized = - node.width !== null && - node.height !== null && - typeof node.width !== 'undefined' && - typeof node.height !== 'undefined'; - - return ( - - ); - })} +
); }; diff --git a/src/store/index.ts b/src/store/index.ts index 1d9a85e9..eb5555ee 100644 --- a/src/store/index.ts +++ b/src/store/index.ts @@ -27,7 +27,7 @@ import { EdgeChange, NodePositionChange, } from '../types'; -import { isNode, isEdge, getRectOfNodes, getNodesInside, getConnectedEdges } from '../utils/graph'; +import { isNode, isEdge, getRectOfNodes, getNodesInside, getConnectedEdges, flattenNodes } from '../utils/graph'; import { getHandleBounds } from '../components/Nodes/utils'; const { Provider, useStore, useStoreApi } = createContext(); @@ -125,7 +125,7 @@ const createStore = () => const { onNodesChange, nodes, transform } = get(); const initialChanges: NodeChange[] = []; - const nodesToChange: NodeChange[] = nodes.reduce((res, node) => { + const nodesToChange: NodeChange[] = flattenNodes(nodes).reduce((res, node) => { const update = updates.find((u) => u.id === node.id); if (update) { const dimensions = getDimensions(update.nodeElement); @@ -155,11 +155,12 @@ const createStore = () => const { onNodesChange, nodes, nodeExtent } = get(); if (onNodesChange) { - const matchingNodes = nodes.filter((n) => n.id === id || n.isSelected); + const matchingNodes = flattenNodes(nodes).filter((n) => n.id === id || n.isSelected); + const matchingChildNodes = flattenNodes(matchingNodes); //.filter(n => !!n.childNodes).reduce((result, node) => result.concat(node.childNodes!), [])); - if (matchingNodes?.length) { + if (matchingChildNodes?.length) { onNodesChange( - matchingNodes.map((n) => { + matchingChildNodes.map((n) => { const change: NodePositionChange = { id: n.id, type: 'position', @@ -276,7 +277,8 @@ const createStore = () => }, unselectNodesAndEdges: () => { const { nodes, edges, onNodesChange, onEdgesChange } = get(); - const nodesToUnselect = nodes.map((n) => { + + const nodesToUnselect = flattenNodes(nodes).map((n) => { n.isSelected = false; return createNodeOrEdgeSelectionChange(false)(n); }) as NodeChange[]; diff --git a/src/types/index.ts b/src/types/index.ts index 064c4cc4..817657dd 100644 --- a/src/types/index.ts +++ b/src/types/index.ts @@ -86,6 +86,7 @@ export interface Node { width?: number | null; height?: number | null; handleBounds?: NodeHandleBounds; + childNodes?: Node[]; } export enum ArrowHeadType { diff --git a/src/utils/graph.ts b/src/utils/graph.ts index cda69ca1..4895966b 100644 --- a/src/utils/graph.ts +++ b/src/utils/graph.ts @@ -61,7 +61,7 @@ export const getMarkerId = (marker: EdgeMarkerType | undefined): string => { .join('&'); }; -const connectionExists = (edge: Edge, elements: Elements) => { +const connectionExists = (edge: Edge, elements: Edge[]) => { return elements.some( (el) => isEdge(el) && @@ -294,6 +294,10 @@ function applyChanges(changes: NodeChange[] | EdgeChange[], elements: any[]): an const initElements: any[] = []; return elements.reduce((res: any[], item: any) => { + if (item.childNodes) { + item.childNodes = applyChanges(changes, item.childNodes); + } + const currentChange = changes.find((c) => c.id === item.id); if (currentChange) { @@ -338,3 +342,13 @@ export function applyNodeChanges(changes: NodeChange[], nodes: Node[]): Node[] { export function applyEdgeChanges(changes: EdgeChange[], edges: Edge[]): Edge[] { return applyChanges(changes, edges) as Edge[]; } + +export function flattenNodes(nodes: Node[] | undefined): Node[] { + if (!nodes) { + return []; + } + + return nodes.reduce((result, node) => { + return result.concat([node, ...flattenNodes(node.childNodes)]); + }, []); +}