diff --git a/packages/react/src/hooks/useReactFlow.ts b/packages/react/src/hooks/useReactFlow.ts index 4ac0e50c..3b68763e 100644 --- a/packages/react/src/hooks/useReactFlow.ts +++ b/packages/react/src/hooks/useReactFlow.ts @@ -1,4 +1,4 @@ -import { useCallback, useMemo } from 'react'; +import { useMemo } from 'react'; import { evaluateAbsolutePosition, getElementsToRemove, @@ -9,10 +9,12 @@ import { } from '@xyflow/system'; import useViewportHelper from './useViewportHelper'; -import { useStoreApi } from './useStore'; +import { useStore, useStoreApi } from './useStore'; import { useBatchContext } from '../components/BatchProvider'; import { isNode } from '../utils'; -import type { ReactFlowInstance, Instance, Node, Edge, InternalNode } from '../types'; +import type { ReactFlowInstance, Node, Edge, InternalNode, ReactFlowState, GeneralHelpers } from '../types'; + +const selector = (s: ReactFlowState) => !!s.panZoom; /** * Hook for accessing the ReactFlow instance. @@ -27,185 +29,40 @@ export function useReactFlow>( - () => store.getState().nodes.map((n) => ({ ...n })) as NodeType[], - [] - ); + const generalHelper = useMemo>(() => { + const getInternalNode: GeneralHelpers['getInternalNode'] = (id) => + store.getState().nodeLookup.get(id) as InternalNode; - const getInternalNode = useCallback>( - (id) => store.getState().nodeLookup.get(id) as InternalNode, - [] - ); - - const getNode = useCallback>( - (id) => getInternalNode(id)?.internals.userNode as NodeType, - [getInternalNode] - ); - - const getEdges = useCallback>(() => { - const { edges = [] } = store.getState(); - return edges.map((e) => ({ ...e })) as EdgeType[]; - }, []); - - const getEdge = useCallback>((id) => store.getState().edgeLookup.get(id) as EdgeType, []); - - const setNodes = useCallback>((payload) => { - batchContext.nodeQueue.push(payload as NodeType[]); - }, []); - - const setEdges = useCallback>((payload) => { - batchContext.edgeQueue.push(payload as EdgeType[]); - }, []); - - const addNodes = useCallback>((payload) => { - const newNodes = Array.isArray(payload) ? payload : [payload]; - batchContext.nodeQueue.push((nodes) => [...nodes, ...newNodes]); - }, []); - - const addEdges = useCallback>((payload) => { - const newEdges = Array.isArray(payload) ? payload : [payload]; - batchContext.edgeQueue.push((edges) => [...edges, ...newEdges]); - }, []); - - const toObject = useCallback>(() => { - const { nodes = [], edges = [], transform } = store.getState(); - const [x, y, zoom] = transform; - return { - nodes: nodes.map((n) => ({ ...n })) as NodeType[], - edges: edges.map((e) => ({ ...e })) as EdgeType[], - viewport: { - x, - y, - zoom, - }, - }; - }, []); - - const deleteElements = useCallback( - async ({ nodes: nodesToRemove = [], edges: edgesToRemove = [] }) => { - const { - nodes, - edges, - hasDefaultNodes, - hasDefaultEdges, - onNodesDelete, - onEdgesDelete, - onNodesChange, - onEdgesChange, - onDelete, - onBeforeDelete, - } = store.getState(); - const { nodes: matchingNodes, edges: matchingEdges } = await getElementsToRemove({ - nodesToRemove, - edgesToRemove, - nodes, - edges, - onBeforeDelete, - }); - - const hasMatchingEdges = matchingEdges.length > 0; - const hasMatchingNodes = matchingNodes.length > 0; - - if (hasMatchingEdges) { - if (hasDefaultEdges) { - const nextEdges = edges.filter((e) => !matchingEdges.some((mE) => mE.id === e.id)); - store.getState().setEdges(nextEdges); - } - - onEdgesDelete?.(matchingEdges); - onEdgesChange?.( - matchingEdges.map((edge) => ({ - id: edge.id, - type: 'remove', - })) - ); - } - - if (hasMatchingNodes) { - if (hasDefaultNodes) { - const nextNodes = nodes.filter((n) => !matchingNodes.some((mN) => mN.id === n.id)); - store.getState().setNodes(nextNodes); - } - - onNodesDelete?.(matchingNodes); - onNodesChange?.(matchingNodes.map((node) => ({ id: node.id, type: 'remove' }))); - } - - if (hasMatchingNodes || hasMatchingEdges) { - onDelete?.({ nodes: matchingNodes, edges: matchingEdges }); - } - - return { deletedNodes: matchingNodes, deletedEdges: matchingEdges }; - }, - [] - ); - - const getNodeRect = useCallback((node: NodeType | { id: string }): Rect | null => { - const { nodeLookup, nodeOrigin } = store.getState(); - - const nodeToUse = isNode(node) ? node : nodeLookup.get(node.id)!; - const position = nodeToUse.parentId - ? evaluateAbsolutePosition(nodeToUse.position, nodeToUse.parentId, nodeLookup, nodeOrigin) - : nodeToUse.position; - - const nodeWithPosition = { - id: nodeToUse.id, - position, - width: nodeToUse.measured?.width ?? nodeToUse.width, - height: nodeToUse.measured?.height ?? nodeToUse.height, - data: nodeToUse.data, + const setNodes: GeneralHelpers['setNodes'] = (payload) => { + batchContext.nodeQueue.push(payload as NodeType[]); }; - return nodeToRect(nodeWithPosition); - }, []); + const getNodeRect = (node: NodeType | { id: string }): Rect | null => { + const { nodeLookup, nodeOrigin } = store.getState(); - const getIntersectingNodes = useCallback>( - (nodeOrRect, partially = true, nodes) => { - const isRect = isRectObject(nodeOrRect); - const nodeRect = isRect ? nodeOrRect : getNodeRect(nodeOrRect); - const hasNodesOption = nodes !== undefined; + const nodeToUse = isNode(node) ? node : nodeLookup.get(node.id)!; + const position = nodeToUse.parentId + ? evaluateAbsolutePosition(nodeToUse.position, nodeToUse.parentId, nodeLookup, nodeOrigin) + : nodeToUse.position; - if (!nodeRect) { - return []; - } + const nodeWithPosition = { + id: nodeToUse.id, + position, + width: nodeToUse.measured?.width ?? nodeToUse.width, + height: nodeToUse.measured?.height ?? nodeToUse.height, + data: nodeToUse.data, + }; - return (nodes || store.getState().nodes).filter((n) => { - const internalNode = store.getState().nodeLookup.get(n.id); + return nodeToRect(nodeWithPosition); + }; - if (internalNode && !isRect && (n.id === nodeOrRect!.id || !internalNode.internals.positionAbsolute)) { - return false; - } - - const currNodeRect = nodeToRect(hasNodesOption ? n : internalNode!); - const overlappingArea = getOverlappingArea(currNodeRect, nodeRect); - const partiallyVisible = partially && overlappingArea > 0; - - return partiallyVisible || overlappingArea >= nodeRect.width * nodeRect.height; - }) as NodeType[]; - }, - [] - ); - - const isNodeIntersecting = useCallback>( - (nodeOrRect, area, partially = true) => { - const isRect = isRectObject(nodeOrRect); - const nodeRect = isRect ? nodeOrRect : getNodeRect(nodeOrRect); - - if (!nodeRect) { - return false; - } - - const overlappingArea = getOverlappingArea(nodeRect, area); - const partiallyVisible = partially && overlappingArea > 0; - - return partiallyVisible || overlappingArea >= nodeRect.width * nodeRect.height; - }, - [] - ); - - const updateNode = useCallback>( - (id, nodeUpdate, options = { replace: false }) => { + const updateNode: GeneralHelpers['updateNode'] = ( + id, + nodeUpdate, + options = { replace: false } + ) => { setNodes((prevNodes) => prevNodes.map((node) => { if (node.id === id) { @@ -216,59 +73,152 @@ export function useReactFlow>( - (id, dataUpdate, options = { replace: false }) => { - updateNode( - id, - (node) => { - const nextData = typeof dataUpdate === 'function' ? dataUpdate(node) : dataUpdate; - return options.replace ? { ...node, data: nextData } : { ...node, data: { ...node.data, ...nextData } }; - }, - options - ); - }, - [updateNode] - ); + return { + getNodes: () => store.getState().nodes.map((n) => ({ ...n })) as NodeType[], + getNode: (id) => getInternalNode(id)?.internals.userNode as NodeType, + getInternalNode, + getEdges: () => { + const { edges = [] } = store.getState(); + return edges.map((e) => ({ ...e })) as EdgeType[]; + }, + getEdge: (id) => store.getState().edgeLookup.get(id) as EdgeType, + setNodes, + setEdges: (payload) => { + batchContext.edgeQueue.push(payload as EdgeType[]); + }, + addNodes: (payload) => { + const newNodes = Array.isArray(payload) ? payload : [payload]; + batchContext.nodeQueue.push((nodes) => [...nodes, ...newNodes]); + }, + addEdges: (payload) => { + const newEdges = Array.isArray(payload) ? payload : [payload]; + batchContext.edgeQueue.push((edges) => [...edges, ...newEdges]); + }, + toObject: () => { + const { nodes = [], edges = [], transform } = store.getState(); + const [x, y, zoom] = transform; + return { + nodes: nodes.map((n) => ({ ...n })) as NodeType[], + edges: edges.map((e) => ({ ...e })) as EdgeType[], + viewport: { + x, + y, + zoom, + }, + }; + }, + deleteElements: async ({ nodes: nodesToRemove = [], edges: edgesToRemove = [] }) => { + const { + nodes, + edges, + hasDefaultNodes, + hasDefaultEdges, + onNodesDelete, + onEdgesDelete, + onNodesChange, + onEdgesChange, + onDelete, + onBeforeDelete, + } = store.getState(); + const { nodes: matchingNodes, edges: matchingEdges } = await getElementsToRemove({ + nodesToRemove, + edgesToRemove, + nodes, + edges, + onBeforeDelete, + }); + + const hasMatchingEdges = matchingEdges.length > 0; + const hasMatchingNodes = matchingNodes.length > 0; + + if (hasMatchingEdges) { + if (hasDefaultEdges) { + const nextEdges = edges.filter((e) => !matchingEdges.some((mE) => mE.id === e.id)); + store.getState().setEdges(nextEdges); + } + + onEdgesDelete?.(matchingEdges); + onEdgesChange?.( + matchingEdges.map((edge) => ({ + id: edge.id, + type: 'remove', + })) + ); + } + + if (hasMatchingNodes) { + if (hasDefaultNodes) { + const nextNodes = nodes.filter((n) => !matchingNodes.some((mN) => mN.id === n.id)); + store.getState().setNodes(nextNodes); + } + + onNodesDelete?.(matchingNodes); + onNodesChange?.(matchingNodes.map((node) => ({ id: node.id, type: 'remove' }))); + } + + if (hasMatchingNodes || hasMatchingEdges) { + onDelete?.({ nodes: matchingNodes, edges: matchingEdges }); + } + + return { deletedNodes: matchingNodes, deletedEdges: matchingEdges }; + }, + getIntersectingNodes: (nodeOrRect, partially = true, nodes) => { + const isRect = isRectObject(nodeOrRect); + const nodeRect = isRect ? nodeOrRect : getNodeRect(nodeOrRect); + const hasNodesOption = nodes !== undefined; + + if (!nodeRect) { + return []; + } + + return (nodes || store.getState().nodes).filter((n) => { + const internalNode = store.getState().nodeLookup.get(n.id); + + if (internalNode && !isRect && (n.id === nodeOrRect!.id || !internalNode.internals.positionAbsolute)) { + return false; + } + + const currNodeRect = nodeToRect(hasNodesOption ? n : internalNode!); + const overlappingArea = getOverlappingArea(currNodeRect, nodeRect); + const partiallyVisible = partially && overlappingArea > 0; + + return partiallyVisible || overlappingArea >= nodeRect.width * nodeRect.height; + }) as NodeType[]; + }, + isNodeIntersecting: (nodeOrRect, area, partially = true) => { + const isRect = isRectObject(nodeOrRect); + const nodeRect = isRect ? nodeOrRect : getNodeRect(nodeOrRect); + + if (!nodeRect) { + return false; + } + + const overlappingArea = getOverlappingArea(nodeRect, area); + const partiallyVisible = partially && overlappingArea > 0; + + return partiallyVisible || overlappingArea >= nodeRect.width * nodeRect.height; + }, + updateNode, + updateNodeData: (id, dataUpdate, options = { replace: false }) => { + updateNode( + id, + (node) => { + const nextData = typeof dataUpdate === 'function' ? dataUpdate(node) : dataUpdate; + return options.replace ? { ...node, data: nextData } : { ...node, data: { ...node.data, ...nextData } }; + }, + options + ); + }, + }; + }, []); return useMemo(() => { return { + ...generalHelper, ...viewportHelper, - getNodes, - getNode, - getInternalNode, - getEdges, - getEdge, - setNodes, - setEdges, - addNodes, - addEdges, - toObject, - deleteElements, - getIntersectingNodes, - isNodeIntersecting, - updateNode, - updateNodeData, + viewportInitialized, }; - }, [ - viewportHelper, - getNodes, - getNode, - getInternalNode, - getEdges, - getEdge, - setNodes, - setEdges, - addNodes, - addEdges, - toObject, - deleteElements, - getIntersectingNodes, - isNodeIntersecting, - updateNode, - updateNodeData, - ]); + }, [viewportInitialized]); } diff --git a/packages/react/src/hooks/useViewportHelper.ts b/packages/react/src/hooks/useViewportHelper.ts index eb92b310..ef4f6f38 100644 --- a/packages/react/src/hooks/useViewportHelper.ts +++ b/packages/react/src/hooks/useViewportHelper.ts @@ -7,10 +7,8 @@ import { rendererPointToPoint, } from '@xyflow/system'; -import { useStoreApi, useStore } from '../hooks/useStore'; -import type { ViewportHelperFunctions, ReactFlowState } from '../types'; - -const selector = (s: ReactFlowState) => !!s.panZoom; +import { useStoreApi } from '../hooks/useStore'; +import type { ViewportHelperFunctions } from '../types'; /** * Hook for getting viewport helper functions. @@ -20,9 +18,8 @@ const selector = (s: ReactFlowState) => !!s.panZoom; */ const useViewportHelper = (): ViewportHelperFunctions => { const store = useStoreApi(); - const panZoomInitialized = useStore(selector); - const viewportHelperFunctions = useMemo(() => { + return useMemo(() => { return { zoomIn: (options) => store.getState().panZoom?.scaleBy(1.2, { duration: options?.duration }), zoomOut: (options) => store.getState().panZoom?.scaleBy(1 / 1.2, { duration: options?.duration }), @@ -117,11 +114,8 @@ const useViewportHelper = (): ViewportHelperFunctions => { y: rendererPosition.y + domY, }; }, - viewportInitialized: panZoomInitialized, }; - }, [panZoomInitialized]); - - return viewportHelperFunctions; + }, []); }; export default useViewportHelper; diff --git a/packages/react/src/types/general.ts b/packages/react/src/types/general.ts index 8b55920f..6f1105a7 100644 --- a/packages/react/src/types/general.ts +++ b/packages/react/src/types/general.ts @@ -156,7 +156,6 @@ export type ViewportHelperFunctions = { * const clientPosition = flowToScreenPosition({ x: node.position.x, y: node.position.y }) */ flowToScreenPosition: (flowPosition: XYPosition) => XYPosition; - viewportInitialized: boolean; }; export type OnBeforeDelete = OnBeforeDeleteBase< diff --git a/packages/react/src/types/instance.ts b/packages/react/src/types/instance.ts index 67a6d928..8ce74d1c 100644 --- a/packages/react/src/types/instance.ts +++ b/packages/react/src/types/instance.ts @@ -13,115 +13,70 @@ export type DeleteElementsOptions = { edges?: (Edge | { id: Edge['id'] })[]; }; -export namespace Instance { - export type GetNodes = () => NodeType[]; - export type SetNodes = ( - payload: NodeType[] | ((nodes: NodeType[]) => NodeType[]) - ) => void; - export type AddNodes = (payload: NodeType[] | NodeType) => void; - export type GetNode = (id: string) => NodeType | undefined; - export type GetInternalNode = (id: string) => InternalNode | undefined; - export type GetEdges = () => EdgeType[]; - export type SetEdges = ( - payload: EdgeType[] | ((edges: EdgeType[]) => EdgeType[]) - ) => void; - export type GetEdge = (id: string) => EdgeType | undefined; - export type AddEdges = (payload: EdgeType[] | EdgeType) => void; - export type ToObject = () => ReactFlowJsonObject< - NodeType, - EdgeType - >; - export type DeleteElements = (params: DeleteElementsOptions) => Promise<{ - deletedNodes: Node[]; - deletedEdges: Edge[]; - }>; - export type GetIntersectingNodes = ( - node: NodeType | { id: Node['id'] } | Rect, - partially?: boolean, - nodes?: NodeType[] - ) => NodeType[]; - export type IsNodeIntersecting = ( - node: NodeType | { id: Node['id'] } | Rect, - area: Rect, - partially?: boolean - ) => boolean; - - export type UpdateNode = ( - id: string, - nodeUpdate: Partial | ((node: NodeType) => Partial), - options?: { replace: boolean } - ) => void; - export type UpdateNodeData = ( - id: string, - dataUpdate: object | ((node: NodeType) => object), - options?: { replace: boolean } - ) => void; -} - -export type ReactFlowInstance = { +export type GeneralHelpers = { /** * Returns nodes. * * @returns nodes array */ - getNodes: Instance.GetNodes; + getNodes: () => NodeType[]; /** * Sets nodes. * * @param payload - the nodes to set or a function that receives the current nodes and returns the new nodes */ - setNodes: Instance.SetNodes; + setNodes: (payload: NodeType[] | ((nodes: NodeType[]) => NodeType[])) => void; /** * Adds nodes. * * @param payload - the nodes to add */ - addNodes: Instance.AddNodes; + addNodes: (payload: NodeType[] | NodeType) => void; /** * Returns a node by id. * * @param id - the node id * @returns the node or undefined if no node was found */ - getNode: Instance.GetNode; + getNode: (id: string) => NodeType | undefined; /** * Returns an internal node by id. * * @param id - the node id * @returns the internal node or undefined if no node was found */ - getInternalNode: Instance.GetInternalNode; + getInternalNode: (id: string) => InternalNode | undefined; /** * Returns edges. * * @returns edges array */ - getEdges: Instance.GetEdges; + getEdges: () => EdgeType[]; /** * Sets edges. * * @param payload - the edges to set or a function that receives the current edges and returns the new edges */ - setEdges: Instance.SetEdges; + setEdges: (payload: EdgeType[] | ((edges: EdgeType[]) => EdgeType[])) => void; /** * Adds edges. * * @param payload - the edges to add */ - addEdges: Instance.AddEdges; + addEdges: (payload: EdgeType[] | EdgeType) => void; /** * Returns an edge by id. * * @param id - the edge id * @returns the edge or undefined if no edge was found */ - getEdge: Instance.GetEdge; + getEdge: (id: string) => EdgeType | undefined; /** * Returns the nodes, edges and the viewport as a JSON object. * * @returns the nodes, edges and the viewport as a JSON object */ - toObject: Instance.ToObject; + toObject: () => ReactFlowJsonObject; /** * Deletes nodes and edges. * @@ -130,7 +85,10 @@ export type ReactFlowInstance Promise<{ + deletedNodes: Node[]; + deletedEdges: Edge[]; + }>; /** * Returns all nodes that intersect with the given node or rect. * @@ -140,7 +98,11 @@ export type ReactFlowInstance; + getIntersectingNodes: ( + node: NodeType | { id: Node['id'] } | Rect, + partially?: boolean, + nodes?: NodeType[] + ) => NodeType[]; /** * Checks if the given node or rect intersects with the passed rect. * @@ -150,7 +112,7 @@ export type ReactFlowInstance; + isNodeIntersecting: (node: NodeType | { id: Node['id'] } | Rect, area: Rect, partially?: boolean) => boolean; /** * Updates a node. * @@ -161,7 +123,11 @@ export type ReactFlowInstance ({ position: { x: node.position.x + 10, y: node.position.y } })); */ - updateNode: Instance.UpdateNode; + updateNode: ( + id: string, + nodeUpdate: Partial | ((node: NodeType) => Partial), + options?: { replace: boolean } + ) => void; /** * Updates the data attribute of a node. * @@ -172,6 +138,17 @@ export type ReactFlowInstance; - viewportInitialized: boolean; -} & Omit; + updateNodeData: ( + id: string, + dataUpdate: object | ((node: NodeType) => object), + options?: { replace: boolean } + ) => void; +}; + +export type ReactFlowInstance = GeneralHelpers< + NodeType, + EdgeType +> & + Omit & { + viewportInitialized: boolean; + };