eidos_kernel/workflow/
validate.rs1use std::collections::{BTreeMap, BTreeSet};
2
3use super::TaskDag;
4
5#[derive(Debug, Clone, PartialEq, Eq)]
7pub enum DagValidationError {
8 Empty,
9 DuplicateNodeId(String),
10 MissingDependency { node_id: String, dependency: String },
11 Cycle { nodes: Vec<String> },
12}
13
14pub fn validate_dag(dag: &TaskDag) -> Result<Vec<String>, DagValidationError> {
18 if dag.nodes.is_empty() {
19 return Err(DagValidationError::Empty);
20 }
21
22 let mut ids = BTreeSet::new();
23 for node in &dag.nodes {
24 if !ids.insert(node.id.clone()) {
25 return Err(DagValidationError::DuplicateNodeId(node.id.clone()));
26 }
27 }
28
29 for node in &dag.nodes {
30 for dependency in &node.depends_on {
31 if !ids.contains(dependency) {
32 return Err(DagValidationError::MissingDependency {
33 node_id: node.id.clone(),
34 dependency: dependency.clone(),
35 });
36 }
37 }
38 }
39
40 let mut remaining_deps: BTreeMap<String, BTreeSet<String>> = dag
41 .nodes
42 .iter()
43 .map(|node| {
44 (
45 node.id.clone(),
46 node.depends_on.iter().cloned().collect::<BTreeSet<_>>(),
47 )
48 })
49 .collect();
50 let mut dependents: BTreeMap<String, BTreeSet<String>> = BTreeMap::new();
51 for node in &dag.nodes {
52 for dependency in &node.depends_on {
53 dependents
54 .entry(dependency.clone())
55 .or_default()
56 .insert(node.id.clone());
57 }
58 }
59
60 let mut ready = remaining_deps
61 .iter()
62 .filter_map(|(id, deps)| deps.is_empty().then_some(id.clone()))
63 .collect::<BTreeSet<_>>();
64 let mut order = Vec::with_capacity(dag.nodes.len());
65
66 while let Some(id) = ready.pop_first() {
67 if !remaining_deps.contains_key(&id) {
68 continue;
69 }
70 remaining_deps.remove(&id);
71 order.push(id.clone());
72
73 if let Some(children) = dependents.get(&id) {
74 for child in children {
75 if let Some(deps) = remaining_deps.get_mut(child) {
76 deps.remove(&id);
77 if deps.is_empty() {
78 ready.insert(child.clone());
79 }
80 }
81 }
82 }
83 }
84
85 if remaining_deps.is_empty() {
86 Ok(order)
87 } else {
88 Err(DagValidationError::Cycle {
89 nodes: remaining_deps.keys().cloned().collect(),
90 })
91 }
92}