mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 02:43:49 +08:00
refactor(workflow): align standalone node execution layers (#40447)
(cherry picked from commit 4eb9a24997)
This commit is contained in:
parent
f99d88890f
commit
bedb6fedd6
@ -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()
|
||||
|
||||
@ -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()
|
||||
|
||||
Loading…
Reference in New Issue
Block a user