diff --git a/example/src/Basic/index.tsx b/example/src/Basic/index.tsx index 1476b437..7f91818d 100644 --- a/example/src/Basic/index.tsx +++ b/example/src/Basic/index.tsx @@ -66,6 +66,7 @@ const BasicFlow = () => { maxZoom={4} fitView defaultEdgeOptions={defaultEdgeOptions} + selectNodesOnDrag={false} > diff --git a/src/components/Nodes/DefaultNode.tsx b/src/components/Nodes/DefaultNode.tsx index 63921133..353bd7ca 100644 --- a/src/components/Nodes/DefaultNode.tsx +++ b/src/components/Nodes/DefaultNode.tsx @@ -4,14 +4,11 @@ import Handle from '../../components/Handle'; import { NodeProps, Position } from '../../types'; const DefaultNode = ({ - id, data, isConnectable, targetPosition = Position.Top, sourcePosition = Position.Bottom, }: NodeProps) => { - console.log('render', id); - return ( <> diff --git a/src/components/Nodes/utils.ts b/src/components/Nodes/utils.ts index b0e6ea7e..2b9b5137 100644 --- a/src/components/Nodes/utils.ts +++ b/src/components/Nodes/utils.ts @@ -1,5 +1,5 @@ import { MouseEvent } from 'react'; -import { GetState } from 'zustand'; +import { GetState, SetState } from 'zustand'; import { HandleElement, Node, Position, ReactFlowState } from '../../types'; import { getDimensions } from '../../utils'; @@ -55,3 +55,29 @@ export function getMouseHandler( handler(event, { ...node }); }; } + +// this handler is called by +// 1. the click handler when node is not draggable or selectNodesOnDrag = false +// or +// 2. the on drag start handler when node is draggable and selectNodesOnDrag = true +export function handleNodeClick({ + id, + store, +}: { + id: string; + store: { + getState: GetState; + setState: SetState; + }; +}) { + const { addSelectedNodes, unselectNodesAndEdges, multiSelectionActive, nodeInternals } = store.getState(); + const node = nodeInternals.get(id)!; + + store.setState({ nodesSelectionActive: false }); + + if (!node.selected) { + addSelectedNodes([id]); + } else if (node.selected && multiSelectionActive) { + unselectNodesAndEdges({ nodes: [node] }); + } +} diff --git a/src/components/Nodes/wrapNode.tsx b/src/components/Nodes/wrapNode.tsx index 513b3d44..b1293f56 100644 --- a/src/components/Nodes/wrapNode.tsx +++ b/src/components/Nodes/wrapNode.tsx @@ -1,17 +1,13 @@ import React, { useEffect, useRef, memo, ComponentType, MouseEvent } from 'react'; import cc from 'classcat'; -import shallow from 'zustand/shallow'; import { useStore, useStoreApi } from '../../store'; import { Provider } from '../../contexts/NodeIdContext'; import { NodeProps, WrapNodeProps, ReactFlowState } from '../../types'; import useDrag from '../../hooks/useDrag'; -import { getMouseHandler } from './utils'; +import { getMouseHandler, handleNodeClick } from './utils'; -const selector = (s: ReactFlowState) => ({ - addSelectedNodes: s.addSelectedNodes, - updateNodeDimensions: s.updateNodeDimensions, -}); +const selector = (s: ReactFlowState) => s.updateNodeDimensions; export default (NodeComponent: ComponentType) => { const NodeWrapper = ({ @@ -27,9 +23,9 @@ export default (NodeComponent: ComponentType) => { onMouseLeave, onContextMenu, onNodeDoubleClick, - onNodeDragStart, - onNodeDrag, - onNodeDragStop, + onDragStart, + onDrag, + onDragStop, style, className, isDraggable, @@ -47,8 +43,8 @@ export default (NodeComponent: ComponentType) => { noDragClassName, }: WrapNodeProps) => { const store = useStoreApi(); - const { addSelectedNodes, updateNodeDimensions } = useStore(selector, shallow); - const nodeElement = useRef(null); + const updateNodeDimensions = useStore(selector); + const nodeRef = useRef(null); const prevSourcePosition = useRef(sourcePosition); const prevTargetPosition = useRef(targetPosition); const prevType = useRef(type); @@ -60,12 +56,12 @@ export default (NodeComponent: ComponentType) => { const onContextMenuHandler = getMouseHandler(id, store.getState, onContextMenu); const onNodeDoubleClickHandler = getMouseHandler(id, store.getState, onNodeDoubleClick); const onSelectNodeHandler = (event: MouseEvent) => { - if (isSelectable) { - store.setState({ nodesSelectionActive: false }); - - if (!selected) { - addSelectedNodes([id]); - } + if (isSelectable && (!selectNodesOnDrag || !isDraggable)) { + // this handler gets called within the drag start event when selectNodesOnDrag=true + handleNodeClick({ + id, + store, + }); } if (onClick) { @@ -75,8 +71,8 @@ export default (NodeComponent: ComponentType) => { }; useEffect(() => { - if (nodeElement.current && !hidden) { - const currNode = nodeElement.current; + if (nodeRef.current && !hidden) { + const currNode = nodeRef.current; resizeObserver?.observe(currNode); return () => resizeObserver?.unobserve(currNode); @@ -89,7 +85,7 @@ export default (NodeComponent: ComponentType) => { const sourcePosChanged = prevSourcePosition.current !== sourcePosition; const targetPosChanged = prevTargetPosition.current !== targetPosition; - if (nodeElement.current && (typeChanged || sourcePosChanged || targetPosChanged)) { + if (nodeRef.current && (typeChanged || sourcePosChanged || targetPosChanged)) { if (typeChanged) { prevType.current = type; } @@ -99,15 +95,15 @@ export default (NodeComponent: ComponentType) => { if (targetPosChanged) { prevTargetPosition.current = targetPosition; } - updateNodeDimensions([{ id, nodeElement: nodeElement.current, forceUpdate: true }]); + updateNodeDimensions([{ id, nodeElement: nodeRef.current, forceUpdate: true }]); } }, [id, type, sourcePosition, targetPosition]); const dragging = useDrag({ - onStart: onNodeDragStart, - onDrag: onNodeDrag, - onStop: onNodeDragStop, - nodeRef: nodeElement, + onStart: onDragStart, + onDrag: onDrag, + onStop: onDragStop, + nodeRef, disabled: !isDraggable, noDragClassName, handleSelector: dragHandle, @@ -120,22 +116,20 @@ export default (NodeComponent: ComponentType) => { return null; } - const nodeClasses = cc([ - 'react-flow__node', - `react-flow__node-${type}`, - noPanClassName, - className, - { - selected, - selectable: isSelectable, - parent: isParent, - }, - ]); - return (
; export type UseDragData = { dx: number; dy: number }; @@ -76,30 +77,23 @@ function useDrag({ } else { const dragHandler = drag() .on('start', (event: UseDragEvent) => { - const { nodeInternals, addSelectedNodes, unselectNodesAndEdges, multiSelectionActive } = store.getState(); + const { nodeInternals, multiSelectionActive, unselectNodesAndEdges } = store.getState(); parentPos.current = getParentNodePosition(nodeInternals, nodeId); - // this part is the regular drag handler for a single node - // it selects the dragged node and deselects all other nodes if multiSelectionActive = false - if (nodeId && isSelectable) { - const node = nodeInternals.get(nodeId)!; - - if (selectNodesOnDrag) { - store.setState({ nodesSelectionActive: false }); - - if (!node.selected) { - addSelectedNodes([nodeId]); - } - } else if (!selectNodesOnDrag && !node.selected) { - if (multiSelectionActive) { - addSelectedNodes([nodeId]); - } else { - unselectNodesAndEdges(); - store.setState({ nodesSelectionActive: false }); - } + if (!selectNodesOnDrag && !multiSelectionActive && nodeId) { + if (!nodeInternals.get(nodeId)?.selected) { + // we need to reset selected nodes when selectNodesOnDrag=false + unselectNodesAndEdges(); } } + if (nodeId && isSelectable && selectNodesOnDrag) { + handleNodeClick({ + id: nodeId, + store, + }); + } + const mousePos = getMousePosition(event); dragItems.current = getDragItems(nodeInternals, mousePos, nodeId); diff --git a/src/store/index.ts b/src/store/index.ts index 419ecbab..727a6515 100644 --- a/src/store/index.ts +++ b/src/store/index.ts @@ -14,6 +14,7 @@ import { NodeSelectionChange, NodePositionChange, NodeDragItem, + UnselectNodesAndEdgesParams, } from '../types'; import { getHandleBounds } from '../components/Nodes/utils'; import { createSelectionChange, getSelectionChanges } from '../utils/changes'; @@ -152,19 +153,22 @@ const createStore = () => set, }); }, - unselectNodesAndEdges: () => { - const { nodeInternals, edges } = get(); - const nodes = Array.from(nodeInternals.values()); + unselectNodesAndEdges: ({ nodes, edges }: UnselectNodesAndEdgesParams = {}) => { + const { nodeInternals, edges: storeEdges } = get(); + const nodesToUnselect = nodes ? nodes : Array.from(nodeInternals.values()); + const edgesToUnselect = edges ? edges : storeEdges; - const nodesToUnselect = nodes.map((n) => { + const changedNodes = nodesToUnselect.map((n) => { n.selected = false; return createSelectionChange(n.id, false); }) as NodeSelectionChange[]; - const edgesToUnselect = edges.map((edge) => createSelectionChange(edge.id, false)) as EdgeSelectionChange[]; + const changedEdges = edgesToUnselect.map((edge) => + createSelectionChange(edge.id, false) + ) as EdgeSelectionChange[]; updateNodesAndEdgesSelections({ - changedNodes: nodesToUnselect, - changedEdges: edgesToUnselect, + changedNodes, + changedEdges, get, set, }); diff --git a/src/types/general.ts b/src/types/general.ts index 978de52a..1a36991f 100644 --- a/src/types/general.ts +++ b/src/types/general.ts @@ -105,6 +105,11 @@ export type FitBoundsOptions = ViewportHelperFunctionOptions & { padding?: number; }; +export type UnselectNodesAndEdgesParams = { + nodes?: Node[]; + edges?: Edge[]; +}; + export interface ViewportHelperFunctions { zoomIn: ZoomInOut; zoomOut: ZoomInOut; @@ -187,7 +192,7 @@ export type ReactFlowActions = { updateNodeDimensions: (updates: NodeDimensionUpdate[]) => void; updateNodePositions: (nodeDragItems: NodeDragItem[]) => void; resetSelectedElements: () => void; - unselectNodesAndEdges: () => void; + unselectNodesAndEdges: (params?: UnselectNodesAndEdgesParams) => void; addSelectedNodes: (nodeIds: string[]) => void; addSelectedEdges: (edgeIds: string[]) => void; setMinZoom: (minZoom: number) => void; diff --git a/src/types/nodes.ts b/src/types/nodes.ts index 6c416ef5..7a5db412 100644 --- a/src/types/nodes.ts +++ b/src/types/nodes.ts @@ -71,9 +71,9 @@ export interface WrapNodeProps { onMouseMove?: NodeMouseHandler; onMouseLeave?: NodeMouseHandler; onContextMenu?: NodeMouseHandler; - onNodeDragStart?: NodeDragHandler; - onNodeDrag?: NodeDragHandler; - onNodeDragStop?: NodeDragHandler; + onDragStart?: NodeDragHandler; + onDrag?: NodeDragHandler; + onDragStop?: NodeDragHandler; style?: CSSProperties; className?: string; sourcePosition: Position; diff --git a/src/utils/changes.ts b/src/utils/changes.ts index 9279a3f0..2e6a1cf0 100644 --- a/src/utils/changes.ts +++ b/src/utils/changes.ts @@ -45,8 +45,7 @@ function handleParentExpand(res: any[], updateItem: any) { } function applyChanges(changes: any[], elements: any[]): any[] { - // unfortunately we need this hack to handle the setNodes and setEdges function of the - // useReactFlow hook. + // 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); }