Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 4 additions & 8 deletions web/src/components/menus/CommandMenu.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ import {
exportWorkflowBundle,
importWorkflowBundle
} from "../../utils/workflowBundle";
import { useNodes } from "../../contexts/NodeContext";
import { useNodes, useNodeStoreRef } from "../../contexts/NodeContext";
import { create } from "zustand";
import { shallow } from "zustand/shallow";
import { isDevelopment } from "../../lib/env";
Expand Down Expand Up @@ -108,17 +108,12 @@ const styles = () =>

const WorkflowCommands = memo(function WorkflowCommands() {
const executeAndClose = useCommandMenu((state) => state.executeAndClose);
// Optimization: use shallow equality to prevent the CommandMenu from
// re-rendering 60 times a second on unrelated node position updates
const store = useNodeStoreRef();
const {
nodes,
edges,
currentWorkflow,
workflowJSON,
autoLayout
} = useNodes((state) => ({
nodes: state.nodes,
edges: state.edges,
currentWorkflow: state.workflow,
workflowJSON: state.workflowJSON,
autoLayout: state.autoLayout
Expand All @@ -143,8 +138,9 @@ const WorkflowCommands = memo(function WorkflowCommands() {
const bundleInputRef = useRef<HTMLInputElement>(null);

const runWorkflow = useCallback(() => {
const { nodes, edges } = store.getState();
run({}, currentWorkflow, nodes, edges);
}, [run, currentWorkflow, nodes, edges]);
}, [run, currentWorkflow, store]);

const downloadWorkflow = useCallback(() => {
const blob = new Blob([workflowJSON()], { type: "application/json" });
Expand Down
5 changes: 4 additions & 1 deletion web/src/hooks/__tests__/useAlignNodes.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,10 @@ jest.mock("../../contexts/NodeContext", () => ({
getSelectedNodes: jest.fn(() => [])
};
return selector(mockState);
})
}),
useNodeStoreRef: jest.fn(() => ({
getState: () => ({ nodes: mockNodes })
}))
}));

const createMockNode = (
Expand Down
5 changes: 4 additions & 1 deletion web/src/hooks/__tests__/useDuplicate.test.ts
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import { renderHook, act } from "@testing-library/react";
import { useDuplicateNodes } from "../useDuplicate";
import { useReactFlow } from "@xyflow/react";
import { useNodes } from "../../contexts/NodeContext";
import { useNodes, useNodeStoreRef } from "../../contexts/NodeContext";

jest.mock("@xyflow/react");
jest.mock("../../contexts/NodeContext");
Expand Down Expand Up @@ -41,6 +41,9 @@ describe("useDuplicateNodes", () => {
}
return mockUseNodesReturn;
});
(useNodeStoreRef as jest.Mock).mockReturnValue({
getState: () => ({ nodes: mockUseNodesReturn.nodes, edges: mockUseNodesReturn.edges })
});
(useReactFlow as jest.Mock).mockReturnValue({
getNodesBounds: mockGetNodesBounds,
});
Expand Down
122 changes: 34 additions & 88 deletions web/src/hooks/__tests__/useSelectConnected.test.ts
Original file line number Diff line number Diff line change
@@ -1,14 +1,16 @@
import { renderHook, act } from "@testing-library/react";
import { useSelectConnected } from "../useSelectConnected";
import { useNodes } from "../../contexts/NodeContext";
import { useNodes, useNodeStoreRef } from "../../contexts/NodeContext";
import { Node, Edge } from "@xyflow/react";
import { NodeData } from "../../stores/NodeData";

jest.mock("../../contexts/NodeContext", () => ({
useNodes: jest.fn()
useNodes: jest.fn(),
useNodeStoreRef: jest.fn()
}));

const mockUseNodes = useNodes as jest.MockedFunction<typeof useNodes>;
const mockUseNodeStoreRef = useNodeStoreRef as jest.MockedFunction<typeof useNodeStoreRef>;

const createMockNodeData = (): NodeData => ({
properties: {},
Expand Down Expand Up @@ -53,24 +55,32 @@ describe("useSelectConnected", () => {

const createMockSetSelectedNodes = () => jest.fn();

const setupMocks = (
nodes: Node<NodeData>[],
edges: Edge[],
getSelectedNodes: () => Node<NodeData>[],
setSelectedNodes: jest.Mock
) => {
mockUseNodeStoreRef.mockReturnValue({
getState: () => ({ nodes, edges })
} as ReturnType<typeof useNodeStoreRef>);
mockUseNodes.mockReturnValue({
getSelectedNodes,
setSelectedNodes
});
};

beforeEach(() => {
jest.clearAllMocks();
});

describe("direction: both", () => {
it("should select all connected nodes when direction is 'both'", () => {
const setSelectedNodes = createMockSetSelectedNodes();

mockUseNodes.mockReturnValue({
nodes: mockNodes,
edges: mockEdges,
getSelectedNodes: () => [mockNodes[1]],
setSelectedNodes
});
setupMocks(mockNodes, mockEdges, () => [mockNodes[1]], setSelectedNodes);

const { result } = renderHook(() => useSelectConnected({ direction: "both" }));

expect(result.current.connectedNodeCount).toBe(3);
expect(result.current.getConnectedNodeIds()).toEqual([
"input-node",
"process-node-2",
Expand All @@ -80,31 +90,19 @@ describe("useSelectConnected", () => {

it("should select connected nodes when multiple nodes are selected", () => {
const setSelectedNodes = createMockSetSelectedNodes();

mockUseNodes.mockReturnValue({
nodes: mockNodes,
edges: mockEdges,
getSelectedNodes: () => [mockNodes[1], mockNodes[2]],
setSelectedNodes
});
setupMocks(mockNodes, mockEdges, () => [mockNodes[1], mockNodes[2]], setSelectedNodes);

const { result } = renderHook(() => useSelectConnected({ direction: "both" }));

expect(result.current.connectedNodeCount).toBe(2);
const connectedIds = result.current.getConnectedNodeIds();
expect(connectedIds).toContain("input-node");
expect(connectedIds).toContain("output-node");
expect(connectedIds.length).toBe(2);
});

it("should call setSelectedNodes with all connected nodes", () => {
const setSelectedNodes = createMockSetSelectedNodes();

mockUseNodes.mockReturnValue({
nodes: mockNodes,
edges: mockEdges,
getSelectedNodes: () => [mockNodes[1]],
setSelectedNodes
});
setupMocks(mockNodes, mockEdges, () => [mockNodes[1]], setSelectedNodes);

const { result } = renderHook(() => useSelectConnected({ direction: "both" }));

Expand All @@ -126,53 +124,33 @@ describe("useSelectConnected", () => {
describe("direction: upstream", () => {
it("should only select upstream nodes when direction is 'upstream'", () => {
const setSelectedNodes = createMockSetSelectedNodes();

mockUseNodes.mockReturnValue({
nodes: mockNodes,
edges: mockEdges,
getSelectedNodes: () => [mockNodes[2]],
setSelectedNodes
});
setupMocks(mockNodes, mockEdges, () => [mockNodes[2]], setSelectedNodes);

const { result } = renderHook(() => useSelectConnected({ direction: "upstream" }));

expect(result.current.connectedNodeCount).toBe(2);
const connectedIds = result.current.getConnectedNodeIds();
expect(connectedIds).toContain("input-node");
expect(connectedIds).toContain("process-node-1");
expect(connectedIds.length).toBe(2);
});

it("should not include selected nodes in upstream result", () => {
const setSelectedNodes = createMockSetSelectedNodes();

mockUseNodes.mockReturnValue({
nodes: mockNodes,
edges: mockEdges,
getSelectedNodes: () => [mockNodes[0]],
setSelectedNodes
});
setupMocks(mockNodes, mockEdges, () => [mockNodes[0]], setSelectedNodes);

const { result } = renderHook(() => useSelectConnected({ direction: "upstream" }));

expect(result.current.connectedNodeCount).toBe(0);
expect(result.current.getConnectedNodeIds()).toEqual([]);
});
});

describe("direction: downstream", () => {
it("should only select downstream nodes when direction is 'downstream'", () => {
const setSelectedNodes = createMockSetSelectedNodes();

mockUseNodes.mockReturnValue({
nodes: mockNodes,
edges: mockEdges,
getSelectedNodes: () => [mockNodes[1]],
setSelectedNodes
});
setupMocks(mockNodes, mockEdges, () => [mockNodes[1]], setSelectedNodes);

const { result } = renderHook(() => useSelectConnected({ direction: "downstream" }));

expect(result.current.connectedNodeCount).toBe(2);
expect(result.current.getConnectedNodeIds()).toEqual([
"process-node-2",
"output-node"
Expand All @@ -181,47 +159,27 @@ describe("useSelectConnected", () => {

it("should not include selected nodes in downstream result", () => {
const setSelectedNodes = createMockSetSelectedNodes();

mockUseNodes.mockReturnValue({
nodes: mockNodes,
edges: mockEdges,
getSelectedNodes: () => [mockNodes[3]],
setSelectedNodes
});
setupMocks(mockNodes, mockEdges, () => [mockNodes[3]], setSelectedNodes);

const { result } = renderHook(() => useSelectConnected({ direction: "downstream" }));

expect(result.current.connectedNodeCount).toBe(0);
expect(result.current.getConnectedNodeIds()).toEqual([]);
});
});

describe("empty selection", () => {
it("should return empty array when no nodes are selected", () => {
const setSelectedNodes = createMockSetSelectedNodes();

mockUseNodes.mockReturnValue({
nodes: mockNodes,
edges: mockEdges,
getSelectedNodes: () => [],
setSelectedNodes
});
setupMocks(mockNodes, mockEdges, () => [], setSelectedNodes);

const { result } = renderHook(() => useSelectConnected({ direction: "both" }));

expect(result.current.connectedNodeCount).toBe(0);
expect(result.current.getConnectedNodeIds()).toEqual([]);
});

it("should not call setSelectedNodes when no nodes are selected", () => {
const setSelectedNodes = createMockSetSelectedNodes();

mockUseNodes.mockReturnValue({
nodes: mockNodes,
edges: mockEdges,
getSelectedNodes: () => [],
setSelectedNodes
});
setupMocks(mockNodes, mockEdges, () => [], setSelectedNodes);

const { result } = renderHook(() => useSelectConnected({ direction: "both" }));

Expand All @@ -236,17 +194,11 @@ describe("useSelectConnected", () => {
describe("default direction", () => {
it("should default to 'both' direction", () => {
const setSelectedNodes = createMockSetSelectedNodes();

mockUseNodes.mockReturnValue({
nodes: mockNodes,
edges: mockEdges,
getSelectedNodes: () => [mockNodes[1]],
setSelectedNodes
});
setupMocks(mockNodes, mockEdges, () => [mockNodes[1]], setSelectedNodes);

const { result } = renderHook(() => useSelectConnected());

expect(result.current.connectedNodeCount).toBe(3);
expect(result.current.getConnectedNodeIds().length).toBe(3);
});
});

Expand All @@ -269,18 +221,12 @@ describe("useSelectConnected", () => {
];

const setSelectedNodes = createMockSetSelectedNodes();

mockUseNodes.mockReturnValue({
nodes: branchedNodes,
edges: branchedEdges,
getSelectedNodes: () => [branchedNodes[0]],
setSelectedNodes
});
setupMocks(branchedNodes, branchedEdges, () => [branchedNodes[0]], setSelectedNodes);

const { result } = renderHook(() => useSelectConnected({ direction: "downstream" }));

expect(result.current.connectedNodeCount).toBe(4);
const connectedIds = result.current.getConnectedNodeIds();
expect(connectedIds.length).toBe(4);
expect(connectedIds).toContain("branch-a");
expect(connectedIds).toContain("branch-b");
expect(connectedIds).toContain("merged");
Expand Down
13 changes: 6 additions & 7 deletions web/src/hooks/nodes/useGroupIntoSubgraph.ts
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import { useCallback } from "react";
import { useReactFlow, type Edge as RFEdge, type Node as RFNode } from "@xyflow/react";
import { shallow } from "zustand/shallow";
import { useNodes } from "../../contexts/NodeContext";
import { useNodes, useNodeStoreRef } from "../../contexts/NodeContext";
import useMetadataStore from "../../stores/MetadataStore";
import { reactFlowNodeToGraphNode } from "../../stores/reactFlowNodeToGraphNode";
import { reactFlowEdgeToGraphEdge } from "../../stores/reactFlowEdgeToGraphEdge";
Expand Down Expand Up @@ -175,15 +175,14 @@ function buildPlan(
export function useGroupIntoSubgraph() {
const reactFlowInstance = useReactFlow();
const getMetadata = useMetadataStore((s) => s.getMetadata);
const { createNode, addNode, addEdge, deleteEdges, deleteNodes, nodes, edges } =
const store = useNodeStoreRef();
const { createNode, addNode, addEdge, deleteEdges, deleteNodes } =
useNodes((s) => ({
createNode: s.createNode,
addNode: s.addNode,
addEdge: s.addEdge,
deleteEdges: s.deleteEdges,
deleteNodes: s.deleteNodes,
nodes: s.nodes,
edges: s.edges
deleteNodes: s.deleteNodes
}), shallow);

return useCallback(
Expand All @@ -195,6 +194,7 @@ export function useGroupIntoSubgraph() {
return null;
}

const { nodes, edges } = store.getState();
const idSet = new Set(selectedIds);
const selectedNodes = nodes.filter((n) => idSet.has(n.id)) as RFNode<NodeData>[];
if (selectedNodes.length === 0) return null;
Expand Down Expand Up @@ -239,8 +239,7 @@ export function useGroupIntoSubgraph() {
addEdge,
deleteEdges,
deleteNodes,
nodes,
edges
store
]
);
}
Expand Down
6 changes: 3 additions & 3 deletions web/src/hooks/nodes/useNodeContextMenu.ts
Original file line number Diff line number Diff line change
Expand Up @@ -59,7 +59,7 @@ export function useNodeContextMenu(): UseNodeContextMenuReturn {
deleteNode,
getSelectedNodes,
toggleBypass,
nodes,
findNode,
setSelectedNodes
} = useNodes((state) => ({
updateNodeData: state.updateNodeData,
Expand All @@ -68,11 +68,11 @@ export function useNodeContextMenu(): UseNodeContextMenuReturn {
deleteNode: state.deleteNode,
getSelectedNodes: state.getSelectedNodes,
toggleBypass: state.toggleBypass,
nodes: state.nodes,
findNode: state.findNode,
setSelectedNodes: state.setSelectedNodes
}), shallow);

const rawNode = nodeId ? nodes.find((n) => n.id === nodeId) : undefined;
const rawNode = nodeId ? findNode(nodeId) : undefined;
const node = rawNode as Node<NodeData> | null;
const nodeData = node?.data;
const { writeClipboard } = useClipboard();
Expand Down
Loading
Loading