Skip to main content

substrate_dag/
lib.rs

1#![forbid(unsafe_code)]
2#![warn(missing_docs)]
3
4//! DAG workflow orchestration via [`WorkflowPort`] and petgraph.
5
6use 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/// [`WorkflowPort`] backed by petgraph directed graphs.
15#[derive(Debug, Default, Clone, Copy)]
16pub struct DagWorkflow;
17
18impl DagWorkflow {
19    /// Create a new workflow engine.
20    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)| 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}