Skip to main content

oxigdal_workflow/dag/
graph.rs

1//! DAG construction and validation.
2
3use crate::error::{DagError, Result};
4use petgraph::Direction;
5use petgraph::graph::{DiGraph, NodeIndex};
6use petgraph::visit::EdgeRef;
7use serde::{Deserialize, Serialize};
8use std::collections::{HashMap, HashSet, VecDeque};
9use std::hash::{Hash, Hasher};
10
11/// A task node in the workflow DAG.
12#[derive(Debug, Clone, Serialize, Deserialize)]
13pub struct TaskNode {
14    /// Unique task identifier.
15    pub id: String,
16    /// Task name.
17    pub name: String,
18    /// Task description.
19    pub description: Option<String>,
20    /// Task configuration as JSON.
21    pub config: serde_json::Value,
22    /// Retry policy.
23    pub retry: RetryPolicy,
24    /// Timeout in seconds.
25    pub timeout_secs: Option<u64>,
26    /// Resource requirements.
27    pub resources: ResourceRequirements,
28    /// Custom metadata.
29    pub metadata: HashMap<String, String>,
30}
31
32/// Retry policy for task execution.
33#[derive(Debug, Clone, Serialize, Deserialize)]
34pub struct RetryPolicy {
35    /// Maximum number of retry attempts.
36    pub max_attempts: u32,
37    /// Delay between retries in milliseconds.
38    pub delay_ms: u64,
39    /// Backoff multiplier for exponential backoff.
40    pub backoff_multiplier: f64,
41    /// Maximum delay in milliseconds.
42    pub max_delay_ms: u64,
43}
44
45impl Default for RetryPolicy {
46    fn default() -> Self {
47        Self {
48            max_attempts: 3,
49            delay_ms: 1000,
50            backoff_multiplier: 2.0,
51            max_delay_ms: 60000,
52        }
53    }
54}
55
56/// Resource requirements for a task.
57#[derive(Debug, Clone, Serialize, Deserialize)]
58pub struct ResourceRequirements {
59    /// CPU cores required (can be fractional).
60    pub cpu_cores: f64,
61    /// Memory required in MB.
62    pub memory_mb: u64,
63    /// GPU required.
64    pub gpu: bool,
65    /// Disk space required in MB.
66    pub disk_mb: u64,
67    /// Custom resource requirements.
68    pub custom: HashMap<String, f64>,
69}
70
71impl Default for ResourceRequirements {
72    fn default() -> Self {
73        Self {
74            cpu_cores: 1.0,
75            memory_mb: 1024,
76            gpu: false,
77            disk_mb: 1024,
78            custom: HashMap::new(),
79        }
80    }
81}
82
83impl PartialEq for TaskNode {
84    fn eq(&self, other: &Self) -> bool {
85        self.id == other.id
86    }
87}
88
89impl Eq for TaskNode {}
90
91impl Hash for TaskNode {
92    fn hash<H: Hasher>(&self, state: &mut H) {
93        self.id.hash(state);
94    }
95}
96
97/// An edge representing a dependency between tasks.
98#[derive(Debug, Clone, Serialize, Deserialize)]
99pub struct TaskEdge {
100    /// Edge type (data dependency, control dependency, etc.).
101    pub edge_type: EdgeType,
102    /// Condition for edge activation.
103    pub condition: Option<String>,
104}
105
106/// Type of dependency edge.
107#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
108pub enum EdgeType {
109    /// Data dependency - output of one task is input to another.
110    Data,
111    /// Control dependency - one task must complete before another starts.
112    Control,
113    /// Conditional - edge is only followed if condition is met.
114    Conditional,
115}
116
117impl Default for TaskEdge {
118    fn default() -> Self {
119        Self {
120            edge_type: EdgeType::Control,
121            condition: None,
122        }
123    }
124}
125
126#[derive(Debug, Clone, Serialize, Deserialize)]
127
128/// Workflow DAG structure.
129pub struct WorkflowDag {
130    /// Underlying directed graph.
131    pub(crate) graph: DiGraph<TaskNode, TaskEdge>,
132    /// Mapping from task ID to node index.
133    pub(crate) task_map: HashMap<String, NodeIndex>,
134}
135
136impl WorkflowDag {
137    /// Create a new empty workflow DAG.
138    pub fn new() -> Self {
139        Self {
140            graph: DiGraph::new(),
141            task_map: HashMap::new(),
142        }
143    }
144
145    /// Add a task to the DAG.
146    pub fn add_task(&mut self, task: TaskNode) -> Result<NodeIndex> {
147        if self.task_map.contains_key(&task.id) {
148            return Err(
149                DagError::InvalidNode(format!("Task '{}' already exists in DAG", task.id)).into(),
150            );
151        }
152
153        let node_index = self.graph.add_node(task.clone());
154        self.task_map.insert(task.id.clone(), node_index);
155        Ok(node_index)
156    }
157
158    /// Add a dependency edge between two tasks.
159    pub fn add_dependency(
160        &mut self,
161        from_task_id: &str,
162        to_task_id: &str,
163        edge: TaskEdge,
164    ) -> Result<()> {
165        let from_idx = self
166            .task_map
167            .get(from_task_id)
168            .ok_or_else(|| DagError::invalid_node(from_task_id))?;
169
170        let to_idx = self
171            .task_map
172            .get(to_task_id)
173            .ok_or_else(|| DagError::invalid_node(to_task_id))?;
174
175        self.graph.add_edge(*from_idx, *to_idx, edge);
176        Ok(())
177    }
178
179    /// Get a task by ID.
180    pub fn get_task(&self, task_id: &str) -> Option<&TaskNode> {
181        self.task_map
182            .get(task_id)
183            .and_then(|idx| self.graph.node_weight(*idx))
184    }
185
186    /// Get a task by ID (mutable).
187    pub fn get_task_mut(&mut self, task_id: &str) -> Option<&mut TaskNode> {
188        self.task_map
189            .get(task_id)
190            .and_then(|idx| self.graph.node_weight_mut(*idx))
191    }
192
193    /// Get task dependencies (tasks that must complete before this task).
194    pub fn get_dependencies(&self, task_id: &str) -> Vec<String> {
195        if let Some(&idx) = self.task_map.get(task_id) {
196            self.graph
197                .edges_directed(idx, Direction::Incoming)
198                .filter_map(|edge| {
199                    self.graph
200                        .node_weight(edge.source())
201                        .map(|task| task.id.clone())
202                })
203                .collect()
204        } else {
205            Vec::new()
206        }
207    }
208
209    /// Get task dependents (tasks that depend on this task).
210    pub fn get_dependents(&self, task_id: &str) -> Vec<String> {
211        if let Some(&idx) = self.task_map.get(task_id) {
212            self.graph
213                .edges_directed(idx, Direction::Outgoing)
214                .filter_map(|edge| {
215                    self.graph
216                        .node_weight(edge.target())
217                        .map(|task| task.id.clone())
218                })
219                .collect()
220        } else {
221            Vec::new()
222        }
223    }
224
225    /// Validate the DAG structure.
226    pub fn validate(&self) -> Result<()> {
227        // Check if DAG is empty
228        if self.graph.node_count() == 0 {
229            return Err(DagError::EmptyDag.into());
230        }
231
232        // Check for cycles
233        self.check_cycles()?;
234
235        // Check for unreachable nodes
236        self.check_reachability()?;
237
238        Ok(())
239    }
240
241    /// Check for cycles in the DAG using DFS.
242    fn check_cycles(&self) -> Result<()> {
243        let mut visited = HashSet::new();
244        let mut rec_stack = HashSet::new();
245
246        for node_idx in self.graph.node_indices() {
247            if !visited.contains(&node_idx)
248                && let Some(cycle_path) =
249                    self.dfs_cycle_check(node_idx, &mut visited, &mut rec_stack)
250            {
251                return Err(DagError::cycle(cycle_path).into());
252            }
253        }
254
255        Ok(())
256    }
257
258    /// DFS-based cycle detection.
259    fn dfs_cycle_check(
260        &self,
261        node: NodeIndex,
262        visited: &mut HashSet<NodeIndex>,
263        rec_stack: &mut HashSet<NodeIndex>,
264    ) -> Option<String> {
265        visited.insert(node);
266        rec_stack.insert(node);
267
268        for neighbor in self.graph.neighbors(node) {
269            if !visited.contains(&neighbor) {
270                if let Some(path) = self.dfs_cycle_check(neighbor, visited, rec_stack) {
271                    return Some(path);
272                }
273            } else if rec_stack.contains(&neighbor) {
274                // Cycle detected, construct path
275                let current_task = self.graph.node_weight(node).map(|t| &t.id)?;
276                let next_task = self.graph.node_weight(neighbor).map(|t| &t.id)?;
277                return Some(format!("{} -> {}", current_task, next_task));
278            }
279        }
280
281        rec_stack.remove(&node);
282        None
283    }
284
285    /// Check if all nodes are reachable from root nodes.
286    fn check_reachability(&self) -> Result<()> {
287        // Find root nodes (nodes with no incoming edges)
288        let root_nodes: Vec<NodeIndex> = self
289            .graph
290            .node_indices()
291            .filter(|&idx| self.graph.edges_directed(idx, Direction::Incoming).count() == 0)
292            .collect();
293
294        if root_nodes.is_empty() {
295            // If no root nodes, check if the graph has cycles (all nodes have incoming edges)
296            return Ok(());
297        }
298
299        // BFS from all root nodes to find reachable nodes
300        let mut reachable = HashSet::new();
301        let mut queue = VecDeque::from(root_nodes);
302
303        while let Some(node) = queue.pop_front() {
304            if reachable.insert(node) {
305                for neighbor in self.graph.neighbors(node) {
306                    if !reachable.contains(&neighbor) {
307                        queue.push_back(neighbor);
308                    }
309                }
310            }
311        }
312
313        // Check if all nodes are reachable
314        for node_idx in self.graph.node_indices() {
315            if !reachable.contains(&node_idx)
316                && let Some(task) = self.graph.node_weight(node_idx)
317            {
318                return Err(DagError::UnreachableNode(task.id.clone()).into());
319            }
320        }
321
322        Ok(())
323    }
324
325    /// Get all tasks in the DAG.
326    pub fn tasks(&self) -> Vec<&TaskNode> {
327        self.graph
328            .node_indices()
329            .filter_map(|idx| self.graph.node_weight(idx))
330            .collect()
331    }
332
333    /// Get the number of tasks in the DAG.
334    pub fn task_count(&self) -> usize {
335        self.graph.node_count()
336    }
337
338    /// Get the number of dependencies in the DAG.
339    pub fn dependency_count(&self) -> usize {
340        self.graph.edge_count()
341    }
342
343    /// Get root tasks (tasks with no dependencies).
344    pub fn root_tasks(&self) -> Vec<&TaskNode> {
345        self.graph
346            .node_indices()
347            .filter(|&idx| self.graph.edges_directed(idx, Direction::Incoming).count() == 0)
348            .filter_map(|idx| self.graph.node_weight(idx))
349            .collect()
350    }
351
352    /// Get leaf tasks (tasks with no dependents).
353    pub fn leaf_tasks(&self) -> Vec<&TaskNode> {
354        self.graph
355            .node_indices()
356            .filter(|&idx| self.graph.edges_directed(idx, Direction::Outgoing).count() == 0)
357            .filter_map(|idx| self.graph.node_weight(idx))
358            .collect()
359    }
360
361    /// Get all edges in the DAG as (from_task_id, to_task_id, edge_data) tuples.
362    ///
363    /// This method is useful for visualization and serialization purposes.
364    /// Returns edges in the order they are stored in the graph.
365    pub fn edges(&self) -> Vec<(&str, &str, &TaskEdge)> {
366        self.graph
367            .edge_indices()
368            .filter_map(|edge_idx| {
369                let (from_idx, to_idx) = self.graph.edge_endpoints(edge_idx)?;
370                let from_node = self.graph.node_weight(from_idx)?;
371                let to_node = self.graph.node_weight(to_idx)?;
372                let edge = self.graph.edge_weight(edge_idx)?;
373                Some((from_node.id.as_str(), to_node.id.as_str(), edge))
374            })
375            .collect()
376    }
377
378    /// Get all edges with their edge types as (from_task_id, to_task_id, edge_type) tuples.
379    ///
380    /// A simplified version of `edges()` that only returns edge types.
381    pub fn edge_pairs(&self) -> Vec<(String, String)> {
382        self.graph
383            .edge_indices()
384            .filter_map(|edge_idx| {
385                let (from_idx, to_idx) = self.graph.edge_endpoints(edge_idx)?;
386                let from_node = self.graph.node_weight(from_idx)?;
387                let to_node = self.graph.node_weight(to_idx)?;
388                Some((from_node.id.clone(), to_node.id.clone()))
389            })
390            .collect()
391    }
392
393    /// Get task dependencies along with their edge data.
394    ///
395    /// Returns a vector of (dependency_task_id, edge_data) tuples for the given task.
396    /// Dependencies are tasks that must complete before the specified task can start.
397    pub fn get_dependencies_with_edges(&self, task_id: &str) -> Vec<(String, &TaskEdge)> {
398        if let Some(&idx) = self.task_map.get(task_id) {
399            self.graph
400                .edges_directed(idx, Direction::Incoming)
401                .filter_map(|edge| {
402                    let source_node = self.graph.node_weight(edge.source())?;
403                    Some((source_node.id.clone(), edge.weight()))
404                })
405                .collect()
406        } else {
407            Vec::new()
408        }
409    }
410
411    /// Get task dependents along with their edge data.
412    ///
413    /// Returns a vector of (dependent_task_id, edge_data) tuples for the given task.
414    /// Dependents are tasks that wait for the specified task to complete.
415    pub fn get_dependents_with_edges(&self, task_id: &str) -> Vec<(String, &TaskEdge)> {
416        if let Some(&idx) = self.task_map.get(task_id) {
417            self.graph
418                .edges_directed(idx, Direction::Outgoing)
419                .filter_map(|edge| {
420                    let target_node = self.graph.node_weight(edge.target())?;
421                    Some((target_node.id.clone(), edge.weight()))
422                })
423                .collect()
424        } else {
425            Vec::new()
426        }
427    }
428
429    /// Get the edge data between two specific tasks, if it exists.
430    ///
431    /// Returns `None` if either task does not exist or no edge connects them.
432    pub fn get_edge_between(&self, from_task_id: &str, to_task_id: &str) -> Option<&TaskEdge> {
433        let from_idx = self.task_map.get(from_task_id)?;
434        let to_idx = self.task_map.get(to_task_id)?;
435        self.graph
436            .find_edge(*from_idx, *to_idx)
437            .and_then(|edge_idx| self.graph.edge_weight(edge_idx))
438    }
439
440    /// Check if a dependency exists between two tasks.
441    ///
442    /// Returns `true` if `from_task_id` has a direct edge to `to_task_id`.
443    pub fn has_dependency(&self, from_task_id: &str, to_task_id: &str) -> bool {
444        self.get_edge_between(from_task_id, to_task_id).is_some()
445    }
446
447    /// Check if a task has any dependencies (incoming edges).
448    pub fn has_dependencies(&self, task_id: &str) -> bool {
449        if let Some(&idx) = self.task_map.get(task_id) {
450            self.graph.edges_directed(idx, Direction::Incoming).count() > 0
451        } else {
452            false
453        }
454    }
455
456    /// Check if a task has any dependents (outgoing edges).
457    pub fn has_dependents(&self, task_id: &str) -> bool {
458        if let Some(&idx) = self.task_map.get(task_id) {
459            self.graph.edges_directed(idx, Direction::Outgoing).count() > 0
460        } else {
461            false
462        }
463    }
464
465    /// Get the in-degree of a task (number of dependencies).
466    pub fn in_degree(&self, task_id: &str) -> usize {
467        if let Some(&idx) = self.task_map.get(task_id) {
468            self.graph.edges_directed(idx, Direction::Incoming).count()
469        } else {
470            0
471        }
472    }
473
474    /// Get the out-degree of a task (number of dependents).
475    pub fn out_degree(&self, task_id: &str) -> usize {
476        if let Some(&idx) = self.task_map.get(task_id) {
477            self.graph.edges_directed(idx, Direction::Outgoing).count()
478        } else {
479            0
480        }
481    }
482
483    /// Get all task IDs in the DAG.
484    pub fn task_ids(&self) -> Vec<String> {
485        self.task_map.keys().cloned().collect()
486    }
487
488    /// Check if a task exists in the DAG.
489    pub fn contains_task(&self, task_id: &str) -> bool {
490        self.task_map.contains_key(task_id)
491    }
492
493    /// Remove a task from the DAG along with all its edges.
494    ///
495    /// Returns the removed task, or `None` if the task did not exist.
496    ///
497    /// `petgraph::graph::DiGraph::remove_node` swap-removes: the node currently
498    /// at the last index is moved into the freed slot and re-indexed. We must
499    /// repoint `task_map` for that swapped node, otherwise later lookups for it
500    /// would resolve to a stale/out-of-bounds `NodeIndex`.
501    pub fn remove_task(&mut self, task_id: &str) -> Option<TaskNode> {
502        let node_idx = self.task_map.remove(task_id)?;
503
504        // Index of the node that will be swapped into `node_idx` (if any). This
505        // must be read *before* the removal, while the count still includes the
506        // node being removed.
507        let last_idx = NodeIndex::new(self.graph.node_count() - 1);
508
509        let removed = self.graph.remove_node(node_idx)?;
510
511        // If the removed node was not itself the last node, the previously-last
512        // node now lives at `node_idx`; update its task_map entry to match.
513        if last_idx != node_idx
514            && let Some(moved_task) = self.graph.node_weight(node_idx)
515        {
516            self.task_map.insert(moved_task.id.clone(), node_idx);
517        }
518
519        Some(removed)
520    }
521
522    /// Get edges filtered by edge type.
523    pub fn edges_by_type(&self, edge_type: EdgeType) -> Vec<(&str, &str, &TaskEdge)> {
524        self.graph
525            .edge_indices()
526            .filter_map(|edge_idx| {
527                let edge = self.graph.edge_weight(edge_idx)?;
528                if edge.edge_type != edge_type {
529                    return None;
530                }
531                let (from_idx, to_idx) = self.graph.edge_endpoints(edge_idx)?;
532                let from_node = self.graph.node_weight(from_idx)?;
533                let to_node = self.graph.node_weight(to_idx)?;
534                Some((from_node.id.as_str(), to_node.id.as_str(), edge))
535            })
536            .collect()
537    }
538
539    /// Get a subgraph containing only the specified tasks and edges between them.
540    ///
541    /// Tasks not present in the original DAG are silently ignored.
542    pub fn subgraph(&self, task_ids: &[&str]) -> WorkflowDag {
543        let mut sub = WorkflowDag::new();
544        let id_set: HashSet<&str> = task_ids.iter().copied().collect();
545
546        // Add matching nodes
547        for task_id in task_ids {
548            if let Some(task) = self.get_task(task_id) {
549                // Ignore errors from duplicate insertions if task_ids has duplicates
550                let _ = sub.add_task(task.clone());
551            }
552        }
553
554        // Add edges that connect nodes within the subgraph
555        for (from_id, to_id, edge) in self.edges() {
556            if id_set.contains(from_id) && id_set.contains(to_id) {
557                let _ = sub.add_dependency(from_id, to_id, edge.clone());
558            }
559        }
560
561        sub
562    }
563
564    /// Compute the transitive closure of dependencies for a task.
565    ///
566    /// Returns all tasks that must complete (directly or transitively) before
567    /// the given task can execute.
568    pub fn transitive_dependencies(&self, task_id: &str) -> Vec<String> {
569        let mut visited = HashSet::new();
570        let mut queue = VecDeque::new();
571
572        // Seed with direct dependencies
573        for dep in self.get_dependencies(task_id) {
574            if visited.insert(dep.clone()) {
575                queue.push_back(dep);
576            }
577        }
578
579        while let Some(current) = queue.pop_front() {
580            for dep in self.get_dependencies(&current) {
581                if visited.insert(dep.clone()) {
582                    queue.push_back(dep);
583                }
584            }
585        }
586
587        visited.into_iter().collect()
588    }
589
590    /// Compute the transitive closure of dependents for a task.
591    ///
592    /// Returns all tasks that (directly or transitively) depend on the given task.
593    pub fn transitive_dependents(&self, task_id: &str) -> Vec<String> {
594        let mut visited = HashSet::new();
595        let mut queue = VecDeque::new();
596
597        // Seed with direct dependents
598        for dep in self.get_dependents(task_id) {
599            if visited.insert(dep.clone()) {
600                queue.push_back(dep);
601            }
602        }
603
604        while let Some(current) = queue.pop_front() {
605            for dep in self.get_dependents(&current) {
606                if visited.insert(dep.clone()) {
607                    queue.push_back(dep);
608                }
609            }
610        }
611
612        visited.into_iter().collect()
613    }
614
615    /// Get summary statistics about the DAG structure.
616    pub fn summary(&self) -> DagSummary {
617        let node_count = self.graph.node_count();
618        let edge_count = self.graph.edge_count();
619        let root_count = self.root_tasks().len();
620        let leaf_count = self.leaf_tasks().len();
621
622        let max_in_degree = self
623            .graph
624            .node_indices()
625            .map(|idx| self.graph.edges_directed(idx, Direction::Incoming).count())
626            .max()
627            .unwrap_or(0);
628
629        let max_out_degree = self
630            .graph
631            .node_indices()
632            .map(|idx| self.graph.edges_directed(idx, Direction::Outgoing).count())
633            .max()
634            .unwrap_or(0);
635
636        let data_edges = self.edges_by_type(EdgeType::Data).len();
637        let control_edges = self.edges_by_type(EdgeType::Control).len();
638        let conditional_edges = self.edges_by_type(EdgeType::Conditional).len();
639
640        DagSummary {
641            node_count,
642            edge_count,
643            root_count,
644            leaf_count,
645            max_in_degree,
646            max_out_degree,
647            data_edge_count: data_edges,
648            control_edge_count: control_edges,
649            conditional_edge_count: conditional_edges,
650        }
651    }
652}
653
654/// Summary statistics for a DAG.
655#[derive(Debug, Clone, Serialize, Deserialize)]
656pub struct DagSummary {
657    /// Number of task nodes.
658    pub node_count: usize,
659    /// Number of dependency edges.
660    pub edge_count: usize,
661    /// Number of root tasks (no dependencies).
662    pub root_count: usize,
663    /// Number of leaf tasks (no dependents).
664    pub leaf_count: usize,
665    /// Maximum number of dependencies any single task has.
666    pub max_in_degree: usize,
667    /// Maximum number of dependents any single task has.
668    pub max_out_degree: usize,
669    /// Number of data dependency edges.
670    pub data_edge_count: usize,
671    /// Number of control dependency edges.
672    pub control_edge_count: usize,
673    /// Number of conditional dependency edges.
674    pub conditional_edge_count: usize,
675}
676
677impl Default for WorkflowDag {
678    fn default() -> Self {
679        Self::new()
680    }
681}
682
683#[cfg(test)]
684mod tests {
685    use super::*;
686
687    fn create_test_task(id: &str, name: &str) -> TaskNode {
688        TaskNode {
689            id: id.to_string(),
690            name: name.to_string(),
691            description: None,
692            config: serde_json::json!({}),
693            retry: RetryPolicy::default(),
694            timeout_secs: Some(60),
695            resources: ResourceRequirements::default(),
696            metadata: HashMap::new(),
697        }
698    }
699
700    #[test]
701    fn test_add_task() {
702        let mut dag = WorkflowDag::new();
703        let task = create_test_task("task1", "Task 1");
704        let result = dag.add_task(task);
705        assert!(result.is_ok());
706        assert_eq!(dag.task_count(), 1);
707    }
708
709    #[test]
710    fn test_duplicate_task() {
711        let mut dag = WorkflowDag::new();
712        let task1 = create_test_task("task1", "Task 1");
713        let task2 = create_test_task("task1", "Task 1 Duplicate");
714
715        dag.add_task(task1).ok();
716        let result = dag.add_task(task2);
717        assert!(result.is_err());
718    }
719
720    #[test]
721    fn test_add_dependency() {
722        let mut dag = WorkflowDag::new();
723        dag.add_task(create_test_task("task1", "Task 1")).ok();
724        dag.add_task(create_test_task("task2", "Task 2")).ok();
725
726        let result = dag.add_dependency("task1", "task2", TaskEdge::default());
727        assert!(result.is_ok());
728        assert_eq!(dag.dependency_count(), 1);
729    }
730
731    #[test]
732    fn test_cycle_detection() {
733        let mut dag = WorkflowDag::new();
734        dag.add_task(create_test_task("task1", "Task 1")).ok();
735        dag.add_task(create_test_task("task2", "Task 2")).ok();
736        dag.add_task(create_test_task("task3", "Task 3")).ok();
737
738        // Create a cycle: task1 -> task2 -> task3 -> task1
739        dag.add_dependency("task1", "task2", TaskEdge::default())
740            .ok();
741        dag.add_dependency("task2", "task3", TaskEdge::default())
742            .ok();
743        dag.add_dependency("task3", "task1", TaskEdge::default())
744            .ok();
745
746        let result = dag.validate();
747        assert!(result.is_err());
748    }
749
750    #[test]
751    fn test_valid_dag() {
752        let mut dag = WorkflowDag::new();
753        dag.add_task(create_test_task("task1", "Task 1")).ok();
754        dag.add_task(create_test_task("task2", "Task 2")).ok();
755        dag.add_task(create_test_task("task3", "Task 3")).ok();
756
757        // Create a valid DAG: task1 -> task2, task1 -> task3
758        dag.add_dependency("task1", "task2", TaskEdge::default())
759            .ok();
760        dag.add_dependency("task1", "task3", TaskEdge::default())
761            .ok();
762
763        let result = dag.validate();
764        assert!(result.is_ok());
765    }
766
767    #[test]
768    fn test_root_and_leaf_tasks() {
769        let mut dag = WorkflowDag::new();
770        dag.add_task(create_test_task("task1", "Task 1")).ok();
771        dag.add_task(create_test_task("task2", "Task 2")).ok();
772        dag.add_task(create_test_task("task3", "Task 3")).ok();
773
774        dag.add_dependency("task1", "task2", TaskEdge::default())
775            .ok();
776        dag.add_dependency("task2", "task3", TaskEdge::default())
777            .ok();
778
779        let roots = dag.root_tasks();
780        assert_eq!(roots.len(), 1);
781        assert_eq!(roots[0].id, "task1");
782
783        let leaves = dag.leaf_tasks();
784        assert_eq!(leaves.len(), 1);
785        assert_eq!(leaves[0].id, "task3");
786    }
787
788    #[test]
789    fn test_edges() {
790        let mut dag = WorkflowDag::new();
791        dag.add_task(create_test_task("t1", "Task 1")).ok();
792        dag.add_task(create_test_task("t2", "Task 2")).ok();
793        dag.add_task(create_test_task("t3", "Task 3")).ok();
794
795        dag.add_dependency("t1", "t2", TaskEdge::default()).ok();
796        dag.add_dependency(
797            "t2",
798            "t3",
799            TaskEdge {
800                edge_type: EdgeType::Data,
801                condition: None,
802            },
803        )
804        .ok();
805
806        let edges = dag.edges();
807        assert_eq!(edges.len(), 2);
808
809        // Check first edge
810        let (from, to, edge) = &edges[0];
811        assert_eq!(*from, "t1");
812        assert_eq!(*to, "t2");
813        assert_eq!(edge.edge_type, EdgeType::Control);
814
815        // Check second edge
816        let (from, to, edge) = &edges[1];
817        assert_eq!(*from, "t2");
818        assert_eq!(*to, "t3");
819        assert_eq!(edge.edge_type, EdgeType::Data);
820    }
821
822    #[test]
823    fn test_get_dependencies_with_edges() {
824        let mut dag = WorkflowDag::new();
825        dag.add_task(create_test_task("t1", "Task 1")).ok();
826        dag.add_task(create_test_task("t2", "Task 2")).ok();
827        dag.add_task(create_test_task("t3", "Task 3")).ok();
828
829        dag.add_dependency(
830            "t1",
831            "t3",
832            TaskEdge {
833                edge_type: EdgeType::Data,
834                condition: None,
835            },
836        )
837        .ok();
838        dag.add_dependency("t2", "t3", TaskEdge::default()).ok();
839
840        let deps = dag.get_dependencies_with_edges("t3");
841        assert_eq!(deps.len(), 2);
842
843        // Both t1 and t2 should be dependencies of t3
844        let dep_ids: Vec<&str> = deps.iter().map(|(id, _)| id.as_str()).collect();
845        assert!(dep_ids.contains(&"t1"));
846        assert!(dep_ids.contains(&"t2"));
847
848        // No dependencies for root task
849        let root_deps = dag.get_dependencies_with_edges("t1");
850        assert!(root_deps.is_empty());
851
852        // Non-existent task returns empty
853        let missing_deps = dag.get_dependencies_with_edges("nonexistent");
854        assert!(missing_deps.is_empty());
855    }
856
857    #[test]
858    fn test_get_dependents_with_edges() {
859        let mut dag = WorkflowDag::new();
860        dag.add_task(create_test_task("t1", "Task 1")).ok();
861        dag.add_task(create_test_task("t2", "Task 2")).ok();
862        dag.add_task(create_test_task("t3", "Task 3")).ok();
863
864        dag.add_dependency("t1", "t2", TaskEdge::default()).ok();
865        dag.add_dependency("t1", "t3", TaskEdge::default()).ok();
866
867        let dependents = dag.get_dependents_with_edges("t1");
868        assert_eq!(dependents.len(), 2);
869
870        let dep_ids: Vec<&str> = dependents.iter().map(|(id, _)| id.as_str()).collect();
871        assert!(dep_ids.contains(&"t2"));
872        assert!(dep_ids.contains(&"t3"));
873    }
874
875    #[test]
876    fn test_get_edge_between() {
877        let mut dag = WorkflowDag::new();
878        dag.add_task(create_test_task("t1", "Task 1")).ok();
879        dag.add_task(create_test_task("t2", "Task 2")).ok();
880        dag.add_task(create_test_task("t3", "Task 3")).ok();
881
882        dag.add_dependency(
883            "t1",
884            "t2",
885            TaskEdge {
886                edge_type: EdgeType::Data,
887                condition: Some("output.ready".to_string()),
888            },
889        )
890        .ok();
891
892        let edge = dag.get_edge_between("t1", "t2");
893        assert!(edge.is_some());
894        let edge = edge.expect("Edge should exist");
895        assert_eq!(edge.edge_type, EdgeType::Data);
896        assert_eq!(edge.condition.as_deref(), Some("output.ready"));
897
898        // Reverse direction should not exist
899        assert!(dag.get_edge_between("t2", "t1").is_none());
900        // Non-connected nodes
901        assert!(dag.get_edge_between("t1", "t3").is_none());
902    }
903
904    #[test]
905    fn test_has_dependency() {
906        let mut dag = WorkflowDag::new();
907        dag.add_task(create_test_task("t1", "Task 1")).ok();
908        dag.add_task(create_test_task("t2", "Task 2")).ok();
909
910        dag.add_dependency("t1", "t2", TaskEdge::default()).ok();
911
912        assert!(dag.has_dependency("t1", "t2"));
913        assert!(!dag.has_dependency("t2", "t1"));
914        assert!(!dag.has_dependency("t1", "nonexistent"));
915    }
916
917    #[test]
918    fn test_has_dependencies_and_dependents() {
919        let mut dag = WorkflowDag::new();
920        dag.add_task(create_test_task("t1", "Task 1")).ok();
921        dag.add_task(create_test_task("t2", "Task 2")).ok();
922        dag.add_task(create_test_task("t3", "Task 3")).ok();
923
924        dag.add_dependency("t1", "t2", TaskEdge::default()).ok();
925        dag.add_dependency("t2", "t3", TaskEdge::default()).ok();
926
927        // t1: root, has dependents but no dependencies
928        assert!(!dag.has_dependencies("t1"));
929        assert!(dag.has_dependents("t1"));
930
931        // t2: middle, has both
932        assert!(dag.has_dependencies("t2"));
933        assert!(dag.has_dependents("t2"));
934
935        // t3: leaf, has dependencies but no dependents
936        assert!(dag.has_dependencies("t3"));
937        assert!(!dag.has_dependents("t3"));
938    }
939
940    #[test]
941    fn test_in_out_degree() {
942        let mut dag = WorkflowDag::new();
943        dag.add_task(create_test_task("t1", "Task 1")).ok();
944        dag.add_task(create_test_task("t2", "Task 2")).ok();
945        dag.add_task(create_test_task("t3", "Task 3")).ok();
946        dag.add_task(create_test_task("t4", "Task 4")).ok();
947
948        // t1 -> t3, t2 -> t3, t3 -> t4
949        dag.add_dependency("t1", "t3", TaskEdge::default()).ok();
950        dag.add_dependency("t2", "t3", TaskEdge::default()).ok();
951        dag.add_dependency("t3", "t4", TaskEdge::default()).ok();
952
953        assert_eq!(dag.in_degree("t1"), 0);
954        assert_eq!(dag.out_degree("t1"), 1);
955        assert_eq!(dag.in_degree("t3"), 2);
956        assert_eq!(dag.out_degree("t3"), 1);
957        assert_eq!(dag.in_degree("t4"), 1);
958        assert_eq!(dag.out_degree("t4"), 0);
959        // Non-existent
960        assert_eq!(dag.in_degree("nonexistent"), 0);
961    }
962
963    #[test]
964    fn test_task_ids_and_contains() {
965        let mut dag = WorkflowDag::new();
966        dag.add_task(create_test_task("t1", "Task 1")).ok();
967        dag.add_task(create_test_task("t2", "Task 2")).ok();
968
969        let ids = dag.task_ids();
970        assert_eq!(ids.len(), 2);
971        assert!(dag.contains_task("t1"));
972        assert!(dag.contains_task("t2"));
973        assert!(!dag.contains_task("t3"));
974    }
975
976    #[test]
977    fn test_remove_task() {
978        let mut dag = WorkflowDag::new();
979        dag.add_task(create_test_task("t1", "Task 1")).ok();
980        dag.add_task(create_test_task("t2", "Task 2")).ok();
981        dag.add_dependency("t1", "t2", TaskEdge::default()).ok();
982
983        assert_eq!(dag.task_count(), 2);
984        assert_eq!(dag.dependency_count(), 1);
985
986        let removed = dag.remove_task("t1");
987        assert!(removed.is_some());
988        assert_eq!(removed.as_ref().map(|t| t.id.as_str()), Some("t1"));
989        assert!(!dag.contains_task("t1"));
990
991        // Removing non-existent should return None
992        assert!(dag.remove_task("nonexistent").is_none());
993    }
994
995    #[test]
996    fn test_remove_task_keeps_task_map_consistent_after_swap() {
997        // Regression: petgraph swap-removes, moving the last-inserted node into
998        // the freed slot. task_map must be repointed for that swapped node.
999        let mut dag = WorkflowDag::new();
1000        dag.add_task(create_test_task("t1", "Task 1")).ok();
1001        dag.add_task(create_test_task("t2", "Task 2")).ok();
1002        dag.add_task(create_test_task("t3", "Task 3")).ok();
1003
1004        // Edges touching the task that will be swapped (t3, the last node).
1005        dag.add_dependency("t2", "t3", TaskEdge::default()).ok();
1006        dag.add_dependency("t1", "t3", TaskEdge::default()).ok();
1007
1008        // Remove a NON-last task; petgraph moves t3 into t1's old slot.
1009        let removed = dag.remove_task("t1");
1010        assert_eq!(removed.as_ref().map(|t| t.id.as_str()), Some("t1"));
1011        assert!(!dag.contains_task("t1"));
1012
1013        // t3 (the swapped node) must still resolve to the correct TaskNode and
1014        // report accurate degree/dependency information.
1015        let t3 = dag.get_task("t3").expect("t3 must still be reachable");
1016        assert_eq!(t3.id, "t3");
1017        assert_eq!(t3.name, "Task 3");
1018        assert_eq!(dag.in_degree("t3"), 1); // only the t2 -> t3 edge remains
1019        assert!(dag.has_dependency("t2", "t3"));
1020        assert_eq!(dag.get_dependencies("t3"), vec!["t2".to_string()]);
1021
1022        // t2 must also remain correct.
1023        let t2 = dag.get_task("t2").expect("t2 must still be reachable");
1024        assert_eq!(t2.name, "Task 2");
1025
1026        // A follow-up removal must keep indices consistent as well.
1027        let removed_t2 = dag.remove_task("t2");
1028        assert_eq!(removed_t2.as_ref().map(|t| t.id.as_str()), Some("t2"));
1029        let t3_again = dag.get_task("t3").expect("t3 must survive second removal");
1030        assert_eq!(t3_again.id, "t3");
1031        assert_eq!(dag.in_degree("t3"), 0);
1032        assert_eq!(dag.task_count(), 1);
1033    }
1034
1035    #[test]
1036    fn test_edges_by_type() {
1037        let mut dag = WorkflowDag::new();
1038        dag.add_task(create_test_task("t1", "Task 1")).ok();
1039        dag.add_task(create_test_task("t2", "Task 2")).ok();
1040        dag.add_task(create_test_task("t3", "Task 3")).ok();
1041
1042        dag.add_dependency(
1043            "t1",
1044            "t2",
1045            TaskEdge {
1046                edge_type: EdgeType::Data,
1047                condition: None,
1048            },
1049        )
1050        .ok();
1051        dag.add_dependency("t1", "t3", TaskEdge::default()).ok();
1052
1053        let data_edges = dag.edges_by_type(EdgeType::Data);
1054        assert_eq!(data_edges.len(), 1);
1055        assert_eq!(data_edges[0].0, "t1");
1056        assert_eq!(data_edges[0].1, "t2");
1057
1058        let control_edges = dag.edges_by_type(EdgeType::Control);
1059        assert_eq!(control_edges.len(), 1);
1060        assert_eq!(control_edges[0].0, "t1");
1061        assert_eq!(control_edges[0].1, "t3");
1062    }
1063
1064    #[test]
1065    fn test_subgraph() {
1066        let mut dag = WorkflowDag::new();
1067        dag.add_task(create_test_task("t1", "Task 1")).ok();
1068        dag.add_task(create_test_task("t2", "Task 2")).ok();
1069        dag.add_task(create_test_task("t3", "Task 3")).ok();
1070        dag.add_task(create_test_task("t4", "Task 4")).ok();
1071
1072        dag.add_dependency("t1", "t2", TaskEdge::default()).ok();
1073        dag.add_dependency("t2", "t3", TaskEdge::default()).ok();
1074        dag.add_dependency("t3", "t4", TaskEdge::default()).ok();
1075
1076        // Extract subgraph with only t2 and t3
1077        let sub = dag.subgraph(&["t2", "t3"]);
1078        assert_eq!(sub.task_count(), 2);
1079        assert_eq!(sub.dependency_count(), 1);
1080        assert!(sub.contains_task("t2"));
1081        assert!(sub.contains_task("t3"));
1082        assert!(!sub.contains_task("t1"));
1083        assert!(!sub.contains_task("t4"));
1084    }
1085
1086    #[test]
1087    fn test_transitive_dependencies() {
1088        let mut dag = WorkflowDag::new();
1089        dag.add_task(create_test_task("t1", "Task 1")).ok();
1090        dag.add_task(create_test_task("t2", "Task 2")).ok();
1091        dag.add_task(create_test_task("t3", "Task 3")).ok();
1092        dag.add_task(create_test_task("t4", "Task 4")).ok();
1093
1094        dag.add_dependency("t1", "t2", TaskEdge::default()).ok();
1095        dag.add_dependency("t2", "t3", TaskEdge::default()).ok();
1096        dag.add_dependency("t3", "t4", TaskEdge::default()).ok();
1097
1098        let trans_deps = dag.transitive_dependencies("t4");
1099        assert_eq!(trans_deps.len(), 3);
1100        assert!(trans_deps.contains(&"t1".to_string()));
1101        assert!(trans_deps.contains(&"t2".to_string()));
1102        assert!(trans_deps.contains(&"t3".to_string()));
1103
1104        // Root has no transitive dependencies
1105        let root_deps = dag.transitive_dependencies("t1");
1106        assert!(root_deps.is_empty());
1107    }
1108
1109    #[test]
1110    fn test_transitive_dependents() {
1111        let mut dag = WorkflowDag::new();
1112        dag.add_task(create_test_task("t1", "Task 1")).ok();
1113        dag.add_task(create_test_task("t2", "Task 2")).ok();
1114        dag.add_task(create_test_task("t3", "Task 3")).ok();
1115        dag.add_task(create_test_task("t4", "Task 4")).ok();
1116
1117        dag.add_dependency("t1", "t2", TaskEdge::default()).ok();
1118        dag.add_dependency("t2", "t3", TaskEdge::default()).ok();
1119        dag.add_dependency("t3", "t4", TaskEdge::default()).ok();
1120
1121        let trans_dependents = dag.transitive_dependents("t1");
1122        assert_eq!(trans_dependents.len(), 3);
1123        assert!(trans_dependents.contains(&"t2".to_string()));
1124        assert!(trans_dependents.contains(&"t3".to_string()));
1125        assert!(trans_dependents.contains(&"t4".to_string()));
1126
1127        // Leaf has no transitive dependents
1128        let leaf_deps = dag.transitive_dependents("t4");
1129        assert!(leaf_deps.is_empty());
1130    }
1131
1132    #[test]
1133    fn test_summary() {
1134        let mut dag = WorkflowDag::new();
1135        dag.add_task(create_test_task("t1", "Task 1")).ok();
1136        dag.add_task(create_test_task("t2", "Task 2")).ok();
1137        dag.add_task(create_test_task("t3", "Task 3")).ok();
1138        dag.add_task(create_test_task("t4", "Task 4")).ok();
1139
1140        dag.add_dependency(
1141            "t1",
1142            "t2",
1143            TaskEdge {
1144                edge_type: EdgeType::Data,
1145                condition: None,
1146            },
1147        )
1148        .ok();
1149        dag.add_dependency("t1", "t3", TaskEdge::default()).ok();
1150        dag.add_dependency("t2", "t4", TaskEdge::default()).ok();
1151        dag.add_dependency("t3", "t4", TaskEdge::default()).ok();
1152
1153        let summary = dag.summary();
1154        assert_eq!(summary.node_count, 4);
1155        assert_eq!(summary.edge_count, 4);
1156        assert_eq!(summary.root_count, 1);
1157        assert_eq!(summary.leaf_count, 1);
1158        assert_eq!(summary.max_in_degree, 2); // t4 has 2 incoming
1159        assert_eq!(summary.max_out_degree, 2); // t1 has 2 outgoing
1160        assert_eq!(summary.data_edge_count, 1);
1161        assert_eq!(summary.control_edge_count, 3);
1162        assert_eq!(summary.conditional_edge_count, 0);
1163    }
1164
1165    #[test]
1166    fn test_edge_pairs() {
1167        let mut dag = WorkflowDag::new();
1168        dag.add_task(create_test_task("t1", "Task 1")).ok();
1169        dag.add_task(create_test_task("t2", "Task 2")).ok();
1170        dag.add_dependency("t1", "t2", TaskEdge::default()).ok();
1171
1172        let pairs = dag.edge_pairs();
1173        assert_eq!(pairs.len(), 1);
1174        assert_eq!(pairs[0], ("t1".to_string(), "t2".to_string()));
1175    }
1176
1177    #[test]
1178    fn test_get_dependencies_and_dependents() {
1179        let mut dag = WorkflowDag::new();
1180        dag.add_task(create_test_task("t1", "Task 1")).ok();
1181        dag.add_task(create_test_task("t2", "Task 2")).ok();
1182        dag.add_task(create_test_task("t3", "Task 3")).ok();
1183
1184        dag.add_dependency("t1", "t3", TaskEdge::default()).ok();
1185        dag.add_dependency("t2", "t3", TaskEdge::default()).ok();
1186
1187        let deps = dag.get_dependencies("t3");
1188        assert_eq!(deps.len(), 2);
1189        assert!(deps.contains(&"t1".to_string()));
1190        assert!(deps.contains(&"t2".to_string()));
1191
1192        let dependents = dag.get_dependents("t1");
1193        assert_eq!(dependents.len(), 1);
1194        assert!(dependents.contains(&"t3".to_string()));
1195    }
1196}