diff --git a/examples/react/src/examples/Stress/index.tsx b/examples/react/src/examples/Stress/index.tsx index 7e83fac7..c1c24d86 100644 --- a/examples/react/src/examples/Stress/index.tsx +++ b/examples/react/src/examples/Stress/index.tsx @@ -1,7 +1,6 @@ import { useState, useCallback } from 'react'; import { ReactFlow, - ReactFlowInstance, Edge, Node, NodeChange, @@ -18,11 +17,6 @@ import { import { getNodesAndEdges } from './utils'; -const onInit = (reactFlowInstance: ReactFlowInstance) => { - reactFlowInstance.fitView(); - console.log(reactFlowInstance.getNodes()); -}; - const { nodes: initialNodes, edges: initialEdges } = getNodesAndEdges(25, 25); const StressFlow = () => { @@ -64,11 +58,11 @@ const StressFlow = () => { diff --git a/examples/react/src/examples/Stress/utils.ts b/examples/react/src/examples/Stress/utils.ts index 6278e006..9266bf9c 100644 --- a/examples/react/src/examples/Stress/utils.ts +++ b/examples/react/src/examples/Stress/utils.ts @@ -20,8 +20,6 @@ export function getNodesAndEdges(xElements = 10, yElements = 10): ElementsCollec style: { width: 50, height: 30, fontSize: 11 }, data, position, - width: 50, - height: 30, }; initialNodes.push(node); diff --git a/package.json b/package.json index 4411e137..274670c2 100644 --- a/package.json +++ b/package.json @@ -9,7 +9,7 @@ "preinstall": "npx only-allow pnpm", "dev": "turbo run dev --parallel --concurrency 12", "dev:svelte": "turbo run dev --filter=svelte --filter=system", - "dev:react": "turbo run dev --filter=react", + "dev:react": "turbo run dev --filter=react-examples ", "test:svelte": "pnpm --filter=playwright run test:svelte", "test:svelte:ui": "pnpm --filter=playwright run test:svelte:ui", "test:react": "pnpm --filter=playwright run test:react", diff --git a/packages/react/src/additional-components/MiniMap/MiniMap.tsx b/packages/react/src/additional-components/MiniMap/MiniMap.tsx index 49e02573..f56ae57f 100644 --- a/packages/react/src/additional-components/MiniMap/MiniMap.tsx +++ b/packages/react/src/additional-components/MiniMap/MiniMap.tsx @@ -118,7 +118,7 @@ function MiniMap({ const onSvgNodeClick = onNodeClick ? useCallback((event: MouseEvent, nodeId: string) => { - const node = store.getState().nodes.find((n) => n.id === nodeId)!; + const node = store.getState().nodeLookup.get(nodeId)!; onNodeClick(event, node); }, []) : undefined; diff --git a/packages/react/src/additional-components/NodeResizer/ResizeControl.tsx b/packages/react/src/additional-components/NodeResizer/ResizeControl.tsx index 3e038fd0..36783fcf 100644 --- a/packages/react/src/additional-components/NodeResizer/ResizeControl.tsx +++ b/packages/react/src/additional-components/NodeResizer/ResizeControl.tsx @@ -65,8 +65,8 @@ function ResizeControl({ const dragHandler = drag() .on('start', (event: ResizeDragEvent) => { - const { nodes, transform, snapGrid, snapToGrid } = store.getState(); - const node = nodes.find((n) => n.id === id); + const { nodeLookup, transform, snapGrid, snapToGrid } = store.getState(); + const node = nodeLookup.get(id); const { xSnapped, ySnapped } = getPointerPosition(event.sourceEvent, { transform, snapGrid, snapToGrid }); prevValues.current = { @@ -86,9 +86,9 @@ function ResizeControl({ onResizeStart?.(event, { ...prevValues.current }); }) .on('drag', (event: ResizeDragEvent) => { - const { nodes, transform, snapGrid, snapToGrid, triggerNodeChanges } = store.getState(); + const { nodeLookup, transform, snapGrid, snapToGrid, triggerNodeChanges } = store.getState(); const { xSnapped, ySnapped } = getPointerPosition(event.sourceEvent, { transform, snapGrid, snapToGrid }); - const node = nodes.find((n) => n.id === id); + const node = nodeLookup.get(id); if (node) { const changes: NodeChange[] = []; diff --git a/packages/react/src/additional-components/NodeToolbar/NodeToolbar.tsx b/packages/react/src/additional-components/NodeToolbar/NodeToolbar.tsx index fc0c309f..ad581bcf 100644 --- a/packages/react/src/additional-components/NodeToolbar/NodeToolbar.tsx +++ b/packages/react/src/additional-components/NodeToolbar/NodeToolbar.tsx @@ -87,12 +87,12 @@ function NodeToolbar({ const nodeIds = Array.isArray(nodeId) ? nodeId : [nodeId || contextNodeId || '']; return nodeIds.reduce((acc, id) => { - const node = state.nodes.find((n) => n.id === id); + const node = state.nodeLookup.get(id); if (node) { acc.push(node); } return acc; - }, [] as Node[]); + }, []); }, [nodeId, contextNodeId] ); diff --git a/packages/react/src/components/ConnectionLine/index.tsx b/packages/react/src/components/ConnectionLine/index.tsx index 439cc342..ed1a2115 100644 --- a/packages/react/src/components/ConnectionLine/index.tsx +++ b/packages/react/src/components/ConnectionLine/index.tsx @@ -43,7 +43,7 @@ const ConnectionLine = ({ const { fromNode, handleId, toX, toY, connectionMode } = useStore( useCallback( (s: ReactFlowStore) => ({ - fromNode: s.nodes.find((n) => n.id === nodeId), + fromNode: s.nodeLookup.get(nodeId), handleId: s.connectionStartHandle?.handleId, toX: (s.connectionPosition.x - s.transform[0]) / s.transform[2], toY: (s.connectionPosition.y - s.transform[1]) / s.transform[2], diff --git a/packages/react/src/components/Edges/BaseEdge.tsx b/packages/react/src/components/Edges/BaseEdge.tsx index dd496f7e..91cfd58a 100644 --- a/packages/react/src/components/Edges/BaseEdge.tsx +++ b/packages/react/src/components/Edges/BaseEdge.tsx @@ -1,7 +1,6 @@ import { isNumeric } from '@xyflow/system'; import type { BaseEdgeProps } from '../../types'; - import EdgeText from './EdgeText'; const BaseEdge = ({ diff --git a/packages/react/src/components/Edges/wrapEdge.tsx b/packages/react/src/components/Edges/wrapEdge.tsx index 750254ff..b0a876b0 100644 --- a/packages/react/src/components/Edges/wrapEdge.tsx +++ b/packages/react/src/components/Edges/wrapEdge.tsx @@ -1,4 +1,4 @@ -import { memo, useState, useMemo, useRef, type ComponentType, type KeyboardEvent } from 'react'; +import { memo, useState, useMemo, useRef, type ComponentType, type KeyboardEvent, useCallback } from 'react'; import cc from 'classcat'; import { shallow } from 'zustand/shallow'; import { getMarkerId, elementSelectionKeys, XYHandle, type Connection, getEdgePosition } from '@xyflow/system'; @@ -53,26 +53,30 @@ export default (EdgeComponent: ComponentType) => { const [updateHover, setUpdateHover] = useState(false); const [updating, setUpdating] = useState(false); const store = useStoreApi(); - const edgePosition = useStore((state) => { - const sourceNode = state.nodes.find((n) => n.id === source); - const targetNode = state.nodes.find((n) => n.id === target); + const edgePosition = useStore( + useCallback( + (state) => { + const sourceNode = state.nodeLookup.get(source); + const targetNode = state.nodeLookup.get(target); - if (!sourceNode || !targetNode) { - return null; - } + if (!sourceNode || !targetNode) { + return null; + } - const pos = getEdgePosition({ - id, - sourceNode, - targetNode, - sourceHandle: sourceHandleId || null, - targetHandle: targetHandleId || null, - connectionMode: state.connectionMode, - onError: state.onError, - }); - - return pos; - }, shallow); + return getEdgePosition({ + id, + sourceNode, + targetNode, + sourceHandle: sourceHandleId || null, + targetHandle: targetHandleId || null, + connectionMode: state.connectionMode, + onError: state.onError, + }); + }, + [source, target] + ), + shallow + ); const markerStartUrl = useMemo(() => `url(#${getMarkerId(markerStart, rfId)})`, [markerStart, rfId]); const markerEndUrl = useMemo(() => `url(#${getMarkerId(markerEnd, rfId)})`, [markerEnd, rfId]); diff --git a/packages/react/src/components/Handle/index.tsx b/packages/react/src/components/Handle/index.tsx index 925bc93e..29231b71 100644 --- a/packages/react/src/components/Handle/index.tsx +++ b/packages/react/src/components/Handle/index.tsx @@ -1,3 +1,8 @@ +/* + * The Handle component is used to connect nodes. When the user mousedowns a handle, we start the connection process. + * The user can then drag the connection to another handle or node. When the user releases the mouse, we check if the + * connection is valid and if so, we call the onConnect callback. + */ import { memo, HTMLAttributes, forwardRef, MouseEvent as ReactMouseEvent, TouchEvent as ReactTouchEvent } from 'react'; import cc from 'classcat'; import { shallow } from 'zustand/shallow'; diff --git a/packages/react/src/components/Nodes/utils.ts b/packages/react/src/components/Nodes/utils.ts index a70ccd34..b9a82189 100644 --- a/packages/react/src/components/Nodes/utils.ts +++ b/packages/react/src/components/Nodes/utils.ts @@ -12,7 +12,7 @@ export function getMouseHandler( return handler === undefined ? handler : (event: MouseEvent) => { - const node = getState().nodes.find((n) => n.id === id)!; + const node = getState().nodeLookup.get(id)!; handler(event, { ...node }); }; } @@ -35,8 +35,8 @@ export function handleNodeClick({ unselect?: boolean; nodeRef?: RefObject; }) { - const { addSelectedNodes, unselectNodesAndEdges, multiSelectionActive, nodes, onError } = store.getState(); - const node = nodes.find((n) => n.id === id)!; + const { addSelectedNodes, unselectNodesAndEdges, multiSelectionActive, nodeLookup, onError } = store.getState(); + const node = nodeLookup.get(id); if (!node) { onError?.('012', errorMessages['error012'](id)); diff --git a/packages/react/src/components/Nodes/wrapNode.tsx b/packages/react/src/components/Nodes/wrapNode.tsx index f96361ce..e50ffcf7 100644 --- a/packages/react/src/components/Nodes/wrapNode.tsx +++ b/packages/react/src/components/Nodes/wrapNode.tsx @@ -146,7 +146,7 @@ export default (NodeComponent: ComponentType) => { if (targetPosChanged) { prevTargetPosition.current = targetPosition; } - store.getState().updateNodeDimensions([{ id, nodeElement: nodeRef.current, forceUpdate: true }]); + store.getState().updateNodeDimensions(new Map([[id, { id, nodeElement: nodeRef.current, forceUpdate: true }]])); } }, [id, type, sourcePosition, targetPosition]); diff --git a/packages/react/src/components/SelectionListener/index.tsx b/packages/react/src/components/SelectionListener/index.tsx index b71231f6..4ef1e130 100644 --- a/packages/react/src/components/SelectionListener/index.tsx +++ b/packages/react/src/components/SelectionListener/index.tsx @@ -1,3 +1,9 @@ +/* + * This is a helper component for calling the onSelectionChange listener. + * It will only be mounted if the user has passed an onSelectionChange listener + * or is using the useOnSelectionChange hook. + * @TODO: Now that we have the onNodesChange and on EdgesChange listeners, do we still need this component? + */ import { memo, useEffect } from 'react'; import { shallow } from 'zustand/shallow'; @@ -24,8 +30,6 @@ function areEqual(a: SelectorSlice, b: SelectorSlice) { ); } -// This is just a helper component for calling the onSelectionChange listener. -// @TODO: Now that we have the onNodesChange and on EdgesChange listeners, do we still need this component? const SelectionListener = memo(({ onSelectionChange }: SelectionListenerProps) => { const store = useStoreApi(); const { selectedNodes, selectedEdges } = useStore(selector, areEqual); diff --git a/packages/react/src/components/StoreUpdater/index.tsx b/packages/react/src/components/StoreUpdater/index.tsx index a0ce26dd..05a8aa50 100644 --- a/packages/react/src/components/StoreUpdater/index.tsx +++ b/packages/react/src/components/StoreUpdater/index.tsx @@ -1,3 +1,8 @@ +/* + * This component helps us to update the store with the vlues coming from the user. + * We distinguish between values we can update directly with `useDirectStoreUpdater` (like `snapGrid`) + * and values that have a dedicated setter function in the store (like `setNodes`). + */ import { useEffect } from 'react'; import { StoreApi } from 'zustand'; import { shallow } from 'zustand/shallow'; @@ -70,10 +75,10 @@ const selector = (s: ReactFlowState) => ({ reset: s.reset, }); -function useStoreUpdater(value: T | undefined, setStoreState: (param: T) => void) { +function useStoreUpdater(value: T | undefined, setStoreAction: (param: T) => void) { useEffect(() => { if (typeof value !== 'undefined') { - setStoreState(value); + setStoreAction(value); } }, [value]); } diff --git a/packages/react/src/container/EdgeRenderer/index.tsx b/packages/react/src/container/EdgeRenderer/index.tsx index fe2de112..24825fe6 100644 --- a/packages/react/src/container/EdgeRenderer/index.tsx +++ b/packages/react/src/container/EdgeRenderer/index.tsx @@ -63,6 +63,8 @@ const EdgeRenderer = ({ children, }: EdgeRendererProps) => { const { edgesFocusable, edgesUpdatable, elementsSelectable, onError } = useStore(selector, shallow); + // we are grouping edges by zIndex here in order to be able to render them in the correct order + // each zIndex gets its own svg element const edgeTree = useVisibleEdges(onlyRenderVisibleElements, elevateEdgesOnSelect); return ( @@ -70,7 +72,7 @@ const EdgeRenderer = ({ {edgeTree.map(({ level, edges, isMaxLevel }) => ( {isMaxLevel && } - + <> {edges.map((edge) => { let edgeType = edge.type || 'default'; @@ -132,7 +134,7 @@ const EdgeRenderer = ({ /> ); })} - + ))} {children} diff --git a/packages/react/src/container/NodeRenderer/index.tsx b/packages/react/src/container/NodeRenderer/index.tsx index f2b3ebe3..98f9fc39 100644 --- a/packages/react/src/container/NodeRenderer/index.tsx +++ b/packages/react/src/container/NodeRenderer/index.tsx @@ -48,11 +48,16 @@ const NodeRenderer = (props: NodeRendererProps) => { } const observer = new ResizeObserver((entries: ResizeObserverEntry[]) => { - const updates = entries.map((entry: ResizeObserverEntry) => ({ - id: entry.target.getAttribute('data-id') as string, - nodeElement: entry.target as HTMLDivElement, - forceUpdate: true, - })); + const updates = new Map(); + + entries.forEach((entry: ResizeObserverEntry) => { + const id = entry.target.getAttribute('data-id') as string; + updates.set(id, { + id, + nodeElement: entry.target as HTMLDivElement, + forceUpdate: true, + }); + }); updateNodeDimensions(updates); }); diff --git a/packages/react/src/hooks/useReactFlow.ts b/packages/react/src/hooks/useReactFlow.ts index 301f49f5..e72c3a4f 100644 --- a/packages/react/src/hooks/useReactFlow.ts +++ b/packages/react/src/hooks/useReactFlow.ts @@ -35,7 +35,7 @@ export default function useReactFlow(): ReactFlo }, []); const getNode = useCallback>((id) => { - return store.getState().nodes.find((n) => n.id === id); + return store.getState().nodeLookup.get(id); }, []); const getEdges = useCallback>(() => { @@ -181,7 +181,7 @@ export default function useReactFlow(): ReactFlo nodeOrRect: Node | { id: Node['id'] } | Rect ): [Rect | null, Node | null | undefined, boolean] => { const isRect = isRectObject(nodeOrRect); - const node = isRect ? null : store.getState().nodes.find((n) => n.id === nodeOrRect.id); + const node = isRect ? null : store.getState().nodeLookup.get(nodeOrRect.id); if (!isRect && !node) { [null, null, isRect]; diff --git a/packages/react/src/hooks/useUpdateNodeInternals.ts b/packages/react/src/hooks/useUpdateNodeInternals.ts index 9b4bcf61..d14f35f3 100644 --- a/packages/react/src/hooks/useUpdateNodeInternals.ts +++ b/packages/react/src/hooks/useUpdateNodeInternals.ts @@ -8,17 +8,16 @@ function useUpdateNodeInternals(): UpdateNodeInternals { return useCallback((id: string | string[]) => { const { domNode, updateNodeDimensions } = store.getState(); - const updateIds = Array.isArray(id) ? id : [id]; - const updates = updateIds.reduce((res, updateId) => { + const updates = new Map(); + + updateIds.forEach((updateId) => { const nodeElement = domNode?.querySelector(`.react-flow__node[data-id="${updateId}"]`) as HTMLDivElement; if (nodeElement) { - res.push({ id: updateId, nodeElement, forceUpdate: true }); + updates.set(updateId, { id: updateId, nodeElement, forceUpdate: true }); } - - return res; - }, []); + }); requestAnimationFrame(() => updateNodeDimensions(updates)); }, []); diff --git a/packages/react/src/hooks/useVisibleEdges.ts b/packages/react/src/hooks/useVisibleEdges.ts index 6d7f808a..67d02c92 100644 --- a/packages/react/src/hooks/useVisibleEdges.ts +++ b/packages/react/src/hooks/useVisibleEdges.ts @@ -12,8 +12,8 @@ function useVisibleEdges(onlyRenderVisible: boolean, elevateEdgesOnSelect: boole const visibleEdges = onlyRenderVisible && s.width && s.height ? s.edges.filter((e) => { - const sourceNode = s.nodes.find((n) => n.id === e.source); - const targetNode = s.nodes.find((n) => n.id === e.target); + const sourceNode = s.nodeLookup.get(e.source); + const targetNode = s.nodeLookup.get(e.target); return ( sourceNode && @@ -29,7 +29,7 @@ function useVisibleEdges(onlyRenderVisible: boolean, elevateEdgesOnSelect: boole }) : s.edges; - return groupEdgesByZLevel(visibleEdges, s.nodes, elevateEdgesOnSelect); + return groupEdgesByZLevel(visibleEdges, s.nodeLookup, elevateEdgesOnSelect); }, [onlyRenderVisible, elevateEdgesOnSelect] ), diff --git a/packages/react/src/store/index.ts b/packages/react/src/store/index.ts index ad42fd6e..a6898f6d 100644 --- a/packages/react/src/store/index.ts +++ b/packages/react/src/store/index.ts @@ -41,18 +41,19 @@ const createRFStore = ({ (set, get) => ({ ...getInitialState({ nodes, edges, width, height, fitView }), setNodes: (nodes: Node[]) => { - const { nodes: storeNodes, nodeOrigin, elevateNodesOnSelect } = get(); - const nextNodes = updateNodes(nodes, storeNodes, { nodeOrigin, elevateNodesOnSelect }); + const { nodeLookup, nodeOrigin, elevateNodesOnSelect } = get(); + // Whenver new nodes are set, we need to calculate the absolute positions of the nodes + // and update the nodeLookup. + const nextNodes = updateNodes(nodes, nodeLookup, { nodeOrigin, elevateNodesOnSelect }); set({ nodes: nextNodes }); }, - getNodes: () => { - return get().nodes; - }, setEdges: (edges: Edge[]) => { const { defaultEdgeOptions = {} } = get(); set({ edges: edges.map((e) => ({ ...defaultEdgeOptions, ...e })) }); }, + // when the user works with an uncontrolled flow, + // we set a flag `hasDefaultNodes` / `hasDefaultEdges` setDefaultNodesAndEdges: (nodes?: Node[], edges?: Edge[]) => { const hasDefaultNodes = typeof nodes !== 'undefined'; const hasDefaultEdges = typeof edges !== 'undefined'; @@ -68,7 +69,7 @@ const createRFStore = ({ }; if (hasDefaultNodes) { - nextState.nodes = updateNodes(nodes, [], { + nextState.nodes = updateNodes(nodes, new Map(), { nodeOrigin: get().nodeOrigin, elevateNodesOnSelect: get().elevateNodesOnSelect, }); @@ -79,14 +80,27 @@ const createRFStore = ({ set(nextState); }, + // Every node gets registerd at a ResizeObserver. Whenever a node + // changes its dimensions, this function is called to measure the + // new dimensions and update the nodes. updateNodeDimensions: (updates) => { - const { onNodesChange, fitView, nodes, fitViewOnInit, fitViewDone, fitViewOnInitOptions, domNode, nodeOrigin } = - get(); + const { + onNodesChange, + fitView, + nodes, + nodeLookup, + fitViewOnInit, + fitViewDone, + fitViewOnInitOptions, + domNode, + nodeOrigin, + } = get(); const changes: NodeDimensionChange[] = []; const updatedNodes = updateNodeDimensionsSystem( updates, nodes, + nodeLookup, domNode, nodeOrigin, (id: string, dimensions: Dimensions) => { @@ -102,8 +116,9 @@ const createRFStore = ({ return; } - const nextNodes = updateAbsolutePositions(updatedNodes, nodeOrigin); + const nextNodes = updateAbsolutePositions(updatedNodes, nodeLookup, nodeOrigin); + // we call fitView once initially after all dimensions are set let nextFitViewDone = fitViewDone; if (!fitViewDone && fitViewOnInit) { nextFitViewDone = fitView(nextNodes, { @@ -112,6 +127,11 @@ const createRFStore = ({ }); } + // here we are cirmumventing the onNodesChange handler + // in order to be able to display nodes even if the user + // has not provided an onNodesChange handler. + // Nodes are only rendered if they have a width and height + // attribute which they get from this handler. set({ nodes: nextNodes, fitViewDone: nextFitViewDone }); if (changes?.length > 0) { @@ -138,12 +158,12 @@ const createRFStore = ({ }, triggerNodeChanges: (changes) => { - const { onNodesChange, nodes, hasDefaultNodes, nodeOrigin, elevateNodesOnSelect } = get(); + const { onNodesChange, nodeLookup, nodes, hasDefaultNodes, nodeOrigin, elevateNodesOnSelect } = get(); if (changes?.length) { if (hasDefaultNodes) { const updatedNodes = applyNodeChanges(changes, nodes); - const nextNodes = updateNodes(updatedNodes, nodes, { + const nextNodes = updateNodes(updatedNodes, nodeLookup, { nodeOrigin, elevateNodesOnSelect, }); diff --git a/packages/react/src/store/initialState.ts b/packages/react/src/store/initialState.ts index 8f8fa450..c87c29f1 100644 --- a/packages/react/src/store/initialState.ts +++ b/packages/react/src/store/initialState.ts @@ -22,7 +22,8 @@ const getInitialState = ({ height?: number; fitView?: boolean; } = {}): ReactFlowStore => { - const nextNodes = updateNodes(nodes, [], { nodeOrigin: [0, 0], elevateNodesOnSelect: false }); + const nodeLookup = new Map(); + const nextNodes = updateNodes(nodes, nodeLookup, { nodeOrigin: [0, 0], elevateNodesOnSelect: false }); let transform: Transform = [0, 0, 1]; @@ -43,6 +44,7 @@ const getInitialState = ({ height: 0, transform, nodes: nextNodes, + nodeLookup, edges: edges, onNodesChange: null, onEdgesChange: null, diff --git a/packages/react/src/types/store.ts b/packages/react/src/types/store.ts index 3662f862..02a7d162 100644 --- a/packages/react/src/types/store.ts +++ b/packages/react/src/types/store.ts @@ -46,6 +46,7 @@ export type ReactFlowStore = { height: number; transform: Transform; nodes: Node[]; + nodeLookup: Map; edges: Edge[]; onNodesChange: OnNodesChange | null; onEdgesChange: OnEdgesChange | null; @@ -138,10 +139,9 @@ export type ReactFlowStore = { export type ReactFlowActions = { setNodes: (nodes: Node[]) => void; - getNodes: () => Node[]; setEdges: (edges: Edge[]) => void; setDefaultNodesAndEdges: (nodes?: Node[], edges?: Edge[]) => void; - updateNodeDimensions: (updates: NodeDimensionUpdate[]) => void; + updateNodeDimensions: (updates: Map) => void; updateNodePositions: UpdateNodePositions; resetSelectedElements: () => void; unselectNodesAndEdges: (params?: UnselectNodesAndEdgesParams) => void; diff --git a/packages/react/src/utils/changes.ts b/packages/react/src/utils/changes.ts index a51eaba4..61ff2b19 100644 --- a/packages/react/src/utils/changes.ts +++ b/packages/react/src/utils/changes.ts @@ -42,24 +42,39 @@ export function handleParentExpand(res: any[], updateItem: any) { } } +// This function applies changes to nodes or edges that are triggered by React Flow internally. +// When you drag a node for example, React Flow will send a position change update. +// This function then applies the changes and returns the updated elements. function applyChanges(changes: any[], elements: any[]): any[] { // we need this hack to handle the setNodes and setEdges function of the useReactFlow hook for controlled flows if (changes.some((c) => c.type === 'reset')) { return changes.filter((c) => c.type === 'reset').map((c) => c.item); } + + let remainingChanges = changes; const initElements: any[] = changes.filter((c) => c.type === 'add').map((c) => c.item); return elements.reduce((res: any[], item: any) => { - const currentChanges = changes.filter((c) => c.id === item.id); + const nextChanges: any[] = []; + const _remainingChanges: any[] = []; - if (currentChanges.length === 0) { + remainingChanges.forEach((c) => { + if (c.id === item.id) { + nextChanges.push(c); + } else { + _remainingChanges.push(c); + } + }); + remainingChanges = _remainingChanges; + + if (nextChanges.length === 0) { res.push(item); return res; } const updateItem = { ...item }; - for (const currentChange of currentChanges) { + for (const currentChange of nextChanges) { if (currentChange) { switch (currentChange.type) { case 'select': { diff --git a/packages/svelte/src/lib/actions/drag/index.ts b/packages/svelte/src/lib/actions/drag/index.ts index e936d3a5..d08a17be 100644 --- a/packages/svelte/src/lib/actions/drag/index.ts +++ b/packages/svelte/src/lib/actions/drag/index.ts @@ -30,6 +30,7 @@ export default function drag(domNode: Element, params: UseDragParams) { return { nodes: get(store.nodes), + nodeLookup: get(store.nodeLookup), edges: get(store.edges), nodeExtent: get(store.nodeExtent), snapGrid: snapGrid ? snapGrid : [0, 0], diff --git a/packages/svelte/src/lib/components/BaseEdge/BaseEdge.svelte b/packages/svelte/src/lib/components/BaseEdge/BaseEdge.svelte index 4e5b39d1..d1c60f8f 100644 --- a/packages/svelte/src/lib/components/BaseEdge/BaseEdge.svelte +++ b/packages/svelte/src/lib/components/BaseEdge/BaseEdge.svelte @@ -23,6 +23,7 @@ { - const updates = entries.map((entry: ResizeObserverEntry) => ({ - id: entry.target.getAttribute('data-id') as string, - nodeElement: entry.target as HTMLDivElement, - forceUpdate: true - })); + const updates = new Map(); + + entries.forEach((entry: ResizeObserverEntry) => { + const id = entry.target.getAttribute('data-id') as string; + + updates.set(id, { + id, + nodeElement: entry.target as HTMLDivElement, + forceUpdate: true + }); + }); + updateNodeDimensions(updates); }); diff --git a/packages/svelte/src/lib/hooks/useUpdateNodeInternals.ts b/packages/svelte/src/lib/hooks/useUpdateNodeInternals.ts index 9de90763..c66bb688 100644 --- a/packages/svelte/src/lib/hooks/useUpdateNodeInternals.ts +++ b/packages/svelte/src/lib/hooks/useUpdateNodeInternals.ts @@ -9,17 +9,17 @@ export function useUpdateNodeInternals(): UpdateNodeInternals { // @todo: do we want to add this to system? const updateInternals = (id: string | string[]) => { const updateIds = Array.isArray(id) ? id : [id]; - const updates = updateIds.reduce((res, updateId) => { + const updates = new Map(); + + updateIds.forEach((updateId) => { const nodeElement = get(domNode)?.querySelector( `.svelte-flow__node[data-id="${updateId}"]` ) as HTMLDivElement; if (nodeElement) { - res.push({ id: updateId, nodeElement, forceUpdate: true }); + updates.set(updateId, { id: updateId, nodeElement, forceUpdate: true }); } - - return res; - }, []); + }); requestAnimationFrame(() => updateNodeDimensions(updates)); }; diff --git a/packages/svelte/src/lib/store/derived-connection-props.ts b/packages/svelte/src/lib/store/derived-connection-props.ts index 5be1598c..dd4e805c 100644 --- a/packages/svelte/src/lib/store/derived-connection-props.ts +++ b/packages/svelte/src/lib/store/derived-connection-props.ts @@ -56,15 +56,15 @@ export function getDerivedConnectionProps( currentConnection, store.connectionLineType, store.connectionMode, - store.nodes, + store.nodeLookup, store.viewport ], - ([connection, connectionLineType, connectionMode, nodes, viewport]) => { + ([connection, connectionLineType, connectionMode, nodeLookup, viewport]) => { if (!connection.connectionStartHandle?.nodeId) { return initConnectionProps; } - const fromNode = nodes.find((n) => n.id === connection.connectionStartHandle?.nodeId); + const fromNode = nodeLookup.get(connection.connectionStartHandle?.nodeId); const fromHandleBounds = fromNode?.[internalsSymbol]?.handleBounds; const handleBoundsStrict = fromHandleBounds?.[connection.connectionStartHandle.type || 'source'] || []; diff --git a/packages/svelte/src/lib/store/edge-tree.ts b/packages/svelte/src/lib/store/edge-tree.ts index ff099868..0861be01 100644 --- a/packages/svelte/src/lib/store/edge-tree.ts +++ b/packages/svelte/src/lib/store/edge-tree.ts @@ -9,17 +9,18 @@ export function getEdgeTree(store: SvelteFlowStoreState) { [ store.edges, store.nodes, + store.nodeLookup, store.onlyRenderVisibleElements, store.viewport, store.width, store.height ], - ([edges, nodes, onlyRenderVisibleElements, viewport, width, height]) => { + ([edges, , nodeLookup, onlyRenderVisibleElements, viewport, width, height]) => { const visibleEdges = onlyRenderVisibleElements && width && height ? edges.filter((edge) => { - const sourceNode = nodes.find((node) => node.id === edge.source); - const targetNode = nodes.find((node) => node.id === edge.target); + const sourceNode = nodeLookup.get(edge.source); + const targetNode = nodeLookup.get(edge.target); return ( sourceNode && @@ -40,11 +41,11 @@ export function getEdgeTree(store: SvelteFlowStoreState) { ); return derived( - [visibleEdges, store.nodes, store.connectionMode, store.onError], - ([visibleEdges, nodes, connectionMode, onError]) => { + [visibleEdges, store.nodes, store.nodeLookup, store.connectionMode, store.onError], + ([visibleEdges, , nodeLookup, connectionMode, onError]) => { const layoutedEdges = visibleEdges.reduce((res, edge) => { - const sourceNode = nodes.find((node) => node.id === edge.source); - const targetNode = nodes.find((node) => node.id === edge.target); + const sourceNode = nodeLookup.get(edge.source); + const targetNode = nodeLookup.get(edge.target); if (!sourceNode || !targetNode) { return res; @@ -70,7 +71,7 @@ export function getEdgeTree(store: SvelteFlowStoreState) { return res; }, []); - const groupedEdges = groupEdgesByZLevel(layoutedEdges, nodes, false); + const groupedEdges = groupEdgesByZLevel(layoutedEdges, nodeLookup, false); return groupedEdges; } diff --git a/packages/svelte/src/lib/store/index.ts b/packages/svelte/src/lib/store/index.ts index b78e6fb2..27a6aee0 100644 --- a/packages/svelte/src/lib/store/index.ts +++ b/packages/svelte/src/lib/store/index.ts @@ -86,10 +86,11 @@ export function createStore({ }); }; - function updateNodeDimensions(updates: NodeDimensionUpdate[]) { + function updateNodeDimensions(updates: Map) { const nextNodes = updateNodeDimensionsSystem( updates, get(store.nodes), + get(store.nodeLookup), get(store.domNode), get(store.nodeOrigin) ); diff --git a/packages/svelte/src/lib/store/initial-store.ts b/packages/svelte/src/lib/store/initial-store.ts index 7744ac0c..4bb3d7af 100644 --- a/packages/svelte/src/lib/store/initial-store.ts +++ b/packages/svelte/src/lib/store/initial-store.ts @@ -59,7 +59,11 @@ export const getInitialStore = ({ height?: number; fitView?: boolean; }) => { - const nextNodes = updateNodes(nodes, [], { nodeOrigin: [0, 0], elevateNodesOnSelect: false }); + const nodeLookup = new Map(); + const nextNodes = updateNodes(nodes, nodeLookup, { + nodeOrigin: [0, 0], + elevateNodesOnSelect: false + }); let viewport: Viewport = { x: 0, y: 0, zoom: 1 }; @@ -75,7 +79,8 @@ export const getInitialStore = ({ return { flowId: writable(null), - nodes: createNodesStore(nextNodes), + nodes: createNodesStore(nextNodes, nodeLookup), + nodeLookup: readable>(nodeLookup), visibleNodes: readable([]), edges: createEdgesStore(edges), edgeTree: readable[]>([]), diff --git a/packages/svelte/src/lib/store/types.ts b/packages/svelte/src/lib/store/types.ts index ae854cfd..8d41d63f 100644 --- a/packages/svelte/src/lib/store/types.ts +++ b/packages/svelte/src/lib/store/types.ts @@ -27,7 +27,7 @@ export type SvelteFlowStoreActions = { setTranslateExtent: (extent: CoordinateExtent) => void; fitView: (options?: FitViewOptions) => boolean; updateNodePositions: UpdateNodePositions; - updateNodeDimensions: (updates: NodeDimensionUpdate[]) => void; + updateNodeDimensions: (updates: Map) => void; unselectNodesAndEdges: (params?: { nodes?: Node[]; edges?: Edge[] }) => void; addSelectedNodes: (ids: string[]) => void; addSelectedEdges: (ids: string[]) => void; diff --git a/packages/svelte/src/lib/store/utils.ts b/packages/svelte/src/lib/store/utils.ts index cfd5b343..67c6d288 100644 --- a/packages/svelte/src/lib/store/utils.ts +++ b/packages/svelte/src/lib/store/utils.ts @@ -111,7 +111,8 @@ export type NodeStoreOptions = { // we are creating a custom store for the internals nodes in order to update the zIndex and positionAbsolute. // The user only passes in relative positions, so we need to calculate the absolute positions based on the parent nodes. export const createNodesStore = ( - nodes: Node[] + nodes: Node[], + nodeLookup: Map ): { subscribe: (this: void, run: Subscriber) => Unsubscriber; update: (this: void, updater: Updater) => void; @@ -125,7 +126,7 @@ export const createNodesStore = ( let elevateNodesOnSelect = true; const _set = (nds: Node[]): Node[] => { - const nextNodes = updateNodes(nds, value, { + const nextNodes = updateNodes(nds, nodeLookup, { elevateNodesOnSelect, defaults }); diff --git a/packages/system/src/utils/dom.ts b/packages/system/src/utils/dom.ts index 4aa51174..74bb14aa 100644 --- a/packages/system/src/utils/dom.ts +++ b/packages/system/src/utils/dom.ts @@ -13,7 +13,6 @@ export function getPointerPosition( ): XYPosition & { xSnapped: number; ySnapped: number } { const { x, y } = getEventPosition(event); const pointerPos = pointToRendererPoint({ x, y }, transform); - const { x: xSnapped, y: ySnapped } = snapToGrid ? snapPosition(pointerPos, snapGrid) : pointerPos; // we need the snapped position in order to be able to skip unnecessary drag events @@ -58,6 +57,9 @@ export const getEventPosition = (event: MouseEvent | TouchEvent, bounds?: DOMRec }; }; +// The handle bounds are calculated relative to the node element. +// We store them in the internals object of the node in order to avoid +// unnecessary recalculations. export const getHandleBounds = ( selector: string, nodeElement: HTMLDivElement, diff --git a/packages/system/src/utils/edges/general.ts b/packages/system/src/utils/edges/general.ts index e15d9269..bdffbbd8 100644 --- a/packages/system/src/utils/edges/general.ts +++ b/packages/system/src/utils/edges/general.ts @@ -33,7 +33,7 @@ export type GroupedEdges = { export function groupEdgesByZLevel( edges: EdgeType[], - nodes: NodeBase[], + nodeLookup: Map, elevateEdgesOnSelect = false ): GroupedEdges[] { let maxLevel = -1; @@ -43,8 +43,8 @@ export function groupEdgesByZLevel( let z = hasZIndex ? edge.zIndex! : 0; if (elevateEdgesOnSelect) { - const targetNode = nodes.find((n) => n.id === edge.target); - const sourceNode = nodes.find((n) => n.id === edge.source); + const targetNode = nodeLookup.get(edge.target); + const sourceNode = nodeLookup.get(edge.source); const edgeOrConnectedNodeSelected = edge.selected || targetNode?.selected || sourceNode?.selected; const selectedZIndex = Math.max( sourceNode?.[internalsSymbol]?.z || 0, diff --git a/packages/system/src/utils/store.ts b/packages/system/src/utils/store.ts index 3fee5d02..9ed4f871 100644 --- a/packages/system/src/utils/store.ts +++ b/packages/system/src/utils/store.ts @@ -18,19 +18,21 @@ type ParentNodes = Record; export function updateAbsolutePositions( nodes: NodeType[], + nodeLookup: Map, nodeOrigin: NodeOrigin = [0, 0], parentNodes?: ParentNodes ) { return nodes.map((node) => { - if (node.parentNode && !nodes.find((n) => n.id === node.parentNode)) { + if (node.parentNode && !nodeLookup.has(node.parentNode)) { throw new Error(`Parent node ${node.parentNode} not found`); } if (node.parentNode || parentNodes?.[node.id]) { - const parentNode = node.parentNode ? nodes.find((n) => n.id === node.parentNode) : null; + const parentNode = node.parentNode ? nodeLookup.get(node.parentNode) : null; const { x, y, z } = calculateXYZPosition( node, nodes, + nodeLookup, { ...node.position, z: node[internalsSymbol]?.z ?? 0, @@ -62,7 +64,7 @@ type UpdateNodesOptions = { export function updateNodes( nodes: NodeType[], - storeNodes: NodeType[], + nodeLookup: Map, options: UpdateNodesOptions = { nodeOrigin: [0, 0] as NodeOrigin, elevateNodesOnSelect: true, @@ -73,7 +75,7 @@ export function updateNodes( const selectedNodeZ: number = options?.elevateNodesOnSelect ? 1000 : 0; const nextNodes = nodes.map((n) => { - const currentStoreNode = storeNodes.find((storeNode) => n.id === storeNode.id); + const currentStoreNode = nodeLookup.get(n.id); const node: NodeType = { ...options.defaults, ...n, @@ -96,10 +98,12 @@ export function updateNodes( }, }); + nodeLookup.set(node.id, node); + return node; }); - const nodesWithPositions = updateAbsolutePositions(nextNodes, options.nodeOrigin, parentNodes); + const nodesWithPositions = updateAbsolutePositions(nextNodes, nodeLookup, options.nodeOrigin, parentNodes); return nodesWithPositions; } @@ -107,6 +111,7 @@ export function updateNodes( function calculateXYZPosition( node: NodeType, nodes: NodeType[], + nodeLookup: Map, result: XYZPosition, nodeOrigin: NodeOrigin ): XYZPosition { @@ -114,12 +119,13 @@ function calculateXYZPosition( return result; } - const parentNode = nodes.find((n) => n.id === node.parentNode)!; + const parentNode = nodeLookup.get(node.parentNode)!; const parentNodePosition = getNodePositionWithOrigin(parentNode, parentNode?.origin || nodeOrigin); return calculateXYZPosition( parentNode, nodes, + nodeLookup, { x: (result.x ?? 0) + parentNodePosition.x, y: (result.y ?? 0) + parentNodePosition.y, @@ -130,8 +136,9 @@ function calculateXYZPosition( } export function updateNodeDimensions( - updates: NodeDimensionUpdate[], + updates: Map, nodes: NodeBase[], + nodeLookup: Map, domNode: HTMLElement | null, nodeOrigin?: NodeOrigin, onUpdate?: (id: string, dimensions: Dimensions) => void @@ -146,7 +153,8 @@ export function updateNodeDimensions( const { m22: zoom } = new window.DOMMatrixReadOnly(style.transform); const nextNodes = nodes.map((node) => { - const update = updates.find((u) => u.id === node.id); + const update = updates.get(node.id); + if (update) { const dimensions = getDimensions(update.nodeElement); const doUpdate = !!( @@ -158,7 +166,7 @@ export function updateNodeDimensions( if (doUpdate) { onUpdate?.(node.id, dimensions); - return { + const newNode = { ...node, ...dimensions, [internalsSymbol]: { @@ -169,6 +177,10 @@ export function updateNodeDimensions( }, }, }; + + nodeLookup.set(node.id, newNode); + + return newNode; } } diff --git a/packages/system/src/xydrag/XYDrag.ts b/packages/system/src/xydrag/XYDrag.ts index 9271cb9d..69d50b41 100644 --- a/packages/system/src/xydrag/XYDrag.ts +++ b/packages/system/src/xydrag/XYDrag.ts @@ -33,6 +33,7 @@ export type OnDrag = (event: MouseEvent, dragItems: NodeDragItem[], node: NodeBa type StoreItems = { nodes: NodeBase[]; + nodeLookup: Map; edges: EdgeBase[]; nodeExtent: CoordinateExtent; snapGrid: SnapGrid; @@ -103,6 +104,7 @@ export function XYDrag({ function updateNodes({ x, y }: XYPosition) { const { nodes, + nodeLookup, nodeExtent, snapGrid, snapToGrid, @@ -163,11 +165,11 @@ export function XYDrag({ updateNodePositions(dragItems, true, true); const onNodeOrSelectionDrag = nodeId ? onNodeDrag : wrapSelectionDragFunc(onSelectionDrag); - if (dragEvent) { + if (dragEvent && (onDrag || onNodeOrSelectionDrag)) { const [currentNode, currentNodes] = getEventHandlerParams({ nodeId, dragItems, - nodes, + nodeLookup, }); onDrag?.(dragEvent as MouseEvent, dragItems, currentNode, currentNodes); onNodeOrSelectionDrag?.(dragEvent as MouseEvent, currentNode, currentNodes); @@ -197,6 +199,7 @@ export function XYDrag({ function startDrag(event: UseDragEvent) { const { nodes, + nodeLookup, multiSelectionActive, nodesDraggable, transform, @@ -211,7 +214,7 @@ export function XYDrag({ dragStarted = true; if ((!selectNodesOnDrag || !isSelectable) && !multiSelectionActive && nodeId) { - if (!nodes.find((n) => n.id === nodeId)?.selected) { + if (!nodeLookup.get(nodeId)?.selected) { // we need to reset selected nodes when selectNodesOnDrag=false unselectNodesAndEdges(); } @@ -227,11 +230,11 @@ export function XYDrag({ const onNodeOrSelectionDragStart = nodeId ? onNodeDragStart : wrapSelectionDragFunc(onSelectionDragStart); - if (dragItems) { + if (dragItems && (onDragStart || onNodeOrSelectionDragStart)) { const [currentNode, currentNodes] = getEventHandlerParams({ nodeId, dragItems, - nodes, + nodeLookup, }); onDragStart?.(event.sourceEvent as MouseEvent, dragItems, currentNode, currentNodes); onNodeOrSelectionDragStart?.(event.sourceEvent as MouseEvent, currentNode, currentNodes); @@ -288,18 +291,20 @@ export function XYDrag({ cancelAnimationFrame(autoPanId); if (dragItems) { - const { nodes, updateNodePositions, onNodeDragStop, onSelectionDragStop } = getStoreItems(); + const { nodeLookup, updateNodePositions, onNodeDragStop, onSelectionDragStop } = getStoreItems(); const onNodeOrSelectionDragStop = nodeId ? onNodeDragStop : wrapSelectionDragFunc(onSelectionDragStop); updateNodePositions(dragItems, false, false); - const [currentNode, currentNodes] = getEventHandlerParams({ - nodeId, - dragItems, - nodes, - }); - onDragStop?.(event.sourceEvent as MouseEvent, dragItems, currentNode, currentNodes); - onNodeOrSelectionDragStop?.(event.sourceEvent as MouseEvent, currentNode, currentNodes); + if (onDragStop || onNodeOrSelectionDragStop) { + const [currentNode, currentNodes] = getEventHandlerParams({ + nodeId, + dragItems, + nodeLookup, + }); + onDragStop?.(event.sourceEvent as MouseEvent, dragItems, currentNode, currentNodes); + onNodeOrSelectionDragStop?.(event.sourceEvent as MouseEvent, currentNode, currentNodes); + } } }) .filter((event: MouseEvent) => { diff --git a/packages/system/src/xydrag/utils.ts b/packages/system/src/xydrag/utils.ts index fbf0e8f9..274c2213 100644 --- a/packages/system/src/xydrag/utils.ts +++ b/packages/system/src/xydrag/utils.ts @@ -75,14 +75,14 @@ export function getDragItems( export function getEventHandlerParams({ nodeId, dragItems, - nodes, + nodeLookup, }: { nodeId?: string; dragItems: NodeDragItem[]; - nodes: NodeType[]; + nodeLookup: Map; }): [NodeType, NodeType[]] { const extentedDragItems: NodeType[] = dragItems.map((n) => { - const node = nodes.find((node) => node.id === n.id)!; + const node = nodeLookup.get(n.id)!; return { ...node,