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

View File

@@ -1,5 +1,13 @@
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 { 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 {
...viewportHelper,
@@ -240,6 +283,9 @@ export default function useReactFlow<NodeData = any, EdgeData = any>(): ReactFlo
deleteElements,
getIntersectingNodes,
isNodeIntersecting,
getConnectedEdges,
getIncomers,
getOutgoers,
};
}, [
viewportHelper,
@@ -255,5 +301,8 @@ export default function useReactFlow<NodeData = any, EdgeData = any>(): ReactFlo
deleteElements,
getIntersectingNodes,
isNodeIntersecting,
getConnectedEdges,
getIncomers,
getOutgoers,
]);
}

View File

@@ -42,6 +42,9 @@ export namespace Instance {
area: Rect,
partially?: 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> = {

View File

@@ -1,5 +1,7 @@
import { get, type Writable } from 'svelte/store';
import {
getIncomersBase,
getOutgoersBase,
getOverlappingArea,
isRectObject,
nodeToRect,
@@ -46,6 +48,9 @@ export function useSvelteFlow(): {
screenToFlowCoordinate: (position: XYPosition) => XYPosition;
flowToScreenCoordinate: (position: XYPosition) => XYPosition;
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 {
zoomIn,
@@ -99,16 +104,12 @@ export function useSvelteFlow(): {
},
getViewport: () => get(viewport),
setCenter: (x, y, options) => {
const _width = get(width);
const _height = get(height);
const _maxZoom = get(maxZoom);
const nextZoom = typeof options?.zoom !== 'undefined' ? options.zoom : _maxZoom;
const nextZoom = typeof options?.zoom !== 'undefined' ? options.zoom : get(maxZoom);
get(panZoom)?.setViewport(
{
x: _width / 2 - x * nextZoom,
y: _height / 2 - y * nextZoom,
x: get(width) / 2 - x * nextZoom,
y: get(height) / 2 - y * nextZoom,
zoom: nextZoom
},
{ duration: options?.duration }
@@ -116,17 +117,12 @@ export function useSvelteFlow(): {
},
fitView,
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(
bounds,
_width,
_height,
_minZoom,
_maxZoom,
get(width),
get(height),
get(minZoom),
get(maxZoom),
options?.padding ?? 0.1
);
@@ -206,6 +202,7 @@ export function useSvelteFlow(): {
},
screenToFlowCoordinate: (position: XYPosition) => {
const _domNode = get(domNode);
if (_domNode) {
const _snapGrid = get(snapGrid);
const { x, y, zoom } = get(viewport);
@@ -228,6 +225,7 @@ export function useSvelteFlow(): {
},
flowToScreenCoordinate: (position: XYPosition) => {
const _domNode = get(domNode);
if (_domNode) {
const { x, y, zoom } = get(viewport);
const { x: domX, y: domY } = _domNode.getBoundingClientRect();
@@ -242,6 +240,29 @@ export function useSvelteFlow(): {
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
};
}

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);
export const getOutgoersBase = <NodeType extends NodeBase = NodeBase, EdgeType extends EdgeBase = EdgeBase>(
node: NodeType,
node: Partial<NodeType> & { id: string },
nodes: NodeType[],
edges: EdgeType[]
): NodeType[] => {
if (!isNodeBase(node)) {
if (!node.id) {
return [];
}
const outgoerIds = edges.filter((e) => e.source === node.id).map((e) => e.target);
return nodes.filter((n) => outgoerIds.includes(n.id));
const outgoerIds = new Set();
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>(
node: NodeType,
node: Partial<NodeType> & { id: string },
nodes: NodeType[],
edges: EdgeType[]
): NodeType[] => {
if (!isNodeBase(node)) {
if (!node.id) {
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.includes(n.id));
return nodes.filter((n) => incomersIds.has(n.id));
};
export const getNodePositionWithOrigin = (
@@ -161,9 +172,12 @@ export const getConnectedEdgesBase = <NodeType extends NodeBase = NodeBase, Edge
nodes: NodeType[],
edges: 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>>(