diff --git a/packages/core/src/utils/graph.ts b/packages/core/src/utils/graph.ts index 5b509b13..64bf9215 100644 --- a/packages/core/src/utils/graph.ts +++ b/packages/core/src/utils/graph.ts @@ -364,25 +364,28 @@ export function getNodesInside( }) } -export function getConnectedEdges(nodes: N[], edges: E[]) { - const nodeIds = nodes.map((node) => node.id) +export function getConnectedEdges(nodes: N[], edges: E[]) { + const nodeIds = nodes.map((node) => (isString(node) ? node : node.id)) return edges.filter((edge) => nodeIds.includes(edge.source) || nodeIds.includes(edge.target)) } -export function getConnectedNodes(nodes: N[], edges: E[]) { - const nodeIds = nodes.map((node) => node.id) +export function getConnectedNodes(nodes: N[], edges: E[]) { + const nodeIds = nodes.map((node) => (isString(node) ? node : node.id)) + const connectedNodeIds = edges.reduce((acc, edge) => { if (nodeIds.includes(edge.source)) { acc.add(edge.target) } + if (nodeIds.includes(edge.target)) { acc.add(edge.source) } + return acc }, new Set()) - return nodes.filter((node) => connectedNodeIds.has(node.id)) + return nodes.filter((node) => connectedNodeIds.has(isString(node) ? node : node.id)) } export function getTransformForBounds(