Merge pull request #3480 from wbkd/get-connected-edges

Improvements for getConnectedEdges, getIncomers, getOutgoers
This commit is contained in:
Moritz Klack
2023-10-05 20:23:04 +02:00
committed by GitHub
4 changed files with 115 additions and 28 deletions
+50 -1
View File
@@ -1,5 +1,13 @@
import { useCallback, useMemo } from 'react'; import { useCallback, useMemo } from 'react';
import { getElementsToRemove, getOverlappingArea, isRectObject, nodeToRect, type Rect } from '@xyflow/system'; import {
getElementsToRemove,
getIncomersBase,
getOutgoersBase,
getOverlappingArea,
isRectObject,
nodeToRect,
type Rect,
} from '@xyflow/system';
import useViewportHelper from './useViewportHelper'; import useViewportHelper from './useViewportHelper';
import { useStoreApi } from '../hooks/useStore'; import { useStoreApi } from '../hooks/useStore';
@@ -225,6 +233,41 @@ export default function useReactFlow<NodeData = any, EdgeData = any>(): ReactFlo
[] []
); );
const getConnectedEdges = useCallback<Instance.getConnectedEdges>((node) => {
const { edges } = store.getState();
const nodeIds = new Set();
if (typeof node === 'string') {
nodeIds.add(node);
} else if (node.length >= 1) {
node.forEach((n) => {
nodeIds.add(n.id);
});
}
return edges.filter((edge) => nodeIds.has(edge.source) || nodeIds.has(edge.target));
}, []);
const getIncomers = useCallback<Instance.getIncomers>((node) => {
const { nodes, edges } = store.getState();
if (typeof node === 'string') {
return getIncomersBase({ id: node }, nodes, edges);
}
return getIncomersBase(node, nodes, edges);
}, []);
const getOutgoers = useCallback<Instance.getOutgoers>((node) => {
const { nodes, edges } = store.getState();
if (typeof node == 'string') {
return getOutgoersBase({ id: node }, nodes, edges);
}
return getOutgoersBase(node, nodes, edges);
}, []);
return useMemo(() => { return useMemo(() => {
return { return {
...viewportHelper, ...viewportHelper,
@@ -240,6 +283,9 @@ export default function useReactFlow<NodeData = any, EdgeData = any>(): ReactFlo
deleteElements, deleteElements,
getIntersectingNodes, getIntersectingNodes,
isNodeIntersecting, isNodeIntersecting,
getConnectedEdges,
getIncomers,
getOutgoers,
}; };
}, [ }, [
viewportHelper, viewportHelper,
@@ -255,5 +301,8 @@ export default function useReactFlow<NodeData = any, EdgeData = any>(): ReactFlo
deleteElements, deleteElements,
getIntersectingNodes, getIntersectingNodes,
isNodeIntersecting, isNodeIntersecting,
getConnectedEdges,
getIncomers,
getOutgoers,
]); ]);
} }
+3
View File
@@ -42,6 +42,9 @@ export namespace Instance {
area: Rect, area: Rect,
partially?: boolean partially?: boolean
) => boolean; ) => boolean;
export type getConnectedEdges = (id: string | (Partial<Node> & { id: Node['id'] })[]) => Edge[];
export type getIncomers = (node: string | (Partial<Node> & { id: Node['id'] })) => Node[];
export type getOutgoers = (node: string | (Partial<Node> & { id: Node['id'] })) => Node[];
} }
export type ReactFlowInstance<NodeData = any, EdgeData = any> = { export type ReactFlowInstance<NodeData = any, EdgeData = any> = {
+38 -17
View File
@@ -1,5 +1,7 @@
import { get, type Writable } from 'svelte/store'; import { get, type Writable } from 'svelte/store';
import { import {
getIncomersBase,
getOutgoersBase,
getOverlappingArea, getOverlappingArea,
isRectObject, isRectObject,
nodeToRect, nodeToRect,
@@ -46,6 +48,9 @@ export function useSvelteFlow(): {
screenToFlowCoordinate: (position: XYPosition) => XYPosition; screenToFlowCoordinate: (position: XYPosition) => XYPosition;
flowToScreenCoordinate: (position: XYPosition) => XYPosition; flowToScreenCoordinate: (position: XYPosition) => XYPosition;
viewport: Writable<Viewport>; viewport: Writable<Viewport>;
getConnectedEdges: (id: string | (Partial<Node> & { id: Node['id'] })[]) => Edge[];
getIncomers: (node: string | (Partial<Node> & { id: Node['id'] })) => Node[];
getOutgoers: (node: string | (Partial<Node> & { id: Node['id'] })) => Node[];
} { } {
const { const {
zoomIn, zoomIn,
@@ -99,16 +104,12 @@ export function useSvelteFlow(): {
}, },
getViewport: () => get(viewport), getViewport: () => get(viewport),
setCenter: (x, y, options) => { setCenter: (x, y, options) => {
const _width = get(width); const nextZoom = typeof options?.zoom !== 'undefined' ? options.zoom : get(maxZoom);
const _height = get(height);
const _maxZoom = get(maxZoom);
const nextZoom = typeof options?.zoom !== 'undefined' ? options.zoom : _maxZoom;
get(panZoom)?.setViewport( get(panZoom)?.setViewport(
{ {
x: _width / 2 - x * nextZoom, x: get(width) / 2 - x * nextZoom,
y: _height / 2 - y * nextZoom, y: get(height) / 2 - y * nextZoom,
zoom: nextZoom zoom: nextZoom
}, },
{ duration: options?.duration } { duration: options?.duration }
@@ -116,17 +117,12 @@ export function useSvelteFlow(): {
}, },
fitView, fitView,
fitBounds: (bounds: Rect, options?: FitBoundsOptions) => { fitBounds: (bounds: Rect, options?: FitBoundsOptions) => {
const _width = get(width);
const _height = get(height);
const _maxZoom = get(maxZoom);
const _minZoom = get(minZoom);
const [x, y, zoom] = getTransformForBounds( const [x, y, zoom] = getTransformForBounds(
bounds, bounds,
_width, get(width),
_height, get(height),
_minZoom, get(minZoom),
_maxZoom, get(maxZoom),
options?.padding ?? 0.1 options?.padding ?? 0.1
); );
@@ -206,6 +202,7 @@ export function useSvelteFlow(): {
}, },
screenToFlowCoordinate: (position: XYPosition) => { screenToFlowCoordinate: (position: XYPosition) => {
const _domNode = get(domNode); const _domNode = get(domNode);
if (_domNode) { if (_domNode) {
const _snapGrid = get(snapGrid); const _snapGrid = get(snapGrid);
const { x, y, zoom } = get(viewport); const { x, y, zoom } = get(viewport);
@@ -228,6 +225,7 @@ export function useSvelteFlow(): {
}, },
flowToScreenCoordinate: (position: XYPosition) => { flowToScreenCoordinate: (position: XYPosition) => {
const _domNode = get(domNode); const _domNode = get(domNode);
if (_domNode) { if (_domNode) {
const { x, y, zoom } = get(viewport); const { x, y, zoom } = get(viewport);
const { x: domX, y: domY } = _domNode.getBoundingClientRect(); const { x: domX, y: domY } = _domNode.getBoundingClientRect();
@@ -242,6 +240,29 @@ export function useSvelteFlow(): {
return { x: 0, y: 0 }; return { x: 0, y: 0 };
}, },
viewport: viewport getConnectedEdges: (node) => {
const nodeIds = new Set();
if (typeof node === 'string') {
nodeIds.add(node);
} else if (node.length >= 1) {
node.forEach((n) => {
nodeIds.add(n.id);
});
}
return get(edges).filter((edge) => nodeIds.has(edge.source) || nodeIds.has(edge.target));
},
getIncomers: (node) => {
const _node = typeof node === 'string' ? { id: node } : node;
return getIncomersBase(_node, get(nodes), get(edges));
},
getOutgoers: (node) => {
const _node = typeof node === 'string' ? { id: node } : node;
return getOutgoersBase(_node, get(nodes), get(edges));
},
viewport
}; };
} }
+24 -10
View File
@@ -35,29 +35,40 @@ export const isNodeBase = <NodeType extends NodeBase = NodeBase, EdgeType extend
): element is NodeType => 'id' in element && !('source' in element) && !('target' in element); ): element is NodeType => 'id' in element && !('source' in element) && !('target' in element);
export const getOutgoersBase = <NodeType extends NodeBase = NodeBase, EdgeType extends EdgeBase = EdgeBase>( export const getOutgoersBase = <NodeType extends NodeBase = NodeBase, EdgeType extends EdgeBase = EdgeBase>(
node: NodeType, node: Partial<NodeType> & { id: string },
nodes: NodeType[], nodes: NodeType[],
edges: EdgeType[] edges: EdgeType[]
): NodeType[] => { ): NodeType[] => {
if (!isNodeBase(node)) { if (!node.id) {
return []; return [];
} }
const outgoerIds = edges.filter((e) => e.source === node.id).map((e) => e.target); const outgoerIds = new Set();
return nodes.filter((n) => outgoerIds.includes(n.id)); edges.forEach((edge) => {
if (edge.source === node.id) {
outgoerIds.add(edge.target);
}
});
return nodes.filter((n) => outgoerIds.has(n.id));
}; };
export const getIncomersBase = <NodeType extends NodeBase = NodeBase, EdgeType extends EdgeBase = EdgeBase>( export const getIncomersBase = <NodeType extends NodeBase = NodeBase, EdgeType extends EdgeBase = EdgeBase>(
node: NodeType, node: Partial<NodeType> & { id: string },
nodes: NodeType[], nodes: NodeType[],
edges: EdgeType[] edges: EdgeType[]
): NodeType[] => { ): NodeType[] => {
if (!isNodeBase(node)) { if (!node.id) {
return []; return [];
} }
const incomersIds = new Set();
edges.forEach((edge) => {
if (edge.target === node.id) {
incomersIds.add(edge.source);
}
});
const incomersIds = edges.filter((e) => e.target === node.id).map((e) => e.source); return nodes.filter((n) => incomersIds.has(n.id));
return nodes.filter((n) => incomersIds.includes(n.id));
}; };
export const getNodePositionWithOrigin = ( export const getNodePositionWithOrigin = (
@@ -161,9 +172,12 @@ export const getConnectedEdgesBase = <NodeType extends NodeBase = NodeBase, Edge
nodes: NodeType[], nodes: NodeType[],
edges: EdgeType[] edges: EdgeType[]
): EdgeType[] => { ): EdgeType[] => {
const nodeIds = nodes.map((node) => node.id); const nodeIds = new Set();
nodes.forEach((node) => {
nodeIds.add(node.id);
});
return edges.filter((edge) => nodeIds.includes(edge.source) || nodeIds.includes(edge.target)); return edges.filter((edge) => nodeIds.has(edge.source) || nodeIds.has(edge.target));
}; };
export function fitView<Params extends FitViewParamsBase<NodeBase>, Options extends FitViewOptionsBase<NodeBase>>( export function fitView<Params extends FitViewParamsBase<NodeBase>, Options extends FitViewOptionsBase<NodeBase>>(