diff --git a/example/src/Basic/index.tsx b/example/src/Basic/index.tsx index 7f91818d..74d0f69b 100644 --- a/example/src/Basic/index.tsx +++ b/example/src/Basic/index.tsx @@ -9,6 +9,7 @@ import ReactFlow, { useReactFlow, } from 'react-flow-renderer'; +const onNodeDrag = (_: MouseEvent, node: Node) => console.log('drag', node); const onNodeDragStop = (_: MouseEvent, node: Node) => console.log('drag stop', node); const onNodeClick = (_: MouseEvent, node: Node) => console.log('click', node); @@ -61,6 +62,7 @@ const BasicFlow = () => { defaultEdges={initialEdges} onNodeClick={onNodeClick} onNodeDragStop={onNodeDragStop} + onNodeDrag={onNodeDrag} className="react-flow-basic-example" minZoom={0.2} maxZoom={4} diff --git a/example/src/SaveRestore/Controls.tsx b/example/src/SaveRestore/Controls.tsx index 4c4537c7..c26d55d2 100644 --- a/example/src/SaveRestore/Controls.tsx +++ b/example/src/SaveRestore/Controls.tsx @@ -1,5 +1,5 @@ import React, { memo, useCallback, Dispatch, FC } from 'react'; -import { useReactFlow, ReactFlowInstance, Edge, Node, ReactFlowJsonObject } from 'react-flow-renderer'; +import { useReactFlow, Edge, Node, ReactFlowJsonObject } from 'react-flow-renderer'; import localforage from 'localforage'; localforage.config({ @@ -12,20 +12,17 @@ const flowKey = 'example-flow'; const getNodeId = () => `randomnode_${+new Date()}`; type ControlsProps = { - rfInstance?: ReactFlowInstance; setNodes: Dispatch[]>>; setEdges: Dispatch[]>>; }; -const Controls: FC = ({ rfInstance, setNodes, setEdges }) => { - const { setViewport } = useReactFlow(); +const Controls: FC = ({ setNodes, setEdges }) => { + const { setViewport, toObject } = useReactFlow(); const onSave = useCallback(() => { - if (rfInstance) { - const flow = rfInstance.toObject(); - localforage.setItem(flowKey, flow); - } - }, [rfInstance]); + const flow = toObject(); + localforage.setItem(flowKey, flow); + }, [toObject]); const onRestore = useCallback(() => { const restoreFlow = async () => { @@ -33,6 +30,7 @@ const Controls: FC = ({ rfInstance, setNodes, setEdges }) => { if (flow) { const { x, y, zoom } = flow.viewport; + setNodes(flow.nodes || []); setEdges(flow.edges || []); setViewport({ x, y, zoom: zoom || 0 }); diff --git a/example/src/SaveRestore/index.tsx b/example/src/SaveRestore/index.tsx index c0aacccb..a28f2f4a 100644 --- a/example/src/SaveRestore/index.tsx +++ b/example/src/SaveRestore/index.tsx @@ -1,11 +1,10 @@ -import { useState } from 'react'; +import { useCallback } from 'react'; import ReactFlow, { ReactFlowProvider, Node, addEdge, Connection, Edge, - ReactFlowInstance, useNodesState, useEdgesState, } from 'react-flow-renderer'; @@ -22,10 +21,9 @@ const initialNodes: Node[] = [ const initialEdges: Edge[] = [{ id: 'e1-2', source: '1', target: '2' }]; const SaveRestore = () => { - const [rfInstance, setRfInstance] = useState(); const [nodes, setNodes, onNodesChange] = useNodesState(initialNodes); const [edges, setEdges, onEdgesChange] = useEdgesState(initialEdges); - const onConnect = (params: Connection | Edge) => setEdges((eds) => addEdge(params, eds)); + const onConnect = useCallback((params: Connection | Edge) => setEdges((eds) => addEdge(params, eds)), [setEdges]); return ( @@ -35,9 +33,8 @@ const SaveRestore = () => { onNodesChange={onNodesChange} onEdgesChange={onEdgesChange} onConnect={onConnect} - onInit={setRfInstance} > - + ); diff --git a/example/src/Subflow/DebugNode.tsx b/example/src/Subflow/DebugNode.tsx index f7cc4b36..e73631a0 100644 --- a/example/src/Subflow/DebugNode.tsx +++ b/example/src/Subflow/DebugNode.tsx @@ -11,7 +11,7 @@ const idStyle: CSSProperties = { left: 2, }; -const ColorSelectorNode: FC = ({ zIndex, xPos, yPos, id }) => { +const DebugNode: FC = ({ zIndex, xPos, yPos, id }) => { return ( <> @@ -24,4 +24,4 @@ const ColorSelectorNode: FC = ({ zIndex, xPos, yPos, id }) => { ); }; -export default memo(ColorSelectorNode); +export default memo(DebugNode); diff --git a/example/src/Subflow/index.tsx b/example/src/Subflow/index.tsx index 1d7c79d2..6b6c4799 100644 --- a/example/src/Subflow/index.tsx +++ b/example/src/Subflow/index.tsx @@ -15,6 +15,7 @@ import ReactFlow, { } from 'react-flow-renderer'; import DebugNode from './DebugNode'; +const onNodeDrag = (_: MouseEvent, node: Node, nodes: Node[]) => console.log('drag', node, nodes); const onNodeDragStop = (_: MouseEvent, node: Node, nodes: Node[]) => console.log('drag stop', node, nodes); const onNodeClick = (_: MouseEvent, node: Node) => console.log('click', node); const onEdgeClick = (_: MouseEvent, edge: Edge) => console.log('click', edge); @@ -82,6 +83,7 @@ const initialNodes: Node[] = [ position: { x: 25, y: 50 }, className: 'light', parentNode: '5', + extent: 'parent', }, { id: '5b', @@ -167,6 +169,7 @@ const Subflow = () => { onNodeClick={onNodeClick} onEdgeClick={onEdgeClick} onConnect={onConnect} + onNodeDrag={onNodeDrag} onNodeDragStop={onNodeDragStop} className="react-flow-basic-example" defaultZoom={1.5} diff --git a/src/components/Handle/handler.ts b/src/components/Handle/handler.ts index 54a843fc..9c908ece 100644 --- a/src/components/Handle/handler.ts +++ b/src/components/Handle/handler.ts @@ -46,6 +46,24 @@ export function checkElementBelowIsValid( if (elementBelow && (elementBelowIsTarget || elementBelowIsSource)) { result.isHoveringHandle = true; + const elementBelowNodeId = elementBelow.getAttribute('data-nodeid'); + const elementBelowHandleId = elementBelow.getAttribute('data-handleid'); + const connection: Connection = isTarget + ? { + source: elementBelowNodeId, + sourceHandle: elementBelowHandleId, + target: nodeId, + targetHandle: handleId, + } + : { + source: nodeId, + sourceHandle: handleId, + target: elementBelowNodeId, + targetHandle: elementBelowHandleId, + }; + + result.connection = connection; + // in strict mode we don't allow target to target or source to source connections const isValid = connectionMode === ConnectionMode.Strict @@ -53,23 +71,6 @@ export function checkElementBelowIsValid( : true; if (isValid) { - const elementBelowNodeId = elementBelow.getAttribute('data-nodeid'); - const elementBelowHandleId = elementBelow.getAttribute('data-handleid'); - const connection: Connection = isTarget - ? { - source: elementBelowNodeId, - sourceHandle: elementBelowHandleId, - target: nodeId, - targetHandle: handleId, - } - : { - source: nodeId, - sourceHandle: handleId, - target: elementBelowNodeId, - targetHandle: elementBelowHandleId, - }; - - result.connection = connection; result.isValid = isValidConnection(connection); } } @@ -151,9 +152,7 @@ export function handleMouseDown( return resetRecentHandle(recentHoveredHandle); } - const isOwnHandle = connection.source === connection.target; - - if (!isOwnHandle && elementBelow) { + if (connection.source !== connection.target && elementBelow) { recentHoveredHandle = elementBelow; elementBelow.classList.add('react-flow__handle-connecting'); elementBelow.classList.toggle('react-flow__handle-valid', isValid); diff --git a/src/components/Nodes/utils.ts b/src/components/Nodes/utils.ts index 2b9b5137..a3d1a28f 100644 --- a/src/components/Nodes/utils.ts +++ b/src/components/Nodes/utils.ts @@ -4,21 +4,17 @@ import { GetState, SetState } from 'zustand'; import { HandleElement, Node, Position, ReactFlowState } from '../../types'; import { getDimensions } from '../../utils'; -export const getHandleBounds = (nodeElement: HTMLDivElement, scale: number) => { - const bounds = nodeElement.getBoundingClientRect(); +function getTranslateValues(domNode: HTMLDivElement): [number, number] { + if (typeof window === 'undefined' || !window.DOMMatrixReadOnly) { + return [0, 0]; + } - return { - source: getHandleBoundsByHandleType('.source', nodeElement, bounds, scale), - target: getHandleBoundsByHandleType('.target', nodeElement, bounds, scale), - }; -}; + const style = window.getComputedStyle(domNode); + const { m41, m42 } = new window.DOMMatrixReadOnly(style.transform); + return [m41, m42]; +} -export const getHandleBoundsByHandleType = ( - selector: string, - nodeElement: HTMLDivElement, - parentBounds: DOMRect, - k: number -): HandleElement[] | null => { +export const getHandleBounds = (selector: string, nodeElement: HTMLDivElement): HandleElement[] | null => { const handles = nodeElement.querySelectorAll(selector); if (!handles || !handles.length) { @@ -28,17 +24,16 @@ export const getHandleBoundsByHandleType = ( const handlesArray = Array.from(handles) as HTMLDivElement[]; return handlesArray.map((handle): HandleElement => { - const bounds = handle.getBoundingClientRect(); - const dimensions = getDimensions(handle); - const handleId = handle.getAttribute('data-handleid'); - const handlePosition = handle.getAttribute('data-handlepos') as unknown as Position; + // we don't use getBoundingClientRect here, because it includes the transform of the parent (scaled viewport) + // that we would then need to calculate out again in order to get the correct position. + const [translateX, translateY] = getTranslateValues(handle); return { - id: handleId, - position: handlePosition, - x: (bounds.left - parentBounds.left) / k, - y: (bounds.top - parentBounds.top) / k, - ...dimensions, + id: handle.getAttribute('data-handleid'), + position: handle.getAttribute('data-handlepos') as unknown as Position, + x: handle.offsetLeft + nodeElement.clientLeft + translateX, + y: handle.offsetTop + nodeElement.clientTop + translateY, + ...getDimensions(handle), }; }); }; diff --git a/src/container/ZoomPane/index.tsx b/src/container/ZoomPane/index.tsx index 57a1c8a5..69ad08b4 100644 --- a/src/container/ZoomPane/index.tsx +++ b/src/container/ZoomPane/index.tsx @@ -54,6 +54,7 @@ const ZoomPane = ({ noPanClassName, }: ZoomPaneProps) => { const store = useStoreApi(); + const isZoomingOrPanning = useRef(false); const zoomPane = useRef(null); const prevTransform = useRef({ x: 0, y: 0, zoom: 0 }); const { d3Zoom, d3Selection, d3ZoomHandler } = useStore(selector, shallow); @@ -146,12 +147,11 @@ const ZoomPane = ({ useEffect(() => { if (d3Zoom) { - if (selectionKeyPressed) { + if (selectionKeyPressed && !isZoomingOrPanning.current) { d3Zoom.on('zoom', null); - } else { + } else if (!selectionKeyPressed) { d3Zoom.on('zoom', (event: D3ZoomEvent) => { store.setState({ transform: [event.transform.x, event.transform.y, event.transform.k] }); - if (onMove) { const flowTransform = eventToFlowTransform(event.transform); onMove(event.sourceEvent as MouseEvent | TouchEvent, flowTransform); @@ -163,33 +163,31 @@ const ZoomPane = ({ useEffect(() => { if (d3Zoom) { - if (onMoveStart) { - d3Zoom.on('start', (event: D3ZoomEvent) => { + d3Zoom.on('start', (event: D3ZoomEvent) => { + isZoomingOrPanning.current = true; + + if (onMoveStart) { const flowTransform = eventToFlowTransform(event.transform); prevTransform.current = flowTransform; onMoveStart(event.sourceEvent as MouseEvent | TouchEvent, flowTransform); - }); - } else { - d3Zoom.on('start', null); - } + } + }); } }, [d3Zoom, onMoveStart]); useEffect(() => { if (d3Zoom) { - if (onMoveEnd) { - d3Zoom.on('end', (event: D3ZoomEvent) => { - if (viewChanged(prevTransform.current, event.transform)) { - const flowTransform = eventToFlowTransform(event.transform); - prevTransform.current = flowTransform; + d3Zoom.on('end', (event: D3ZoomEvent) => { + isZoomingOrPanning.current = false; - onMoveEnd(event.sourceEvent as MouseEvent | TouchEvent, flowTransform); - } - }); - } else { - d3Zoom.on('end', null); - } + if (onMoveEnd && viewChanged(prevTransform.current, event.transform)) { + const flowTransform = eventToFlowTransform(event.transform); + prevTransform.current = flowTransform; + + onMoveEnd(event.sourceEvent as MouseEvent | TouchEvent, flowTransform); + } + }); } }, [d3Zoom, onMoveEnd]); diff --git a/src/hooks/useDrag/index.ts b/src/hooks/useDrag/index.ts index 15d365f8..b14ad14d 100644 --- a/src/hooks/useDrag/index.ts +++ b/src/hooks/useDrag/index.ts @@ -77,6 +77,7 @@ function useDrag({ } const pointerPos = getPointerPosition(event); + lastPos.current = pointerPos; dragItems.current = getDragItems(nodeInternals, pointerPos, nodeId); if (onStart && dragItems.current) { diff --git a/src/hooks/useDrag/utils.ts b/src/hooks/useDrag/utils.ts index cd0e0c41..9679a5c7 100644 --- a/src/hooks/useDrag/utils.ts +++ b/src/hooks/useDrag/utils.ts @@ -39,7 +39,8 @@ export function getDragItems(nodeInternals: NodeInternals, mousePos: XYPosition, .filter((n) => (n.selected || n.id === nodeId) && (!n.parentNode || !isParentSelected(n, nodeInternals))) .map((n) => ({ id: n.id, - position: n.positionAbsolute || { x: 0, y: 0 }, + position: n.position || { x: 0, y: 0 }, + positionAbsolute: n.positionAbsolute || { x: 0, y: 0 }, distance: { x: mousePos.x - (n.positionAbsolute?.x ?? 0), y: mousePos.y - (n.positionAbsolute?.y ?? 0), @@ -94,14 +95,28 @@ export function updatePosition( ]; } - dragItem.position = currentExtent ? clampPosition(nextPosition, currentExtent as CoordinateExtent) : nextPosition; + let parentPosition = { x: 0, y: 0 }; + + if (dragItem.parentNode) { + const parentNode = nodeInternals.get(dragItem.parentNode); + parentPosition = { x: parentNode?.positionAbsolute?.x ?? 0, y: parentNode?.positionAbsolute?.y ?? 0 }; + } + + dragItem.positionAbsolute = currentExtent + ? clampPosition(nextPosition, currentExtent as CoordinateExtent) + : nextPosition; + + dragItem.position = { + x: dragItem.positionAbsolute.x - parentPosition.x, + y: dragItem.positionAbsolute.y - parentPosition.y, + }; return dragItem; } // returns two params: // 1. the dragged node (or the first of the list, if we are dragging a node selection) -// 2. array of selected nodes (handy when multi selection is active) +// 2. array of selected nodes (for multi selections) export function getEventHandlerParams({ nodeId, dragItems, @@ -116,6 +131,8 @@ export function getEventHandlerParams({ return { ...node, + position: n.position, + positionAbsolute: n.positionAbsolute, }; }); diff --git a/src/hooks/useViewportHelper.ts b/src/hooks/useViewportHelper.ts index aeaabe61..10e21193 100644 --- a/src/hooks/useViewportHelper.ts +++ b/src/hooks/useViewportHelper.ts @@ -38,12 +38,12 @@ const useViewportHelper = (): ViewportHelperFunctions => { zoomIn: (options) => d3Zoom.scaleBy(getD3Transition(d3Selection, options?.duration), 1.2), zoomOut: (options) => d3Zoom.scaleBy(getD3Transition(d3Selection, options?.duration), 1 / 1.2), zoomTo: (zoomLevel, options) => d3Zoom.scaleTo(getD3Transition(d3Selection, options?.duration), zoomLevel), - getZoom: () => { - const [, , zoom] = store.getState().transform; - return zoom; - }, + getZoom: () => store.getState().transform[2], setViewport: (transform, options) => { - const nextTransform = zoomIdentity.translate(transform.x, transform.y).scale(transform.zoom); + const [x, y, zoom] = store.getState().transform; + const nextTransform = zoomIdentity + .translate(transform.x ?? x, transform.y ?? y) + .scale(transform.zoom ?? zoom); d3Zoom.transform(getD3Transition(d3Selection, options?.duration), nextTransform); }, getViewport: () => { diff --git a/src/store/index.ts b/src/store/index.ts index ccfdfecd..f67dde4b 100644 --- a/src/store/index.ts +++ b/src/store/index.ts @@ -43,7 +43,7 @@ const createStore = () => set({ nodeInternals, edges: nextEdges, hasDefaultNodes, hasDefaultEdges }); }, updateNodeDimensions: (updates: NodeDimensionUpdate[]) => { - const { onNodesChange, transform, nodeInternals, fitViewOnInit, fitViewOnInitDone, fitViewOnInitOptions } = get(); + const { onNodesChange, nodeInternals, fitViewOnInit, fitViewOnInitDone, fitViewOnInitOptions } = get(); const changes: NodeDimensionChange[] = updates.reduce((res, update) => { const node = nodeInternals.get(update.id); @@ -57,12 +57,14 @@ const createStore = () => ); if (doUpdate) { - const handleBounds = getHandleBounds(update.nodeElement, transform[2]); nodeInternals.set(node.id, { ...node, [internalsSymbol]: { ...node[internalsSymbol], - handleBounds, + handleBounds: { + source: getHandleBounds('.source', update.nodeElement), + target: getHandleBounds('.target', update.nodeElement), + }, }, ...dimensions, }); @@ -99,16 +101,8 @@ const createStore = () => }; if (positionChanged) { - change.positionAbsolute = node.position; + change.positionAbsolute = node.positionAbsolute; change.position = node.position; - - if (node.parentNode) { - const parentNode = nodeInternals.get(node.parentNode); - change.position = { - x: change.position.x - (parentNode?.positionAbsolute?.x ?? 0), - y: change.position.y - (parentNode?.positionAbsolute?.y ?? 0), - }; - } } return change; diff --git a/src/types/nodes.ts b/src/types/nodes.ts index 0045ae34..31712102 100644 --- a/src/types/nodes.ts +++ b/src/types/nodes.ts @@ -110,8 +110,8 @@ export type NodeBounds = XYPosition & { export type NodeDragItem = { id: string; - // relative node position position: XYPosition; + positionAbsolute: XYPosition; // distance from the mouse cursor to the node when start dragging distance: XYPosition; width?: number | null;