diff --git a/packages/react/src/hooks/useReactFlow.ts b/packages/react/src/hooks/useReactFlow.ts index f66ac38b..da65a4e3 100644 --- a/packages/react/src/hooks/useReactFlow.ts +++ b/packages/react/src/hooks/useReactFlow.ts @@ -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(): ReactFlo [] ); + const getConnectedEdges = useCallback((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((node) => { + const { nodes, edges } = store.getState(); + + if (typeof node === 'string') { + return getIncomersBase({ id: node }, nodes, edges); + } + + return getIncomersBase(node, nodes, edges); + }, []); + + const getOutgoers = useCallback((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(): ReactFlo deleteElements, getIntersectingNodes, isNodeIntersecting, + getConnectedEdges, + getIncomers, + getOutgoers, }; }, [ viewportHelper, @@ -255,5 +301,8 @@ export default function useReactFlow(): ReactFlo deleteElements, getIntersectingNodes, isNodeIntersecting, + getConnectedEdges, + getIncomers, + getOutgoers, ]); } diff --git a/packages/react/src/types/instance.ts b/packages/react/src/types/instance.ts index 531e518e..3d123c1a 100644 --- a/packages/react/src/types/instance.ts +++ b/packages/react/src/types/instance.ts @@ -42,6 +42,9 @@ export namespace Instance { area: Rect, partially?: boolean ) => boolean; + export type getConnectedEdges = (id: string | (Partial & { id: Node['id'] })[]) => Edge[]; + export type getIncomers = (node: string | (Partial & { id: Node['id'] })) => Node[]; + export type getOutgoers = (node: string | (Partial & { id: Node['id'] })) => Node[]; } export type ReactFlowInstance = { diff --git a/packages/svelte/src/lib/hooks/useSvelteFlow.ts b/packages/svelte/src/lib/hooks/useSvelteFlow.ts index 7dce00b5..756b856c 100644 --- a/packages/svelte/src/lib/hooks/useSvelteFlow.ts +++ b/packages/svelte/src/lib/hooks/useSvelteFlow.ts @@ -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; + getConnectedEdges: (id: string | (Partial & { id: Node['id'] })[]) => Edge[]; + getIncomers: (node: string | (Partial & { id: Node['id'] })) => Node[]; + getOutgoers: (node: string | (Partial & { 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 }; } diff --git a/packages/system/src/utils/graph.ts b/packages/system/src/utils/graph.ts index fb7d74c2..17467503 100644 --- a/packages/system/src/utils/graph.ts +++ b/packages/system/src/utils/graph.ts @@ -35,29 +35,40 @@ export const isNodeBase = 'id' in element && !('source' in element) && !('target' in element); export const getOutgoersBase = ( - node: NodeType, + node: Partial & { 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 = ( - node: NodeType, + node: Partial & { 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 = { - 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, Options extends FitViewOptionsBase>(