refactor(workflow): align standalone node execution layers (#40447)

This commit is contained in:
林玮 (Jade Lin) 2026-08-10 20:37:34 +08:00 committed by GitHub
parent b206e691ca
commit 4eb9a24997
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 164 additions and 39 deletions

View File

@ -34,12 +34,12 @@ 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.nodes.container_effects import ContainerAwaitRequest
from graphon.runtime import GraphRuntimeState, VariablePool
from graphon.runtime import GraphRuntimeState, ReadOnlyGraphRuntimeStateWrapper, VariablePool
from graphon.variable_loader import DUMMY_VARIABLE_LOADER, VariableLoader, load_into_variable_pool
from models.workflow import Workflow
@ -289,7 +289,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",
@ -428,7 +428,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:
@ -544,24 +544,51 @@ class WorkflowEntry:
variable_pool.add([variable_node_id] + variable_key_list, input_value)
@staticmethod
def _traced_node_run(node: Node) -> Generator[GraphNodeEventBase | ContainerAwaitRequest, None, None]:
def _run_node_with_layers(
node: Node, *, tenant_id: str
) -> Generator[GraphNodeEventBase | ContainerAwaitRequest, 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.bind_execution_id(str(uuid4()))
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()

View File

@ -1,4 +1,6 @@
import threading
from collections import UserString
from datetime import datetime
from types import SimpleNamespace
from unittest.mock import MagicMock, patch, sentinel
@ -9,11 +11,13 @@ from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
from core.workflow import workflow_entry
from core.workflow.system_variables import default_system_variables
from graphon.entities.base_node_data import BaseNodeData
from graphon.enums import NodeType
from graphon.enums import NodeType, WorkflowNodeExecutionStatus
from graphon.errors import WorkflowNodeRunFailedError
from graphon.file import File, FileTransferMethod, FileType
from graphon.filters import ResponseStreamFilter
from graphon.graph_events import GraphRunFailedEvent
from graphon.graph_events import GraphRunFailedEvent, NodeRunSucceededEvent
from graphon.model_runtime.entities.llm_entities import LLMUsage
from graphon.node_events import NodeRunResult
from graphon.nodes import BuiltinNodeTypes
from graphon.runtime import VariablePool
from graphon.variables.variables import StringVariable
@ -277,7 +281,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"]),
),
):
@ -340,7 +344,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"]),
),
):
@ -409,7 +413,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"]),
),
):
@ -440,7 +444,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"
@ -466,7 +470,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"),
),
):
@ -585,7 +589,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"]),
),
):
@ -748,32 +752,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 bind_execution_id(self, _execution_id):
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 bind_execution_id(self, _execution_id):
return None
@ -781,8 +810,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 bind_execution_id(self, execution_id):
self.execution_id = 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()