""" 拓扑 DAG 工作流执行引擎 Author: XiaoZhe (Commercial Contact: janejulius119@gmail.com / WeChat: julius119) """ import asyncio import logging from collections import defaultdict, deque from typing import Dict, Any, Set, Callable, Awaitable logger = logging.getLogger("VerseFlow.DAGEngine") class DAGNode: def __init__(self, node_id: str, handler: Callable[..., Awaitable[Dict[str, Any]]]): self.node_id = node_id self.handler = handler self.in_edges: Set[str] = set() self.out_edges: Set[str] = set() class DAGWorkflowEngine: """支持并发分裂与动态合并的拓扑 DAG 执行器""" def __init__(self): self.nodes: Dict[str, DAGNode] = {} def add_node(self, node_id: str, handler: Callable[..., Awaitable[Dict[str, Any]]]): if node_id not in self.nodes: self.nodes[node_id] = DAGNode(node_id, handler) def add_dependency(self, parent_id: str, child_id: str): if parent_id in self.nodes and child_id in self.nodes: self.nodes[parent_id].out_edges.add(child_id) self.nodes[child_id].in_edges.add(parent_id) async def execute(self, initial_inputs: Dict[str, Any]) -> Dict[str, Any]: in_degree = {nid: len(node.in_edges) for nid, node in self.nodes.items()} node_outputs: Dict[str, Dict[str, Any]] = defaultdict(dict) ready_queue = deque([nid for nid, deg in in_degree.items() if deg == 0]) running_tasks: Dict[str, asyncio.Task] = {} logger.info(f"开始执行 DAG 工作流, 节点总数: {len(self.nodes)}") while ready_queue or running_tasks: while ready_queue: nid = ready_queue.popleft() parent_data = { pid: node_outputs[pid] for pid in self.nodes[nid].in_edges } merged_inputs = {**initial_inputs, **parent_data} logger.info(f"启动 DAG 节点: {nid}") task = asyncio.create_task(self.nodes[nid].handler(merged_inputs)) running_tasks[nid] = task if not running_tasks: break done, _ = await asyncio.wait( running_tasks.values(), return_when=asyncio.FIRST_COMPLETED ) finished_nids = [] for nid, task in list(running_tasks.items()): if task in done: finished_nids.append(nid) node_outputs[nid] = await task logger.info(f"DAG 节点完成: {nid}") for child_id in self.nodes[nid].out_edges: in_degree[child_id] -= 1 if in_degree[child_id] == 0: ready_queue.append(child_id) for nid in finished_nids: del running_tasks[nid] return dict(node_outputs)