VerseFlow-Studio / src /verseflow /core /dag_engine.py
julius119's picture
Upload 23 files
63ca2a0 verified
Raw History Blame Contribute Delete
2.92 kB
"""
拓扑 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)