Merge pull request #3480 from wbkd/get-connected-edges
Improvements for getConnectedEdges, getIncomers, getOutgoers
This commit is contained in:
@@ -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,
|
||||||
]);
|
]);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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> = {
|
||||||
|
|||||||
@@ -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
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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>>(
|
||||||
|
|||||||
Reference in New Issue
Block a user