Merge pull request #4221 from xyflow/refactor/use-react-flow

refactor(react): cleanup useReactFlow
This commit is contained in:
Moritz Klack
2024-04-29 17:31:47 +02:00
committed by GitHub
4 changed files with 204 additions and 284 deletions

View File

@@ -1,4 +1,4 @@
import { useCallback, useMemo } from 'react';
import { useMemo } from 'react';
import {
EdgeRemoveChange,
evaluateAbsolutePosition,
@@ -11,10 +11,12 @@ import {
} from '@xyflow/system';
import useViewportHelper from './useViewportHelper';
import { useStoreApi } from './useStore';
import { useStore, useStoreApi } from './useStore';
import { useBatchContext } from '../components/BatchProvider';
import { elementToRemoveChange, 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.
@@ -29,172 +31,40 @@ export function useReactFlow<NodeType extends Node = Node, EdgeType extends Edge
const viewportHelper = useViewportHelper();
const store = useStoreApi();
const batchContext = useBatchContext();
const viewportInitialized = useStore(selector);
const getNodes = useCallback<Instance.GetNodes<NodeType>>(
() => store.getState().nodes.map((n) => ({ ...n })) as NodeType[],
[]
);
const generalHelper = useMemo<GeneralHelpers<NodeType, EdgeType>>(() => {
const getInternalNode: GeneralHelpers<NodeType, EdgeType>['getInternalNode'] = (id) =>
store.getState().nodeLookup.get(id) as InternalNode<NodeType>;
const getInternalNode = useCallback<Instance.GetInternalNode<NodeType>>(
(id) => store.getState().nodeLookup.get(id) as InternalNode<NodeType>,
[]
);
const getNode = useCallback<Instance.GetNode<NodeType>>(
(id) => getInternalNode(id)?.internals.userNode as NodeType,
[getInternalNode]
);
const getEdges = useCallback<Instance.GetEdges<EdgeType>>(() => {
const { edges = [] } = store.getState();
return edges.map((e) => ({ ...e })) as EdgeType[];
}, []);
const getEdge = useCallback<Instance.GetEdge<EdgeType>>((id) => store.getState().edgeLookup.get(id) as EdgeType, []);
const setNodes = useCallback<Instance.SetNodes<NodeType>>((payload) => {
batchContext.nodeQueue.push(payload as NodeType[]);
}, []);
const setEdges = useCallback<Instance.SetEdges<EdgeType>>((payload) => {
batchContext.edgeQueue.push(payload as EdgeType[]);
}, []);
const addNodes = useCallback<Instance.AddNodes<NodeType>>((payload) => {
const newNodes = Array.isArray(payload) ? payload : [payload];
batchContext.nodeQueue.push((nodes) => [...nodes, ...newNodes]);
}, []);
const addEdges = useCallback<Instance.AddEdges<EdgeType>>((payload) => {
const newEdges = Array.isArray(payload) ? payload : [payload];
batchContext.edgeQueue.push((edges) => [...edges, ...newEdges]);
}, []);
const toObject = useCallback<Instance.ToObject<NodeType, EdgeType>>(() => {
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<Instance.DeleteElements>(
async ({ nodes: nodesToRemove = [], edges: edgesToRemove = [] }) => {
const {
nodes,
edges,
onNodesDelete,
onEdgesDelete,
triggerNodeChanges,
triggerEdgeChanges,
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) {
const edgeChanges: EdgeRemoveChange[] = matchingEdges.map(elementToRemoveChange);
onEdgesDelete?.(matchingEdges);
triggerEdgeChanges(edgeChanges);
}
if (hasMatchingNodes) {
const nodeChanges: NodeRemoveChange[] = matchingNodes.map(elementToRemoveChange);
onNodesDelete?.(matchingNodes);
triggerNodeChanges(nodeChanges);
}
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<NodeType>(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<NodeType, EdgeType>['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<Instance.GetIntersectingNodes<NodeType>>(
(nodeOrRect, partially = true, nodes) => {
const isRect = isRectObject(nodeOrRect);
const nodeRect = isRect ? nodeOrRect : getNodeRect(nodeOrRect);
const hasNodesOption = nodes !== undefined;
const nodeToUse = isNode<NodeType>(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<Instance.IsNodeIntersecting<NodeType>>(
(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<Instance.UpdateNode<NodeType>>(
(id, nodeUpdate, options = { replace: false }) => {
const updateNode: GeneralHelpers<NodeType, EdgeType>['updateNode'] = (
id,
nodeUpdate,
options = { replace: false }
) => {
setNodes((prevNodes) =>
prevNodes.map((node) => {
if (node.id === id) {
@@ -205,59 +75,139 @@ export function useReactFlow<NodeType extends Node = Node, EdgeType extends Edge
return node;
})
);
},
[setNodes]
);
};
const updateNodeData = useCallback<Instance.UpdateNodeData<NodeType>>(
(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,
onNodesDelete,
onEdgesDelete,
triggerNodeChanges,
triggerEdgeChanges,
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) {
const edgeChanges: EdgeRemoveChange[] = matchingEdges.map(elementToRemoveChange);
onEdgesDelete?.(matchingEdges);
triggerEdgeChanges(edgeChanges);
}
if (hasMatchingNodes) {
const nodeChanges: NodeRemoveChange[] = matchingNodes.map(elementToRemoveChange);
onNodesDelete?.(matchingNodes);
triggerNodeChanges(nodeChanges);
}
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]);
}

View File

@@ -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<ViewportHelperFunctions>(() => {
return useMemo<ViewportHelperFunctions>(() => {
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;

View File

@@ -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<NodeType extends Node = Node, EdgeType extends Edge = Edge> = OnBeforeDeleteBase<

View File

@@ -13,115 +13,70 @@ export type DeleteElementsOptions = {
edges?: (Edge | { id: Edge['id'] })[];
};
export namespace Instance {
export type GetNodes<NodeType extends Node = Node> = () => NodeType[];
export type SetNodes<NodeType extends Node = Node> = (
payload: NodeType[] | ((nodes: NodeType[]) => NodeType[])
) => void;
export type AddNodes<NodeType extends Node = Node> = (payload: NodeType[] | NodeType) => void;
export type GetNode<NodeType extends Node = Node> = (id: string) => NodeType | undefined;
export type GetInternalNode<NodeType extends Node = Node> = (id: string) => InternalNode<NodeType> | undefined;
export type GetEdges<EdgeType extends Edge = Edge> = () => EdgeType[];
export type SetEdges<EdgeType extends Edge = Edge> = (
payload: EdgeType[] | ((edges: EdgeType[]) => EdgeType[])
) => void;
export type GetEdge<EdgeType extends Edge = Edge> = (id: string) => EdgeType | undefined;
export type AddEdges<EdgeType extends Edge = Edge> = (payload: EdgeType[] | EdgeType) => void;
export type ToObject<NodeType extends Node = Node, EdgeType extends Edge = Edge> = () => ReactFlowJsonObject<
NodeType,
EdgeType
>;
export type DeleteElements = (params: DeleteElementsOptions) => Promise<{
deletedNodes: Node[];
deletedEdges: Edge[];
}>;
export type GetIntersectingNodes<NodeType extends Node = Node> = (
node: NodeType | { id: Node['id'] } | Rect,
partially?: boolean,
nodes?: NodeType[]
) => NodeType[];
export type IsNodeIntersecting<NodeType extends Node = Node> = (
node: NodeType | { id: Node['id'] } | Rect,
area: Rect,
partially?: boolean
) => boolean;
export type UpdateNode<NodeType extends Node = Node> = (
id: string,
nodeUpdate: Partial<NodeType> | ((node: NodeType) => Partial<NodeType>),
options?: { replace: boolean }
) => void;
export type UpdateNodeData<NodeType extends Node = Node> = (
id: string,
dataUpdate: object | ((node: NodeType) => object),
options?: { replace: boolean }
) => void;
}
export type ReactFlowInstance<NodeType extends Node = Node, EdgeType extends Edge = Edge> = {
export type GeneralHelpers<NodeType extends Node = Node, EdgeType extends Edge = Edge> = {
/**
* Returns nodes.
*
* @returns nodes array
*/
getNodes: Instance.GetNodes<NodeType>;
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<NodeType>;
setNodes: (payload: NodeType[] | ((nodes: NodeType[]) => NodeType[])) => void;
/**
* Adds nodes.
*
* @param payload - the nodes to add
*/
addNodes: Instance.AddNodes<NodeType>;
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<NodeType>;
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<NodeType>;
getInternalNode: (id: string) => InternalNode<NodeType> | undefined;
/**
* Returns edges.
*
* @returns edges array
*/
getEdges: Instance.GetEdges<EdgeType>;
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<EdgeType>;
setEdges: (payload: EdgeType[] | ((edges: EdgeType[]) => EdgeType[])) => void;
/**
* Adds edges.
*
* @param payload - the edges to add
*/
addEdges: Instance.AddEdges<EdgeType>;
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<EdgeType>;
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<NodeType, EdgeType>;
toObject: () => ReactFlowJsonObject<NodeType, EdgeType>;
/**
* Deletes nodes and edges.
*
@@ -130,7 +85,10 @@ export type ReactFlowInstance<NodeType extends Node = Node, EdgeType extends Edg
*
* @returns a promise that resolves with the deleted nodes and edges
*/
deleteElements: Instance.DeleteElements;
deleteElements: (params: DeleteElementsOptions) => Promise<{
deletedNodes: Node[];
deletedEdges: Edge[];
}>;
/**
* Returns all nodes that intersect with the given node or rect.
*
@@ -140,7 +98,11 @@ export type ReactFlowInstance<NodeType extends Node = Node, EdgeType extends Edg
*
* @returns an array of intersecting nodes
*/
getIntersectingNodes: Instance.GetIntersectingNodes<NodeType>;
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<NodeType extends Node = Node, EdgeType extends Edg
*
* @returns true if the node or rect intersects with the given area
*/
isNodeIntersecting: Instance.IsNodeIntersecting<NodeType>;
isNodeIntersecting: (node: NodeType | { id: Node['id'] } | Rect, area: Rect, partially?: boolean) => boolean;
/**
* Updates a node.
*
@@ -161,7 +123,11 @@ export type ReactFlowInstance<NodeType extends Node = Node, EdgeType extends Edg
* @example
* updateNode('node-1', (node) => ({ position: { x: node.position.x + 10, y: node.position.y } }));
*/
updateNode: Instance.UpdateNode<NodeType>;
updateNode: (
id: string,
nodeUpdate: Partial<NodeType> | ((node: NodeType) => Partial<NodeType>),
options?: { replace: boolean }
) => void;
/**
* Updates the data attribute of a node.
*
@@ -172,6 +138,17 @@ export type ReactFlowInstance<NodeType extends Node = Node, EdgeType extends Edg
* @example
* updateNodeData('node-1', { label: 'A new label' });
*/
updateNodeData: Instance.UpdateNodeData<NodeType>;
viewportInitialized: boolean;
} & Omit<ViewportHelperFunctions, 'initialized'>;
updateNodeData: (
id: string,
dataUpdate: object | ((node: NodeType) => object),
options?: { replace: boolean }
) => void;
};
export type ReactFlowInstance<NodeType extends Node = Node, EdgeType extends Edge = Edge> = GeneralHelpers<
NodeType,
EdgeType
> &
Omit<ViewportHelperFunctions, 'initialized'> & {
viewportInitialized: boolean;
};