From 1525af39cfbed907cbd5268846b3c0f919881631 Mon Sep 17 00:00:00 2001 From: moklick Date: Tue, 19 Oct 2021 14:57:34 +0200 Subject: [PATCH] refactor(elements): render only visible elements --- cypress/integration/flow/graph-utils.spec.js | 64 +------- example/src/Basic/index.tsx | 4 +- example/src/CustomNode/ColorSelectorNode.tsx | 25 ++++ example/src/CustomNode/index.tsx | 147 +++++++++++++++++++ example/src/Stress/index.tsx | 29 ++-- example/src/Stress/utils.ts | 2 +- example/src/index.tsx | 5 + src/components/Handle/handler.ts | 8 +- src/components/Handle/index.tsx | 8 +- src/components/Nodes/wrapNode.tsx | 14 +- src/components/SelectionListener/index.tsx | 18 +-- src/container/EdgeRenderer/index.tsx | 60 +++++--- src/container/EdgeRenderer/utils.ts | 20 ++- src/container/FlowRenderer/index.tsx | 3 - src/container/GraphView/index.tsx | 12 +- src/container/NodeRenderer/index.tsx | 19 ++- src/container/ReactFlow/index.tsx | 6 +- src/index.ts | 1 - src/store/index.ts | 72 ++++----- src/types/index.ts | 8 +- src/utils/graph.ts | 68 ++++----- 21 files changed, 376 insertions(+), 217 deletions(-) create mode 100644 example/src/CustomNode/ColorSelectorNode.tsx create mode 100644 example/src/CustomNode/index.tsx diff --git a/cypress/integration/flow/graph-utils.spec.js b/cypress/integration/flow/graph-utils.spec.js index e582231b..e6df1455 100644 --- a/cypress/integration/flow/graph-utils.spec.js +++ b/cypress/integration/flow/graph-utils.spec.js @@ -1,4 +1,4 @@ -import { isNode, isEdge, getOutgoers, getIncomers, removeElements, addEdge } from '../../../dist/ReactFlow.js'; +import { isNode, isEdge, getOutgoers, getIncomers, addEdge } from '../../../dist/ReactFlow.js'; const nodes = [ { id: '1', type: 'input', data: { label: 'Node 1' }, position: { x: 250, y: 5 } }, @@ -69,66 +69,4 @@ describe('Graph Utils Testing', () => { } }); }); - - describe('tests removeElements function', () => { - it('removes a node', () => { - const nextElements = removeElements([nodes[0]], elements); - - const nextNodes = nextElements.filter((e) => isNode(e)); - const nextEdges = nextElements.filter((e) => isEdge(e)); - - expect(nextNodes.length).to.be.equal(nodes.length - 1); - expect(nextEdges.length).to.be.equal(edges.length - 2); - }); - - it('removes multiple nodes', () => { - const elementsToRemove = [nodes[0], nodes[1]]; - const nextElements = removeElements(elementsToRemove, elements); - const nextNodes = nextElements.filter((e) => isNode(e)); - const nextEdges = nextElements.filter((e) => isEdge(e)); - - expect(nextNodes.length).to.be.equal(nodes.length - 2); - expect(nextEdges.length).to.be.equal(0); - }); - - it('removes no node', () => { - const nextElementsNoRemove = removeElements([], elements); - expect(nextElementsNoRemove.length).to.be.equal(elements.length); - }); - - it('tries to removes node that does not exist', () => { - const nextElementsNoRemove = removeElements([{ id: 'id-that-does-not-exist' }], elements); - expect(nextElementsNoRemove.length).to.be.equal(elements.length); - }); - - it('removes an edge', () => { - const nextElements = removeElements([edges[0]], elements); - - const nextNodes = nextElements.filter((e) => isNode(e)); - const nextEdges = nextElements.filter((e) => isEdge(e)); - - expect(nextNodes.length).to.be.equal(nodes.length); - expect(nextEdges.length).to.be.equal(edges.length - 1); - }); - - it('removes multiple edges', () => { - const nextElements = removeElements([edges[0], edges[1]], elements); - - const nextNodes = nextElements.filter((e) => isNode(e)); - const nextEdges = nextElements.filter((e) => isEdge(e)); - - expect(nextNodes.length).to.be.equal(nodes.length); - expect(nextEdges.length).to.be.equal(edges.length - 2); - }); - - it('removes node and edge', () => { - const nextElements = removeElements([nodes[0], edges[0]], elements); - - const nextNodes = nextElements.filter((e) => isNode(e)); - const nextEdges = nextElements.filter((e) => isEdge(e)); - - expect(nextNodes.length).to.be.equal(nodes.length - 1); - expect(nextEdges.length).to.be.equal(edges.length - 2); - }); - }); }); diff --git a/example/src/Basic/index.tsx b/example/src/Basic/index.tsx index f47d372a..e13d4a7a 100644 --- a/example/src/Basic/index.tsx +++ b/example/src/Basic/index.tsx @@ -36,8 +36,8 @@ const BasicFlow = () => { const [nodes, setNodes] = useState(initialNodes); const [edges, setEdges] = useState(initialEdges); - const onConnect = useCallback((params: Edge | Connection, nds: Node[]) => { - setEdges((eds) => addEdge(params, nds, eds)); + const onConnect = useCallback((params: Edge | Connection) => { + setEdges((eds) => addEdge(params, eds)); }, []); const onLoad = useCallback((reactFlowInstance: OnLoadParams) => setRfInstance(reactFlowInstance), []); diff --git a/example/src/CustomNode/ColorSelectorNode.tsx b/example/src/CustomNode/ColorSelectorNode.tsx new file mode 100644 index 00000000..7ad20340 --- /dev/null +++ b/example/src/CustomNode/ColorSelectorNode.tsx @@ -0,0 +1,25 @@ +import React, { memo, FC, CSSProperties } from 'react'; + +import { Handle, Position, NodeProps, Connection, Edge } from 'react-flow-renderer'; + +const targetHandleStyle: CSSProperties = { background: '#555' }; +const sourceHandleStyleA: CSSProperties = { ...targetHandleStyle, top: 10 }; +const sourceHandleStyleB: CSSProperties = { ...targetHandleStyle, bottom: 10, top: 'auto' }; + +const onConnect = (params: Connection | Edge) => console.log('handle onConnect', params); + +const ColorSelectorNode: FC = ({ data, isConnectable }) => { + return ( + <> + +
+ Custom Color Picker Node: {data.color} +
+ + + + + ); +}; + +export default memo(ColorSelectorNode); diff --git a/example/src/CustomNode/index.tsx b/example/src/CustomNode/index.tsx new file mode 100644 index 00000000..5c786f58 --- /dev/null +++ b/example/src/CustomNode/index.tsx @@ -0,0 +1,147 @@ +import { useState, useEffect, MouseEvent, useCallback } from 'react'; +import { ChangeEvent } from 'react'; + +import ReactFlow, { + addEdge, + MiniMap, + Controls, + Node, + OnLoadParams, + Position, + SnapGrid, + Connection, + Edge, + NodeChange, + applyNodeChanges, + applyEdgeChanges, + EdgeChange, +} from 'react-flow-renderer'; + +import ColorSelectorNode from './ColorSelectorNode'; + +const onLoad = (reactFlowInstance: OnLoadParams) => console.log('flow loaded:', reactFlowInstance); +const onNodeDragStop = (_: MouseEvent, node: Node) => console.log('drag stop', node); +const onNodeClick = (_: MouseEvent, node: Node) => console.log('click', node); + +const initBgColor = '#1A192B'; + +const connectionLineStyle = { stroke: '#fff' }; +const snapGrid: SnapGrid = [16, 16]; +const nodeTypes = { + selectorNode: ColorSelectorNode, +}; + +const CustomNodeFlow = () => { + const [nodes, setNodes] = useState([]); + const [edges, setEdges] = useState([]); + const [bgColor, setBgColor] = useState(initBgColor); + + useEffect(() => { + const onChange = (event: ChangeEvent) => { + setNodes((nds) => + nds.map((node) => { + if (node.id !== '2') { + return node; + } + + const color = event.target.value; + + setBgColor(color); + + return { + ...node, + data: { + ...node.data, + color, + }, + }; + }) + ); + }; + + setNodes([ + { + id: '1', + type: 'input', + data: { label: 'An input node' }, + position: { x: 0, y: 50 }, + sourcePosition: Position.Right, + }, + { + id: '2', + type: 'selectorNode', + data: { onChange: onChange, color: initBgColor }, + style: { border: '1px solid #777', padding: 10 }, + position: { x: 250, y: 50 }, + }, + { + id: '3', + type: 'output', + data: { label: 'Output A' }, + position: { x: 550, y: 25 }, + targetPosition: Position.Left, + }, + { + id: '4', + type: 'output', + data: { label: 'Output B' }, + position: { x: 550, y: 100 }, + targetPosition: Position.Left, + }, + ]); + + setEdges([ + { id: 'e1-2', source: '1', target: '2', animated: true, style: { stroke: '#fff' } }, + { id: 'e2a-3', source: '2', sourceHandle: 'a', target: '3', animated: true, style: { stroke: '#fff' } }, + { id: 'e2b-4', source: '2', sourceHandle: 'b', target: '4', animated: true, style: { stroke: '#fff' } }, + ]); + }, []); + + const onConnect = (params: Connection | Edge) => + setEdges((eds) => addEdge({ ...params, animated: true, style: { stroke: '#fff' } }, eds)); + + const onNodesChange = useCallback((changes: NodeChange[]) => { + setNodes((ns) => applyNodeChanges(changes, ns)); + }, []); + + const onEdgesChange = useCallback((changes: EdgeChange[]) => { + setEdges((es) => applyEdgeChanges(changes, es)); + }, []); + + return ( + + { + if (n.type === 'input') return '#0041d0'; + if (n.type === 'selectorNode') return bgColor; + if (n.type === 'output') return '#ff0072'; + + return '#eee'; + }} + nodeColor={(n: Node): string => { + if (n.type === 'selectorNode') return bgColor; + + return '#fff'; + }} + /> + + + ); +}; + +export default CustomNodeFlow; diff --git a/example/src/Stress/index.tsx b/example/src/Stress/index.tsx index f3cf1df2..e562ed50 100644 --- a/example/src/Stress/index.tsx +++ b/example/src/Stress/index.tsx @@ -8,25 +8,27 @@ import ReactFlow, { Node, NodeChange, applyNodeChanges, + Connection, + addEdge, } from 'react-flow-renderer'; -import { getElements } from './utils'; +import { getNodesAndEdges } from './utils'; const buttonWrapperStyles: CSSProperties = { position: 'absolute', right: 10, top: 10, zIndex: 4 }; const onLoad = (reactFlowInstance: OnLoadParams) => { reactFlowInstance.fitView(); - console.log(reactFlowInstance.getElements()); + console.log(reactFlowInstance.getNodes()); }; -const initialElements = getElements(30, 30); +const { nodes: initialNodes, edges: initialEdges } = getNodesAndEdges(30, 30); const StressFlow = () => { - const [nodes, setNodes] = useState(initialElements.nodes); - const [edges, setEdges] = useState(initialElements.edges); - // const onElementsRemove = (elementsToRemove: Elements) => setElements((els) => removeElements(elementsToRemove, els)); - // const onConnect = (params: Connection | Edge, nds: Node[]) => setElements((els) => addEdge(params, els)); - + const [nodes, setNodes] = useState(initialNodes); + const [edges, setEdges] = useState(initialEdges); + const onConnect = useCallback((params: Edge | Connection) => { + setEdges((eds) => addEdge(params, eds)); + }, []); const updatePos = () => { setNodes((nds) => { return nds.map((n) => { @@ -43,7 +45,7 @@ const StressFlow = () => { const updateElements = () => { const grid = Math.ceil(Math.random() * 10); - const initialElements = getElements(grid, grid); + const initialElements = getNodesAndEdges(grid, grid); setNodes(initialElements.nodes); setEdges(initialElements.edges); }; @@ -53,7 +55,14 @@ const StressFlow = () => { }, []); return ( - + diff --git a/example/src/Stress/utils.ts b/example/src/Stress/utils.ts index 36cf2949..c5c6e307 100644 --- a/example/src/Stress/utils.ts +++ b/example/src/Stress/utils.ts @@ -5,7 +5,7 @@ type ElementsCollection = { edges: Edge[]; }; -export function getElements(xElements: number = 10, yElements: number = 10): ElementsCollection { +export function getNodesAndEdges(xElements: number = 10, yElements: number = 10): ElementsCollection { const initialNodes = []; const initialEdges: Edge[] = []; let nodeId = 1; diff --git a/example/src/index.tsx b/example/src/index.tsx index d6acbeba..6e08eabd 100644 --- a/example/src/index.tsx +++ b/example/src/index.tsx @@ -5,6 +5,7 @@ import { BrowserRouter as Router, Route, Switch, withRouter } from 'react-router import Basic from './Basic'; import UpdateNode from './UpdateNode'; import Stress from './Stress'; +import CustomNode from './CustomNode'; import './index.css'; @@ -21,6 +22,10 @@ const routes = [ path: '/stress', component: Stress, }, + { + path: '/custom-node', + component: CustomNode, + }, ]; const Header = withRouter(({ history, location }) => { diff --git a/src/components/Handle/handler.ts b/src/components/Handle/handler.ts index 53f7d213..03fd37df 100644 --- a/src/components/Handle/handler.ts +++ b/src/components/Handle/handler.ts @@ -1,8 +1,6 @@ import { MouseEvent as ReactMouseEvent } from 'react'; -import { GetState } from 'zustand'; import { getHostForElement } from '../../utils'; -import { ReactFlowState } from '../../types'; import { ElementId, @@ -105,8 +103,7 @@ export function onMouseDown( onEdgeUpdateEnd?: (evt: MouseEvent) => void, onConnectStart?: OnConnectStartFunc, onConnectStop?: OnConnectStopFunc, - onConnectEnd?: OnConnectEndFunc, - getState?: GetState + onConnectEnd?: OnConnectEndFunc ): void { const reactFlowNode = (event.target as Element).closest('.react-flow'); // when react-flow is used inside a shadow root we can't use document @@ -180,8 +177,7 @@ export function onMouseDown( onConnectStop?.(event); if (isValid) { - const nodes = getState?.().nodes; - onConnect?.(connection, nodes || []); + onConnect?.(connection); } onConnectEnd?.(event); diff --git a/src/components/Handle/index.tsx b/src/components/Handle/index.tsx index 490b1b98..2f594b27 100644 --- a/src/components/Handle/index.tsx +++ b/src/components/Handle/index.tsx @@ -4,7 +4,7 @@ import shallow from 'zustand/shallow'; import { useStore } from '../../store'; import NodeIdContext from '../../contexts/NodeIdContext'; -import { HandleProps, Connection, ElementId, Position, Node, ReactFlowState } from '../../types'; +import { HandleProps, Connection, ElementId, Position, ReactFlowState } from '../../types'; import { onMouseDown, SetSourceIdFunc, SetPosition } from './handler'; @@ -52,9 +52,9 @@ const Handle = forwardRef( const isTarget = type === 'target'; const onConnectExtended = useCallback( - (params: Connection, nodes: Node[]) => { - onConnectAction?.(params, nodes); - onConnect?.(params, nodes); + (params: Connection) => { + onConnectAction?.(params); + onConnect?.(params); }, [onConnectAction, onConnect] ); diff --git a/src/components/Nodes/wrapNode.tsx b/src/components/Nodes/wrapNode.tsx index ebf1106b..b61d0550 100644 --- a/src/components/Nodes/wrapNode.tsx +++ b/src/components/Nodes/wrapNode.tsx @@ -12,6 +12,7 @@ const selector = (s: ReactFlowState) => ({ unsetNodesSelection: s.unsetNodesSelection, updateNodePosition: s.updateNodePosition, updateNodeDimensions: s.updateNodeDimensions, + unselectNodesAndEdges: s.unselectNodesAndEdges, }); export default (NodeComponent: ComponentType) => { @@ -48,10 +49,13 @@ export default (NodeComponent: ComponentType) => { resizeObserver, dragHandle, }: WrapNodeProps) => { - const { addSelectedElements, unsetNodesSelection, updateNodePosition, updateNodeDimensions } = useStore( - selector, - shallow - ); + const { + addSelectedElements, + unselectNodesAndEdges, + unsetNodesSelection, + updateNodePosition, + updateNodeDimensions, + } = useStore(selector, shallow); const nodeElement = useRef(null); const node = useMemo(() => ({ id, type, position: { x: xPos, y: yPos }, data }), [id, type, xPos, yPos, data]); @@ -142,8 +146,8 @@ export default (NodeComponent: ComponentType) => { addSelectedElements([node]); } } else if (!selectNodesOnDrag && !isSelected && isSelectable) { + unselectNodesAndEdges(); unsetNodesSelection(); - addSelectedElements([]); } }, [node, isSelected, selectNodesOnDrag, isSelectable, onNodeDragStart] diff --git a/src/components/SelectionListener/index.tsx b/src/components/SelectionListener/index.tsx index a99a95a0..a4a92441 100644 --- a/src/components/SelectionListener/index.tsx +++ b/src/components/SelectionListener/index.tsx @@ -1,26 +1,26 @@ import { useEffect } from 'react'; import shallow from 'zustand/shallow'; -import { Elements, ReactFlowState } from '../../types'; +import { ReactFlowState, OnSelectionChangeFunc } from '../../types'; import { useStore } from '../../store'; interface SelectionListenerProps { - onSelectionChange: (elements: Elements | null) => void; + onSelectionChange: OnSelectionChangeFunc; } -const selectedElementsSelector = (s: ReactFlowState) => [ - ...s.nodes.filter((n) => n.isSelected), - ...s.edges.filter((e) => e.isSelected), -]; +const selectedElementsSelector = (s: ReactFlowState) => ({ + selectedNodes: s.nodes.filter((n) => n.isSelected), + selectedEdges: s.edges.filter((e) => e.isSelected), +}); // This is just a helper component for calling the onSelectionChange listener. export default ({ onSelectionChange }: SelectionListenerProps) => { - const selectedElements = useStore(selectedElementsSelector, shallow); + const { selectedNodes, selectedEdges } = useStore(selectedElementsSelector, shallow); useEffect(() => { - onSelectionChange(selectedElements); - }, [selectedElements]); + onSelectionChange({ nodes: selectedNodes, edges: selectedEdges }); + }, [selectedNodes, selectedEdges]); return null; }; diff --git a/src/container/EdgeRenderer/index.tsx b/src/container/EdgeRenderer/index.tsx index 739aa4a6..8abe99a5 100644 --- a/src/container/EdgeRenderer/index.tsx +++ b/src/container/EdgeRenderer/index.tsx @@ -4,11 +4,10 @@ import shallow from 'zustand/shallow'; import { useStore } from '../../store'; import ConnectionLine from '../../components/ConnectionLine/index'; import MarkerDefinitions from './MarkerDefinitions'; -import { getEdgePositions, getHandle, getSourceTargetNodes } from './utils'; +import { getEdgePositions, getHandle, getSourceTargetNodes, isEdgeVisible } from './utils'; import { Position, Edge, - Node, Connection, ConnectionLineType, ConnectionLineComponent, @@ -19,8 +18,6 @@ import { } from '../../types'; interface EdgeRendererProps { - nodes: Node[]; - edges: Edge[]; edgeTypes: any; connectionLineType: ConnectionLineType; connectionLineStyle?: CSSProperties; @@ -162,26 +159,12 @@ const Edge = memo( targetPosition ); - // const isVisible = onlyRenderVisibleElements - // ? isEdgeVisible({ - // sourcePos: { x: sourceX, y: sourceY }, - // targetPos: { x: targetX, y: targetY }, - // width, - // height, - // transform, - // }) - // : true; - - // if (!isVisible) { - // return null; - // } - return ( ({ width: s.width, height: s.height, connectionMode: s.connectionMode, + nodes: s.nodes, }); const EdgeRenderer = (props: EdgeRendererProps) => { @@ -247,8 +231,42 @@ const EdgeRenderer = (props: EdgeRendererProps) => { width, height, connectionMode, + nodes, } = useStore(selector, shallow); + const edges = useStore( + useCallback( + (s: ReactFlowState) => { + if (!props.onlyRenderVisibleElements) { + return s.edges; + } + + return s.edges.filter((e) => { + const { sourceNode, targetNode } = getSourceTargetNodes(e, s.nodes); + + return ( + sourceNode?.width && + sourceNode?.height && + targetNode?.width && + targetNode?.height && + isEdgeVisible({ + sourcePos: sourceNode.position, + targetPos: targetNode.position, + sourceWidth: sourceNode.width, + sourceHeight: sourceNode.height, + targetWidth: targetNode.width, + targetHeight: targetNode.height, + width: s.width, + height: s.height, + transform: s.transform, + }) + ); + }); + }, + [props.onlyRenderVisibleElements] + ) + ); + if (!width) { return null; } @@ -260,8 +278,8 @@ const EdgeRenderer = (props: EdgeRendererProps) => { - {props.edges.map((edge: Edge) => { - const { sourceNode, targetNode } = getSourceTargetNodes(edge, props.nodes); + {edges.map((edge: Edge) => { + const { sourceNode, targetNode } = getSourceTargetNodes(edge, nodes); return ( { children: ReactNode; } diff --git a/src/container/GraphView/index.tsx b/src/container/GraphView/index.tsx index 8fddc4f3..24938a74 100644 --- a/src/container/GraphView/index.tsx +++ b/src/container/GraphView/index.tsx @@ -4,14 +4,14 @@ import { useStoreApi } from '../../store'; import FlowRenderer from '../FlowRenderer'; import NodeRenderer from '../NodeRenderer'; import EdgeRenderer from '../EdgeRenderer'; -import { onLoadProject, onLoadGetElements, onLoadToObject } from '../../utils/graph'; +import { onLoadProject, onLoadGetNodes, onLoadGetEdges, onLoadToObject } from '../../utils/graph'; import useZoomPanHelper from '../../hooks/useZoomPanHelper'; import { ReactFlowProps } from '../ReactFlow'; import { NodeTypesType, EdgeTypesType, ConnectionLineType, KeyCode } from '../../types'; -export interface GraphViewProps extends Omit { +export interface GraphViewProps extends Omit { nodeTypes: NodeTypesType; edgeTypes: EdgeTypesType; selectionKeyCode: KeyCode; @@ -26,8 +26,6 @@ export interface GraphViewProps extends Omit ); diff --git a/src/container/NodeRenderer/index.tsx b/src/container/NodeRenderer/index.tsx index a808fb79..ad12d642 100644 --- a/src/container/NodeRenderer/index.tsx +++ b/src/container/NodeRenderer/index.tsx @@ -1,8 +1,9 @@ -import React, { memo, useMemo, ComponentType, MouseEvent } from 'react'; +import React, { memo, useMemo, ComponentType, MouseEvent, useCallback } from 'react'; import shallow from 'zustand/shallow'; import { useStore } from '../../store'; import { Node, NodeTypesType, ReactFlowState, WrapNodeProps } from '../../types'; +import { getNodesInside } from '../../utils/graph'; interface NodeRendererProps { nodeTypes: NodeTypesType; selectNodesOnDrag: boolean; @@ -16,7 +17,6 @@ interface NodeRendererProps { onNodeDrag?: (event: MouseEvent, node: Node) => void; onNodeDragStop?: (event: MouseEvent, node: Node) => void; onlyRenderVisibleElements: boolean; - nodes: Node[]; } const selector = (s: ReactFlowState) => ({ @@ -40,9 +40,16 @@ const NodeRenderer = (props: NodeRendererProps) => { snapToGrid, } = useStore(selector, shallow); - // const visibleNodes = props.onlyRenderVisibleElements - // ? getNodesInside(nodes, { x: 0, y: 0, width, height }, transform, true) - // : nodes; + const nodes = useStore( + useCallback( + (s: ReactFlowState) => { + return props.onlyRenderVisibleElements + ? getNodesInside(s.nodes, { x: 0, y: 0, width: s.width, height: s.height }, s.transform, true) + : s.nodes; + }, + [props.onlyRenderVisibleElements] + ) + ); const transformStyle = useMemo( () => ({ @@ -68,7 +75,7 @@ const NodeRenderer = (props: NodeRendererProps) => { return (
- {props.nodes.map((node) => { + {nodes.map((node) => { const nodeType = node.type || 'default'; const NodeComponent = (props.nodeTypes[nodeType] || props.nodeTypes.default) as ComponentType; diff --git a/src/container/ReactFlow/index.tsx b/src/container/ReactFlow/index.tsx index 3f7b4106..1e7ba986 100644 --- a/src/container/ReactFlow/index.tsx +++ b/src/container/ReactFlow/index.tsx @@ -19,7 +19,7 @@ import { BezierEdge, StepEdge, SmoothStepEdge, StraightEdge } from '../../compon import { createEdgeTypes } from '../EdgeRenderer/utils'; import Wrapper from './Wrapper'; import { - Elements, + OnSelectionChangeFunc, NodeTypesType, EdgeTypesType, OnLoadFunc, @@ -81,7 +81,7 @@ export interface ReactFlowProps extends Omit, 'on onMove?: (flowTransform?: FlowTransform) => void; onMoveStart?: (flowTransform?: FlowTransform) => void; onMoveEnd?: (flowTransform?: FlowTransform) => void; - onSelectionChange?: (elements: Elements | null) => void; + onSelectionChange?: OnSelectionChangeFunc; onSelectionDragStart?: (event: ReactMouseEvent, nodes: Node[]) => void; onSelectionDrag?: (event: ReactMouseEvent, nodes: Node[]) => void; onSelectionDragStop?: (event: ReactMouseEvent, nodes: Node[]) => void; @@ -227,8 +227,6 @@ const ReactFlow = forwardRef(
(); -const unselectElements = (elements: Elements): NodeChange[] | EdgeChange[] => - elements - .filter((e) => e.isSelected) - .map((e) => ({ - id: e.id, - type: 'select', - isSelected: false, - })); +const createNodeOrEdgeSelectionChange = (isSelected: boolean) => (item: Node | Edge) => ({ + id: item.id, + type: 'select', + isSelected, +}); const createStore = () => create((set, get) => ({ @@ -219,8 +215,12 @@ const createStore = () => const selectedEdgeIds = getConnectedEdges(selectedNodes, edges).map((e) => e.id); const selectedNodeIds = selectedNodes.map((n) => n.id); - onNodesChange?.(nodes.map((n) => ({ id: n.id, type: 'select', isSelected: selectedNodeIds.includes(n.id) }))); - onEdgesChange?.(edges.map((e) => ({ id: e.id, type: 'select', isSelected: selectedEdgeIds.includes(e.id) }))); + onNodesChange?.( + nodes.map((n) => createNodeOrEdgeSelectionChange(selectedNodeIds.includes(n.id))(n)) as NodeChange[] + ); + onEdgesChange?.( + edges.map((e) => createNodeOrEdgeSelectionChange(selectedEdgeIds.includes(e.id))(e)) as EdgeChange[] + ); set({ userSelectionRect: nextUserSelectRect, @@ -248,31 +248,22 @@ const createStore = () => set(stateUpdate); }, - addSelectedElements: (elements: Elements) => { + addSelectedElements: (selectedElementsArr: Array) => { const { multiSelectionActive, onNodesChange, onEdgesChange, nodes, edges } = get(); - const selectedElementsArr = Array.isArray(elements) ? elements : [elements]; let changedNodes; let changedEdges; if (multiSelectionActive) { - changedNodes = selectedElementsArr - .filter(isNode) - .map((node) => ({ id: node.id, type: 'select', isSelected: true })); - changedEdges = selectedElementsArr - .filter(isEdge) - .map((edge) => ({ id: edge.id, type: 'select', isSelected: true })); + changedNodes = selectedElementsArr.filter(isNode).map(createNodeOrEdgeSelectionChange(true)); + changedEdges = selectedElementsArr.filter(isEdge).map(createNodeOrEdgeSelectionChange(true)); } else { - changedNodes = nodes.map((node) => ({ - id: node.id, - type: 'select', - isSelected: selectedElementsArr.some((e) => e.id === node.id), - })); - changedEdges = edges.map((edge) => ({ - id: edge.id, - type: 'select', - isSelected: selectedElementsArr.some((e) => e.id === edge.id), - })); + changedNodes = nodes.map((node) => + createNodeOrEdgeSelectionChange(selectedElementsArr.some((e) => e.id === node.id))(node) + ); + changedEdges = edges.map((edge) => + createNodeOrEdgeSelectionChange(selectedElementsArr.some((e) => e.id === edge.id))(edge) + ); } if (changedNodes.length) { @@ -283,6 +274,21 @@ const createStore = () => onEdgesChange?.(changedEdges as EdgeChange[]); } }, + unselectNodesAndEdges: () => { + const { nodes, edges, onNodesChange, onEdgesChange } = get(); + const nodesToUnselect = nodes.map((n) => { + n.isSelected = false; + return createNodeOrEdgeSelectionChange(false)(n); + }) as NodeChange[]; + const edgesToUnselect = edges.map(createNodeOrEdgeSelectionChange(false)) as EdgeChange[]; + + if (nodesToUnselect.length) { + onNodesChange?.(nodesToUnselect); + } + if (edgesToUnselect.length) { + onEdgesChange?.(edgesToUnselect); + } + }, initD3Zoom: ({ d3Zoom, d3Selection, d3ZoomHandler, transform }: InitD3ZoomPayload) => set({ d3Zoom, @@ -312,14 +318,14 @@ const createStore = () => resetSelectedElements: () => { const { nodes, edges, onNodesChange, onEdgesChange } = get(); - const nodesToUnselect = unselectElements(nodes) as NodeChange[]; - const edgesToUnselect = unselectElements(edges) as EdgeChange[]; + const nodesToUnselect = nodes.filter((e) => e.isSelected).map(createNodeOrEdgeSelectionChange(false)); + const edgesToUnselect = edges.filter((e) => e.isSelected).map(createNodeOrEdgeSelectionChange(false)); if (nodesToUnselect.length) { - onNodesChange?.(nodesToUnselect); + onNodesChange?.(nodesToUnselect as NodeChange[]); } if (edgesToUnselect.length) { - onEdgesChange?.(edgesToUnselect); + onEdgesChange?.(edgesToUnselect as EdgeChange[]); } }, setNodeExtent: (nodeExtent: NodeExtent) => diff --git a/src/types/index.ts b/src/types/index.ts index 32f98f6e..c9751832 100644 --- a/src/types/index.ts +++ b/src/types/index.ts @@ -312,7 +312,8 @@ export type OnLoadParams = { zoomTo: (zoomLevel: number) => void; fitView: FitViewFunc; project: ProjectFunc; - getElements: () => Elements; + getNodes: () => Node[]; + getEdges: () => Edge[]; setTransform: (transform: FlowTransform) => void; toObject: ToObjectFunc; }; @@ -353,7 +354,7 @@ export type ConnectionLineComponentProps = { export type ConnectionLineComponent = React.ComponentType; -export type OnConnectFunc = (connection: Connection, nodes: Node[]) => void; +export type OnConnectFunc = (connection: Connection) => void; export type OnConnectStartParams = { nodeId: ElementId | null; handleId: ElementId | null; @@ -490,6 +491,7 @@ export interface ReactFlowState { unsetUserSelection: () => void; unsetNodesSelection: () => void; resetSelectedElements: () => void; + unselectNodesAndEdges: () => void; addSelectedElements: (elements: Elements) => void; updateTransform: (transform: Transform) => void; updateSize: (size: Dimensions) => void; @@ -522,3 +524,5 @@ export interface ReactFlowState { } export type UpdateNodeInternals = (nodeId: ElementId) => void; + +export type OnSelectionChangeFunc = (params: { nodes: Node[]; edges: Edge[] }) => void; diff --git a/src/utils/graph.ts b/src/utils/graph.ts index 109406c5..07ff0662 100644 --- a/src/utils/graph.ts +++ b/src/utils/graph.ts @@ -24,52 +24,38 @@ export const isEdge = (element: Node | Connection | Edge): element is Edge => export const isNode = (element: Node | Connection | Edge): element is Node => 'id' in element && !('source' in element) && !('target' in element); -export const getOutgoers = (node: Node, elements: Elements): Node[] => { +export const getOutgoers = (node: Node, nodes: Node[], edges: Edge[]): Node[] => { if (!isNode(node)) { return []; } - const outgoerIds = elements.filter((e) => isEdge(e) && e.source === node.id).map((e) => (e as Edge).target); - return elements.filter((e) => outgoerIds.includes(e.id)) as Node[]; + const outgoerIds = edges.filter((e) => e.source === node.id).map((e) => e.target); + return nodes.filter((n) => outgoerIds.includes(n.id)); }; -export const getIncomers = (node: Node, elements: Elements): Node[] => { +export const getIncomers = (node: Node, nodes: Node[], edges: Edge[]): Node[] => { if (!isNode(node)) { return []; } - const incomersIds = elements.filter((e) => isEdge(e) && e.target === node.id).map((e) => (e as Edge).source); - return elements.filter((e) => incomersIds.includes(e.id)) as Node[]; -}; - -export const removeElements = (elementsToRemove: Elements, elements: Elements): Elements => { - const nodeIdsToRemove = elementsToRemove.map((n) => n.id); - - return elements.filter((element) => { - const edgeElement = element as Edge; - return !( - nodeIdsToRemove.includes(element.id) || - nodeIdsToRemove.includes(edgeElement.target) || - nodeIdsToRemove.includes(edgeElement.source) - ); - }); + const incomersIds = edges.filter((e) => e.target === node.id).map((e) => e.source); + return nodes.filter((n) => incomersIds.includes(n.id)); }; const getEdgeId = ({ source, sourceHandle, target, targetHandle }: Connection): ElementId => `reactflow__edge-${source}${sourceHandle}-${target}${targetHandle}`; -const connectionExists = (edge: Edge, elements: Elements) => { - return elements.some( - (el) => - isEdge(el) && - el.source === edge.source && - el.target === edge.target && - (el.sourceHandle === edge.sourceHandle || (!el.sourceHandle && !edge.sourceHandle)) && - (el.targetHandle === edge.targetHandle || (!el.targetHandle && !edge.targetHandle)) +const connectionExists = (edge: Edge, edges: Edge[]) => { + return edges.some( + (e) => + edge.source === e.source && + edge.target === e.target && + (edge.sourceHandle === e.sourceHandle || (!edge.sourceHandle && !e.sourceHandle)) && + (edge.targetHandle === e.targetHandle || (!edge.targetHandle && !e.targetHandle)) ); }; -export const addEdge = (edgeParams: Edge | Connection, nodes: Node[], edges: Edge[]): Edge[] => { +export const addEdge = (edgeParams: Edge | Connection, edges: Edge[]): Edge[] => { if (!edgeParams.source || !edgeParams.target) { console.warn("Can't create edge. An edge needs a source and a target."); return edges; @@ -85,7 +71,7 @@ export const addEdge = (edgeParams: Edge | Connection, nodes: Node[], edges: Edg } as Edge; } - if (connectionExists(edge, nodes)) { + if (connectionExists(edge, edges)) { return edges; } @@ -211,7 +197,13 @@ export const getNodesInside = ( const yOverlap = Math.max(0, Math.min(rBox.y2, nBox.y2) - Math.max(rBox.y, nBox.y)); const overlappingArea = Math.ceil(xOverlap * yOverlap); - if (width === null || height === null || isDragging) { + if ( + typeof width === 'undefined' || + typeof height === 'undefined' || + width === null || + height === null || + isDragging + ) { // nodes are initialized with width and height = null return true; } @@ -232,15 +224,19 @@ export const getConnectedEdges = (nodes: Node[], edges: Edge[]): Edge[] => { return edges.filter((edge) => nodeIds.includes(edge.source) || nodeIds.includes(edge.target)); }; -const parseElements = (nodes: Node[], edges: Edge[]): Elements => { - return [...nodes.map((n) => ({ ...n })), ...edges.map((e) => ({ ...e }))]; +export const onLoadGetNodes = (getState: GetState) => { + return (): Node[] => { + const { nodes = [] } = getState(); + + return nodes.map((n) => ({ ...n })); + }; }; -export const onLoadGetElements = (getState: GetState) => { - return (): Elements => { - const { nodes = [], edges = [] } = getState(); +export const onLoadGetEdges = (getState: GetState) => { + return (): Edge[] => { + const { edges = [] } = getState(); - return parseElements(nodes, edges); + return edges.map((e) => ({ ...e })); }; };