Files
DevFlow/crates/df-workflow/src/executor.rs

499 lines
20 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! DAG 执行器 — 按拓扑层级调度执行
use std::collections::HashMap;
use df_types::events::WorkflowEvent;
use df_types::types::NodeId;
use crate::dag::Dag;
use crate::eventbus::EventBus;
use crate::node::{NodeContext, NodeOutput};
use crate::state::StateMachine;
/// DAG 执行器
pub struct DagExecutor {
/// 事件总线
event_bus: EventBus,
/// 节点状态机
state_machine: StateMachine,
/// 工作流执行 ID由调用方传入下沉到每个 NodeContext
execution_id: String,
}
impl DagExecutor {
/// 创建执行器
///
/// `execution_id` 为本次工作流执行的唯一标识,会下沉到每个节点的 NodeContext
/// 用于节点内的事件关联、审计追踪等。
pub fn new(event_bus: EventBus, execution_id: String) -> Self {
Self {
event_bus,
state_machine: StateMachine::new(),
execution_id,
}
}
/// 返回状态机的共享引用Arc clone与 NodeContext.node_status 共享底层 HashMap
///
/// run_workflow 创建执行器后将其注册到 AppState 全局表,
/// cancel_workflow_node IPC 经 execution_id 取出此引用调 set_cancelled
/// 写入直达运行中阻塞节点HumanNode的 is_cancelled 轮询。
pub fn state_machine(&self) -> StateMachine {
self.state_machine.clone()
}
/// 执行 DAG 工作流
///
/// 同一拓扑层内的节点并发执行,层与层之间串行;
/// 本层全部节点完成后统一更新状态,任一失败则中止后续层。
pub async fn run(
&mut self,
dag: &Dag,
initial_config: serde_json::Value,
) -> anyhow::Result<HashMap<NodeId, NodeOutput>> {
let layers = dag.topological_layers()?;
let mut outputs: HashMap<NodeId, NodeOutput> = HashMap::new();
let start = std::time::Instant::now();
tracing::info!("DAG 执行开始,共 {} 层", layers.len());
// 预建入边索引target → 直接前驱列表O(E) 一次构建)
// 避免在内层按节点循环中调用 dag.predecessors()(每次 O(E) 全表扫描,整体 O(V·E))。
let mut adjacency_in: HashMap<NodeId, Vec<NodeId>> = HashMap::new();
for edge in &dag.edges {
adjacency_in
.entry(edge.target.clone())
.or_default()
.push(edge.source.clone());
}
for (layer_idx, layer) in layers.iter().enumerate() {
tracing::info!("执行第 {} 层,共 {} 个节点", layer_idx, layer.len());
// 阶段一:逐节点发 NodeStarted、置为运行中并构建执行 future
let mut node_futures = Vec::with_capacity(layer.len());
for node_id in layer {
let node = dag.nodes.get(node_id).ok_or_else(|| {
anyhow::anyhow!("节点 {} 不存在于 DAG 中", node_id)
})?;
// 发送 NodeStarted 事件
self.event_bus
.send(WorkflowEvent::NodeStarted {
node_id: node_id.clone(),
})
.await;
self.state_machine.set_running(node_id.clone())?;
// 构建节点上下文:通过入边索引 O(入度) 取前驱,而非 O(E) 全表扫描
let mut inputs = HashMap::new();
if let Some(preds) = adjacency_in.get(node_id) {
for pred_id in preds {
if let Some(out) = outputs.get(pred_id) {
inputs.insert(pred_id.clone(), out.clone());
}
}
}
let ctx = NodeContext {
node_id: node_id.clone(),
inputs,
// 节点级覆盖全局级deep_merge(initial_config, node_configs[id])。
// 节点 NodeDef.config 未定义node_configs 缺该 key→ 直接用全局 initial_config
// 与旧行为一致(零行为破坏现有 HumanNode/AiNode 等读全局 config 的调用方)。
config: match dag.node_configs.get(node_id) {
Some(node_cfg) => crate::dag::deep_merge(&initial_config, node_cfg),
None => initial_config.clone(),
},
execution_id: self.execution_id.clone(),
event_bus: self.event_bus.clone(),
// 共享执行器状态机,使节点(如 HumanNode能读取真实状态而非空状态机
node_status: self.state_machine.clone(),
};
// 仅捕获节点共享引用与所有权上下文,避免与 self 借用冲突
let id = node_id.clone();
node_futures.push(async move {
let node_start = std::time::Instant::now();
let result = node.execute(ctx).await;
(id, result, node_start.elapsed().as_millis() as u64)
});
}
// 阶段二:同层节点并发执行,等待全部完成
let results = futures::future::join_all(node_futures).await;
// 阶段三:统一更新状态并发送事件,任一失败则整体返回 Err
let mut first_err: Option<anyhow::Error> = None;
for (node_id, result, duration) in results {
match result {
Ok(output) => {
// 已取消的节点跳过 set_completed:
// Ok 后 join_all 让出执行权,窗口内 cancel_workflow_node IPC → set_cancelled 改 Cancelled,
// 随后此处 set_completed 走 transition 命中(Cancelled→Completed)非法 → bail 致已批准审批报失败(R-P1-3 TOCTOU)
// 与 Err 分支 is_cancelled 短路对称:已取消节点保持 Cancelled 终态
if !self.state_machine.is_cancelled(&node_id) {
self.state_machine.set_completed(node_id.clone())?;
}
self.event_bus
.send(WorkflowEvent::NodeCompleted {
node_id: node_id.clone(),
duration_ms: duration,
})
.await;
outputs.insert(node_id, output);
}
Err(e) => {
let error_msg = e.to_string();
// 已取消的节点(如 HumanNode 审批取消)跳过 set_failed:
// Cancelled 已是终态,transition(Cancelled→Failed) 非法会 bail 致工作流崩溃(B-03b-R1)
if self.state_machine.is_cancelled(&node_id) {
// emit NodeCancelled(非 NodeFailed):取消语义有别于失败,前端按 type 归「取消」非「失败」
self.event_bus
.send(WorkflowEvent::NodeCancelled {
node_id: node_id.clone(),
})
.await;
} else {
self.state_machine.set_failed(node_id.clone())?;
self.event_bus
.send(WorkflowEvent::NodeFailed {
node_id: node_id.clone(),
error: error_msg,
})
.await;
}
// 取消与失败一致:节点返回 Err 即中止后续层(run 返回 Err);
// 仅事件类型区分,工作流结果语义保持不变(零行为破坏)。
if first_err.is_none() {
first_err = Some(e.context(format!("节点 {} 执行失败", node_id)));
}
}
}
}
if let Some(e) = first_err {
return Err(e);
}
}
let total = start.elapsed().as_millis() as u64;
self.event_bus
.send(WorkflowEvent::WorkflowCompleted {
total_duration_ms: total,
})
.await;
tracing::info!("DAG 执行完成,耗时 {}ms", total);
Ok(outputs)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::node::{Node, NodeResult, NodeSchema};
use async_trait::async_trait;
use std::time::Duration;
/// 测试节点sleep 指定毫秒后返回空输出
struct SleepNode {
sleep_ms: u64,
}
#[async_trait]
impl Node for SleepNode {
async fn execute(&self, _ctx: NodeContext) -> NodeResult {
tokio::time::sleep(Duration::from_millis(self.sleep_ms)).await;
Ok(NodeOutput::empty())
}
fn schema(&self) -> NodeSchema {
NodeSchema {
params: serde_json::Value::Null,
output: serde_json::Value::Null,
}
}
fn node_type(&self) -> &str {
"sleep"
}
}
/// 测试节点:直接返回错误
struct FailNode;
/// 测试节点:把 ctx.config 的指定 key 字符串原样回显到 output验证节点级 config 下沉)
struct EchoConfigNode {
key: String,
}
#[async_trait]
impl Node for EchoConfigNode {
async fn execute(&self, ctx: NodeContext) -> NodeResult {
let v = ctx
.config
.get(&self.key)
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
Ok(NodeOutput::from_value(serde_json::json!({ "echo": v })))
}
fn schema(&self) -> NodeSchema {
NodeSchema {
params: serde_json::Value::Null,
output: serde_json::Value::Null,
}
}
fn node_type(&self) -> &str {
"echo_config"
}
}
#[async_trait]
impl Node for FailNode {
async fn execute(&self, _ctx: NodeContext) -> NodeResult {
Err(anyhow::anyhow!("故意失败"))
}
fn schema(&self) -> NodeSchema {
NodeSchema {
params: serde_json::Value::Null,
output: serde_json::Value::Null,
}
}
fn node_type(&self) -> &str {
"fail"
}
}
/// ④-1:节点级 config 下沉验证。
/// DagDef 节点 config={foo:"bar"} + 全局 config={foo:"GLOBAL"} → NodeContext.config.foo=="bar"
/// (节点级覆盖全局级)。同时验证缺节点 config 的节点回退全局 config零行为破坏现有调用方
#[tokio::test]
async fn node_config_overrides_global_in_node_context() {
use crate::registry::NodeRegistry;
// 构造 DagDefn1 节点写 config={foo:"bar"}n2 节点 config 空(验证回退全局)
let mut def = crate::dag_def::DagDef::new();
def.add_node("n1".to_string(), "echo".to_string(), serde_json::json!({ "foo": "bar" }));
def.add_node("n2".to_string(), "echo".to_string(), serde_json::json!({}));
// 注册 echo 工厂 → EchoConfigNode读 config.foo
let mut registry = NodeRegistry::new();
registry.register("echo", |_cfg| {
Box::new(EchoConfigNode { key: "foo".to_string() })
});
let dag = registry.build_dag(&def).expect("build_dag");
let mut executor = DagExecutor::new(EventBus::new(), "test-nodecfg".to_string());
// 全局 config 写 foo=GLOBAL验证被 n1 节点级覆盖n2 回退用全局)
let outputs = executor
.run(&dag, serde_json::json!({ "foo": "GLOBAL" }))
.await
.expect("run");
// n1:节点级 foo="bar" 覆盖全局 "GLOBAL"
assert_eq!(
outputs["n1"].data["echo"],
serde_json::json!("bar"),
"节点级 config 应覆盖全局"
);
// n2:节点 config 空 → 回退全局 foo="GLOBAL"(零行为破坏现有调用方)
assert_eq!(
outputs["n2"].data["echo"],
serde_json::json!("GLOBAL"),
"空节点 config 应回退全局"
);
}
#[tokio::test]
async fn test_same_layer_runs_in_parallel() {
// 两个无依赖节点位于同一层,各 sleep 100ms
let mut dag = Dag::new();
dag.add_node("a".to_string(), Box::new(SleepNode { sleep_ms: 100 }));
dag.add_node("b".to_string(), Box::new(SleepNode { sleep_ms: 100 }));
let mut executor = DagExecutor::new(EventBus::new(), "test-exec".to_string());
let start = std::time::Instant::now();
let outputs = executor.run(&dag, serde_json::Value::Null).await.unwrap();
let elapsed = start.elapsed();
assert_eq!(outputs.len(), 2);
// 串行需要约 200ms并行应明显小于 180ms
assert!(
elapsed < Duration::from_millis(180),
"同层节点应并行执行,实际耗时 {:?}",
elapsed
);
}
#[tokio::test]
async fn test_layer_failure_aborts_following_layers() {
// a(失败) 与 b(成功) 同层c 依赖 a失败后 c 不应执行
let mut dag = Dag::new();
dag.add_node("a".to_string(), Box::new(FailNode));
dag.add_node("b".to_string(), Box::new(SleepNode { sleep_ms: 10 }));
dag.add_node("c".to_string(), Box::new(SleepNode { sleep_ms: 10 }));
dag.add_edge("a".to_string(), "c".to_string());
let mut executor = DagExecutor::new(EventBus::new(), "test-exec".to_string());
let err = executor
.run(&dag, serde_json::Value::Null)
.await
.unwrap_err();
assert!(err.to_string().contains("节点 a 执行失败"));
// 同层成功节点状态正常更新,下游节点保持 Pending
use df_types::types::NodeStatus;
assert_eq!(executor.state_machine.get(&"a".to_string()), NodeStatus::Failed);
assert_eq!(executor.state_machine.get(&"b".to_string()), NodeStatus::Completed);
assert_eq!(executor.state_machine.get(&"c".to_string()), NodeStatus::Pending);
}
/// 测试节点:执行中通过共享 node_status 自取消(模拟外部 IPC set_cancelled),返回 Err
struct CancelSelfNode;
#[async_trait]
impl Node for CancelSelfNode {
async fn execute(&self, ctx: NodeContext) -> NodeResult {
// 通过共享 node_status 置 Cancelled(写共享 HashMap,executor.state_machine 可见)
ctx.node_status.set_cancelled(ctx.node_id.clone());
Err(anyhow::anyhow!("人工审批被取消"))
}
fn schema(&self) -> NodeSchema {
NodeSchema {
params: serde_json::Value::Null,
output: serde_json::Value::Null,
}
}
fn node_type(&self) -> &str {
"cancel_self"
}
}
/// B-03b-R1:取消的节点返回 Err 时,executor 跳过 set_failed(Cancelled→Failed transition 非法会 bail),
/// 状态保持 Cancelled,run 返回取消相关 Err 而非状态转换错误。
#[tokio::test]
async fn test_cancelled_node_skips_set_failed() {
use df_types::types::NodeStatus;
let mut dag = Dag::new();
dag.add_node("x".to_string(), Box::new(CancelSelfNode));
let mut executor = DagExecutor::new(EventBus::new(), "test-cancel".to_string());
let result = executor.run(&dag, serde_json::Value::Null).await;
// run 返回 Err(取消致中止后续层),但不 panic/bail transition
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(
err.contains("取消") || err.contains("执行失败"),
"应返回取消相关错误,实际: {}",
err
);
// 状态保持 Cancelled(未被 set_failed 覆盖为 Failed)—— R1 修法生效的核心验证
assert_eq!(
executor.state_machine.get(&"x".to_string()),
NodeStatus::Cancelled
);
}
/// 测试节点:执行中通过共享 node_status 自取消(模拟外部 IPC set_cancelled),返回 Ok
/// 模拟 HumanNode select! 收合法 HumanApprovalResponse 后 return Ok,但 join_all 让出窗口内
/// cancel_workflow_node IPC 把节点改 Cancelled 的 TOCTOU 场景。
struct CancelSelfThenOkNode;
#[async_trait]
impl Node for CancelSelfThenOkNode {
async fn execute(&self, ctx: NodeContext) -> NodeResult {
ctx.node_status.set_cancelled(ctx.node_id.clone());
Ok(NodeOutput::empty())
}
fn schema(&self) -> NodeSchema {
NodeSchema {
params: serde_json::Value::Null,
output: serde_json::Value::Null,
}
}
fn node_type(&self) -> &str {
"cancel_self_then_ok"
}
}
/// R-P1-3:Ok 后取消的 TOCTOU 防护。节点返回 Ok 但执行期间已被 set_cancelled,
/// executor Ok 分支必须跳过 set_completed(transition Cancelled→Completed 非法会 bail),
/// 状态保持 Cancelled,run 返回 Ok(已批准审批不应误报失败)。
#[tokio::test]
async fn test_cancelled_node_skips_set_completed() {
use df_types::types::NodeStatus;
let mut dag = Dag::new();
dag.add_node("y".to_string(), Box::new(CancelSelfThenOkNode));
let mut executor = DagExecutor::new(EventBus::new(), "test-cancel-ok".to_string());
let result = executor.run(&dag, serde_json::Value::Null).await;
// run 返回 Ok —— Ok 节点不应因 TOCTOU 取消被误判为失败
assert!(
result.is_ok(),
"Ok 后取消应短路 set_completed 而非 bail,实际 err: {:?}",
result.err()
);
// 状态保持 Cancelled(未被 set_completed 覆盖为 Completed,也不 bail)—— R-P1-3 修法核心验证
assert_eq!(
executor.state_machine.get(&"y".to_string()),
NodeStatus::Cancelled
);
}
/// 复核-新③:取消节点(Err 路径)emit 的应是 NodeCancelled(非 NodeFailed),
/// 前端按 type 区分「取消」与「失败」。订阅事件总线收集所有事件断言。
#[tokio::test]
async fn test_cancelled_node_emits_node_cancelled_event() {
use df_types::events::WorkflowEvent;
use df_types::types::NodeStatus;
let mut dag = Dag::new();
dag.add_node("c".to_string(), Box::new(CancelSelfNode));
let bus = EventBus::new();
let mut rx = bus.subscribe();
let mut executor = DagExecutor::new(bus, "test-cancel-event".to_string());
let _ = executor.run(&dag, serde_json::Value::Null).await;
// 收集所有事件(run 已结束,事件总线无新事件)
let mut events: Vec<WorkflowEvent> = Vec::new();
while let Ok(ev) = rx.try_recv() {
events.push(ev);
}
// 必含 NodeCancelled { node_id: "c" }
let cancelled = events.iter().any(|e| matches!(
e,
WorkflowEvent::NodeCancelled { node_id } if node_id == "c"
));
assert!(cancelled, "取消节点应 emit NodeCancelled, 实际事件: {:?}", events);
// 不应含 node "c" 的 NodeFailed(失败事件)—— 语义双标修复的核心验证
let failed = events.iter().any(|e| matches!(
e,
WorkflowEvent::NodeFailed { node_id, .. } if node_id == "c"
));
assert!(!failed, "取消节点不应 emit NodeFailed, 实际事件: {:?}", events);
// 状态保持 Cancelled
assert_eq!(
executor.state_machine.get(&"c".to_string()),
NodeStatus::Cancelled
);
}
}