From bedb6fedd6cdddcf96b938f7c09b8be8d95f1b6e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9E=97=E7=8E=AE=20=28Jade=20Lin=29?= Date: Mon, 10 Aug 2026 20:37:34 +0800 Subject: [PATCH] refactor(workflow): align standalone node execution layers (#40447) (cherry picked from commit 4eb9a24997c1425d917aa295fdbf71ad137033d6) --- api/core/workflow/workflow_entry.py | 51 ++++-- .../workflow/test_workflow_entry_helpers.py | 146 +++++++++++++++--- 2 files changed, 159 insertions(+), 38 deletions(-) diff --git a/api/core/workflow/workflow_entry.py b/api/core/workflow/workflow_entry.py index fb12922ed7f..f89159cd9fe 100644 --- a/api/core/workflow/workflow_entry.py +++ b/api/core/workflow/workflow_entry.py @@ -34,11 +34,11 @@ from graphon.filters import GraphEventFilterContext, ResponseStreamFilter, filte from graphon.graph import Graph from graphon.graph_engine import GraphEngine, GraphEngineConfig from graphon.graph_engine.command_channels import CommandChannel, InMemoryChannel -from graphon.graph_engine.layers import DebugLoggingLayer, ExecutionLimitsLayer -from graphon.graph_events import GraphEngineEvent, GraphNodeEventBase, GraphRunFailedEvent +from graphon.graph_engine.layers import DebugLoggingLayer, ExecutionLimitsLayer, GraphEngineLayer +from graphon.graph_events import GraphEngineEvent, GraphNodeEventBase, GraphRunFailedEvent, is_node_result_event from graphon.nodes import BuiltinNodeTypes from graphon.nodes.base.node import Node -from graphon.runtime import ChildGraphNotFoundError, GraphRuntimeState, VariablePool +from graphon.runtime import ChildGraphNotFoundError, GraphRuntimeState, ReadOnlyGraphRuntimeStateWrapper, VariablePool from graphon.variable_loader import DUMMY_VARIABLE_LOADER, VariableLoader, load_into_variable_pool from models.workflow import Workflow @@ -358,7 +358,7 @@ class WorkflowEntry: node = node_factory.create_node(node_config) try: - generator = cls._traced_node_run(node) + generator = cls._run_node_with_layers(node, tenant_id=workflow.tenant_id) except Exception as e: logger.exception( "error while running node, workflow_id=%s, node_id=%s, node_type=%s, node_version=%s", @@ -497,7 +497,7 @@ class WorkflowEntry: tenant_id=tenant_id, ) - generator = cls._traced_node_run(node) + generator = cls._run_node_with_layers(node, tenant_id=tenant_id) return node, generator except Exception as e: @@ -613,24 +613,49 @@ class WorkflowEntry: variable_pool.add([variable_node_id] + variable_key_list, input_value) @staticmethod - def _traced_node_run(node: Node) -> Generator[GraphNodeEventBase, None, None]: + def _run_node_with_layers(node: Node, *, tenant_id: str) -> Generator[GraphNodeEventBase, None, None]: """ - Wraps a node's run method with OpenTelemetry tracing and returns a generator. + Run a standalone node with the same quota and observability hooks as GraphEngine. """ - # Wrap node.run() with ObservabilityLayer hooks to produce node-level spans - layer = ObservabilityLayer() - layer.on_graph_start() + layers: Sequence[GraphEngineLayer] = ( + LLMQuotaLayer(tenant_id=tenant_id), + ObservabilityLayer(), + ) + command_channel = InMemoryChannel() + runtime_state = ReadOnlyGraphRuntimeStateWrapper(node.graph_runtime_state) + for layer in layers: + layer.initialize(runtime_state, command_channel) + layer.on_graph_start() + node.ensure_execution_id() def _gen(): error: Exception | None = None - layer.on_node_run_start(node) + result_event: GraphNodeEventBase | None = None + layers_finished = False + + def finish_layers() -> None: + nonlocal layers_finished + if layers_finished: + return + layers_finished = True + for layer in layers: + layer.on_node_run_end(node, error, result_event) + for layer in layers: + layer.on_graph_end(error) + try: - yield from node.run() + for layer in layers: + layer.on_node_run_start(node) + for event in node.run(): + if isinstance(event, GraphNodeEventBase) and is_node_result_event(event): + result_event = event + finish_layers() + yield event except Exception as exc: error = exc raise finally: - layer.on_node_run_end(node, error) + finish_layers() return _gen() diff --git a/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py b/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py index 41037233b8c..45673f13451 100644 --- a/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py +++ b/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py @@ -1,5 +1,7 @@ +import threading from collections import UserString from contextlib import nullcontext +from datetime import datetime from types import SimpleNamespace from unittest.mock import MagicMock, patch, sentinel @@ -15,7 +17,7 @@ from graphon.errors import WorkflowNodeRunFailedError from graphon.file import File, FileTransferMethod, FileType from graphon.filters import ResponseStreamFilter from graphon.graph import Graph -from graphon.graph_events import GraphRunFailedEvent +from graphon.graph_events import GraphRunFailedEvent, NodeRunSucceededEvent from graphon.model_runtime.entities.llm_entities import LLMMode, LLMUsage from graphon.node_events import NodeRunResult from graphon.nodes import BuiltinNodeTypes @@ -482,7 +484,7 @@ class TestWorkflowEntrySingleStepRun: patch.object(workflow_entry.WorkflowEntry, "mapping_user_inputs_to_variable_pool"), patch.object( workflow_entry.WorkflowEntry, - "_traced_node_run", + "_run_node_with_layers", return_value=iter(["event"]), ), ): @@ -545,7 +547,7 @@ class TestWorkflowEntrySingleStepRun: ) as mapping_user_inputs_to_variable_pool, patch.object( workflow_entry.WorkflowEntry, - "_traced_node_run", + "_run_node_with_layers", return_value=iter(["event"]), ), ): @@ -614,7 +616,7 @@ class TestWorkflowEntrySingleStepRun: ) as mapping_user_inputs_to_variable_pool, patch.object( workflow_entry.WorkflowEntry, - "_traced_node_run", + "_run_node_with_layers", return_value=iter(["event"]), ), ): @@ -645,7 +647,7 @@ class TestWorkflowEntrySingleStepRun: ) mapping_user_inputs_to_variable_pool.assert_not_called() - def test_wraps_traced_node_run_failures(self): + def test_wraps_layered_node_run_failures(self): class FakeNode: id = "node-id" title = "Node Title" @@ -671,7 +673,7 @@ class TestWorkflowEntrySingleStepRun: patch.object(workflow_entry.WorkflowEntry, "mapping_user_inputs_to_variable_pool"), patch.object( workflow_entry.WorkflowEntry, - "_traced_node_run", + "_run_node_with_layers", side_effect=RuntimeError("boom"), ), ): @@ -790,7 +792,7 @@ class TestWorkflowEntryHelpers: ) as mapping_user_inputs_to_variable_pool, patch.object( workflow_entry.WorkflowEntry, - "_traced_node_run", + "_run_node_with_layers", return_value=iter(["event"]), ), ): @@ -953,32 +955,57 @@ class TestMappingUserInputsBranches: ) -class TestWorkflowEntryTracing: - def test_traced_node_run_reports_success(self): - layer = MagicMock() +class TestWorkflowEntryNodeLayers: + def test_run_node_with_layers_reports_success(self): + quota_layer = MagicMock() + observability_layer = MagicMock() + result_event = NodeRunSucceededEvent( + id="execution-id", + node_id="node-id", + node_type=BuiltinNodeTypes.START, + start_at=datetime.now(), + node_run_result=NodeRunResult(status=WorkflowNodeExecutionStatus.SUCCEEDED), + ) class FakeNode: + graph_runtime_state = sentinel.graph_runtime_state + def ensure_execution_id(self): return None def run(self): - yield "event" + yield result_event - with patch.object(workflow_entry, "ObservabilityLayer", return_value=layer): - events = list(workflow_entry.WorkflowEntry._traced_node_run(FakeNode())) + node = FakeNode() + with ( + patch.object(workflow_entry, "LLMQuotaLayer", return_value=quota_layer) as quota_layer_cls, + patch.object(workflow_entry, "ObservabilityLayer", return_value=observability_layer), + patch.object(workflow_entry, "InMemoryChannel", return_value=sentinel.command_channel), + patch.object( + workflow_entry, + "ReadOnlyGraphRuntimeStateWrapper", + return_value=sentinel.read_only_runtime_state, + ) as runtime_state_wrapper, + ): + events = list(workflow_entry.WorkflowEntry._run_node_with_layers(node, tenant_id="tenant-id")) - assert events == ["event"] - layer.on_graph_start.assert_called_once_with() - layer.on_node_run_start.assert_called_once() - layer.on_node_run_end.assert_called_once_with( - layer.on_node_run_start.call_args.args[0], - None, - ) + assert events == [result_event] + quota_layer_cls.assert_called_once_with(tenant_id="tenant-id") + runtime_state_wrapper.assert_called_once_with(sentinel.graph_runtime_state) + for layer in (quota_layer, observability_layer): + layer.initialize.assert_called_once_with(sentinel.read_only_runtime_state, sentinel.command_channel) + layer.on_graph_start.assert_called_once_with() + layer.on_node_run_start.assert_called_once_with(node) + layer.on_node_run_end.assert_called_once_with(node, None, result_event) + layer.on_graph_end.assert_called_once_with(None) - def test_traced_node_run_reports_errors(self): - layer = MagicMock() + def test_run_node_with_layers_reports_errors(self): + quota_layer = MagicMock() + observability_layer = MagicMock() class FakeNode: + graph_runtime_state = sentinel.graph_runtime_state + def ensure_execution_id(self): return None @@ -986,8 +1013,77 @@ class TestWorkflowEntryTracing: raise RuntimeError("boom") yield - with patch.object(workflow_entry, "ObservabilityLayer", return_value=layer): + node = FakeNode() + with ( + patch.object(workflow_entry, "LLMQuotaLayer", return_value=quota_layer), + patch.object(workflow_entry, "ObservabilityLayer", return_value=observability_layer), + patch.object( + workflow_entry, + "ReadOnlyGraphRuntimeStateWrapper", + return_value=sentinel.read_only_runtime_state, + ), + ): with pytest.raises(RuntimeError, match="boom"): - list(workflow_entry.WorkflowEntry._traced_node_run(FakeNode())) + list(workflow_entry.WorkflowEntry._run_node_with_layers(node, tenant_id="tenant-id")) - assert isinstance(layer.on_node_run_end.call_args.args[1], RuntimeError) + for layer in (quota_layer, observability_layer): + assert layer.on_node_run_end.call_args.args[0] is node + assert isinstance(layer.on_node_run_end.call_args.args[1], RuntimeError) + assert layer.on_node_run_end.call_args.args[2] is None + assert isinstance(layer.on_graph_end.call_args.args[0], RuntimeError) + + def test_run_node_with_layers_deducts_llm_quota(self): + result_event = NodeRunSucceededEvent( + id="execution-id", + node_id="node-id", + node_type=BuiltinNodeTypes.LLM, + start_at=datetime.now(), + node_run_result=NodeRunResult( + status=WorkflowNodeExecutionStatus.SUCCEEDED, + inputs={"model_provider": "openai", "model_name": "gpt-4o"}, + llm_usage=LLMUsage.empty_usage(), + ), + ) + + class FakeNode: + id = "node-id" + node_type = BuiltinNodeTypes.LLM + graph_runtime_state = SimpleNamespace( + stop_event=threading.Event(), + variable_pool=VariablePool(), + ) + node_data = SimpleNamespace( + model=SimpleNamespace(provider="openai", name="gpt-4o"), + error_strategy=None, + retry_config=SimpleNamespace(retry_enabled=False), + ) + + def ensure_execution_id(self): + return "execution-id" + + def run(self): + yield result_event + + with ( + patch.object(workflow_entry, "ObservabilityLayer", return_value=MagicMock()), + patch( + "core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", + autospec=True, + ) as ensure_quota, + patch( + "core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", + autospec=True, + ) as deduct_quota, + ): + generator = workflow_entry.WorkflowEntry._run_node_with_layers(FakeNode(), tenant_id="tenant-id") + event = next(generator) + + assert event is result_event + ensure_quota.assert_called_once_with(tenant_id="tenant-id", provider="openai", model="gpt-4o") + deduct_quota.assert_called_once_with( + tenant_id="tenant-id", + provider="openai", + model="gpt-4o", + usage=result_event.node_run_result.llm_usage, + ) + generator.close()