1#![forbid(unsafe_code)]
2#![warn(missing_docs)]
3
4use std::collections::{HashMap, HashSet};
7
8use petgraph::algo::is_cyclic_directed;
9use petgraph::graph::{DiGraph, NodeIndex};
10use petgraph::Direction;
11use substrate_core::error::{Result, SubstrateError};
12use substrate_core::workflow_port::{Workflow, WorkflowPort};
13
14#[derive(Debug, Default, Clone, Copy)]
16pub struct DagWorkflow;
17
18impl DagWorkflow {
19 pub fn new() -> Self {
21 Self
22 }
23
24 fn build_graph(
25 workflow: &Workflow,
26 ) -> Result<(DiGraph<String, ()>, HashMap<String, NodeIndex>)> {
27 let mut graph = DiGraph::new();
28 let mut index_of = HashMap::new();
29
30 for node in &workflow.nodes {
31 if index_of.contains_key(&node.id) {
32 return Err(SubstrateError::Other(format!(
33 "duplicate node id: {}",
34 node.id
35 )));
36 }
37 let idx = graph.add_node(node.id.clone());
38 index_of.insert(node.id.clone(), idx);
39 }
40
41 for edge in &workflow.edges {
42 let from = index_of.get(&edge.from).ok_or_else(|| {
43 SubstrateError::Other(format!("unknown edge source: {}", edge.from))
44 })?;
45 let to = index_of.get(&edge.to).ok_or_else(|| {
46 SubstrateError::Other(format!("unknown edge target: {}", edge.to))
47 })?;
48 graph.add_edge(*from, *to, ());
49 }
50
51 Ok((graph, index_of))
52 }
53}
54
55impl WorkflowPort for DagWorkflow {
56 fn validate_acyclic(&self, workflow: &Workflow) -> Result<()> {
57 let (graph, _) = Self::build_graph(workflow)?;
58 if is_cyclic_directed(&graph) {
59 return Err(SubstrateError::CycleDetected(
60 "workflow contains a cycle".into(),
61 ));
62 }
63 Ok(())
64 }
65
66 fn topological_order(&self, workflow: &Workflow) -> Result<Vec<String>> {
67 self.validate_acyclic(workflow)?;
68 let (graph, index_of) = Self::build_graph(workflow)?;
69
70 let mut in_degree: HashMap<NodeIndex, usize> =
71 graph.node_indices().map(|n| (n, 0)).collect();
72 for edge in graph.edge_indices() {
73 let (_, target) = graph.edge_endpoints(edge).unwrap();
74 *in_degree.get_mut(&target).unwrap() += 1;
75 }
76
77 let mut ready: Vec<NodeIndex> = in_degree
78 .iter()
79 .filter(|(_, °)| deg == 0)
80 .map(|(&n, _)| n)
81 .collect();
82 ready.sort_by_key(|n| graph[*n].clone());
83
84 let mut order = Vec::with_capacity(graph.node_count());
85 while let Some(n) = ready.first().cloned() {
86 ready.remove(0);
87 order.push(graph[n].clone());
88 for child in graph.neighbors_directed(n, Direction::Outgoing) {
89 let deg = in_degree.get_mut(&child).unwrap();
90 *deg -= 1;
91 if *deg == 0 {
92 ready.push(child);
93 ready.sort_by_key(|i| graph[*i].clone());
94 }
95 }
96 }
97
98 if order.len() != graph.node_count() {
99 return Err(SubstrateError::CycleDetected(
100 "topological sort incomplete — cycle present".into(),
101 ));
102 }
103 let _ = index_of;
104 Ok(order)
105 }
106
107 fn ready_set(&self, workflow: &Workflow, completed: &[String]) -> Result<Vec<String>> {
108 self.validate_acyclic(workflow)?;
109 let (graph, index_of) = Self::build_graph(workflow)?;
110 let done: HashSet<&str> = completed.iter().map(String::as_str).collect();
111
112 let mut ready = Vec::new();
113 for node in &workflow.nodes {
114 if done.contains(node.id.as_str()) {
115 continue;
116 }
117 let idx = index_of[&node.id];
118 let mut preds = graph.neighbors_directed(idx, Direction::Incoming);
119 if preds.all(|p| done.contains(graph[p].as_str())) {
120 ready.push(node.id.clone());
121 }
122 }
123 ready.sort();
124 Ok(ready)
125 }
126}
127
128#[cfg(test)]
129mod tests {
130 use super::*;
131 use substrate_core::workflow_port::{WorkflowEdge, WorkflowNode};
132
133 fn wf(nodes: &[&str], edges: &[(&str, &str)]) -> Workflow {
134 Workflow {
135 nodes: nodes
136 .iter()
137 .map(|id| WorkflowNode { id: (*id).into() })
138 .collect(),
139 edges: edges
140 .iter()
141 .map(|(from, to)| WorkflowEdge {
142 from: (*from).into(),
143 to: (*to).into(),
144 })
145 .collect(),
146 }
147 }
148
149 #[test]
150 fn linear_chain_topological_order() {
151 let dag = DagWorkflow::new();
152 let w = wf(&["a", "b", "c"], &[("a", "b"), ("b", "c")]);
153 let order = dag.topological_order(&w).unwrap();
154 assert_eq!(order, vec!["a", "b", "c"]);
155 }
156
157 #[test]
158 fn diamond_ready_set_after_partial() {
159 let dag = DagWorkflow::new();
160 let w = wf(
161 &["a", "b", "c", "d"],
162 &[("a", "b"), ("a", "c"), ("b", "d"), ("c", "d")],
163 );
164 let ready0 = dag.ready_set(&w, &[]).unwrap();
165 assert_eq!(ready0, vec!["a"]);
166
167 let ready1 = dag.ready_set(&w, &["a".into()]).unwrap();
168 assert_eq!(ready1, vec!["b", "c"]);
169
170 let ready2 = dag
171 .ready_set(&w, &["a".into(), "b".into(), "c".into()])
172 .unwrap();
173 assert_eq!(ready2, vec!["d"]);
174 }
175
176 #[test]
177 fn cycle_rejected() {
178 let dag = DagWorkflow::new();
179 let w = wf(&["a", "b", "c"], &[("a", "b"), ("b", "c"), ("c", "a")]);
180 let err = dag.validate_acyclic(&w).unwrap_err();
181 assert!(matches!(err, SubstrateError::CycleDetected(_)));
182 }
183}