1use 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#[derive(Debug, Clone, Serialize, Deserialize)]
13pub struct TaskNode {
14 pub id: String,
16 pub name: String,
18 pub description: Option<String>,
20 pub config: serde_json::Value,
22 pub retry: RetryPolicy,
24 pub timeout_secs: Option<u64>,
26 pub resources: ResourceRequirements,
28 pub metadata: HashMap<String, String>,
30}
31
32#[derive(Debug, Clone, Serialize, Deserialize)]
34pub struct RetryPolicy {
35 pub max_attempts: u32,
37 pub delay_ms: u64,
39 pub backoff_multiplier: f64,
41 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#[derive(Debug, Clone, Serialize, Deserialize)]
58pub struct ResourceRequirements {
59 pub cpu_cores: f64,
61 pub memory_mb: u64,
63 pub gpu: bool,
65 pub disk_mb: u64,
67 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#[derive(Debug, Clone, Serialize, Deserialize)]
99pub struct TaskEdge {
100 pub edge_type: EdgeType,
102 pub condition: Option<String>,
104}
105
106#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
108pub enum EdgeType {
109 Data,
111 Control,
113 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
128pub struct WorkflowDag {
130 pub(crate) graph: DiGraph<TaskNode, TaskEdge>,
132 pub(crate) task_map: HashMap<String, NodeIndex>,
134}
135
136impl WorkflowDag {
137 pub fn new() -> Self {
139 Self {
140 graph: DiGraph::new(),
141 task_map: HashMap::new(),
142 }
143 }
144
145 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 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 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 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 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 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 pub fn validate(&self) -> Result<()> {
227 if self.graph.node_count() == 0 {
229 return Err(DagError::EmptyDag.into());
230 }
231
232 self.check_cycles()?;
234
235 self.check_reachability()?;
237
238 Ok(())
239 }
240
241 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 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 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 fn check_reachability(&self) -> Result<()> {
287 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 return Ok(());
297 }
298
299 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 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 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 pub fn task_count(&self) -> usize {
335 self.graph.node_count()
336 }
337
338 pub fn dependency_count(&self) -> usize {
340 self.graph.edge_count()
341 }
342
343 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 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 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 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 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 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 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 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 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 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 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 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 pub fn task_ids(&self) -> Vec<String> {
485 self.task_map.keys().cloned().collect()
486 }
487
488 pub fn contains_task(&self, task_id: &str) -> bool {
490 self.task_map.contains_key(task_id)
491 }
492
493 pub fn remove_task(&mut self, task_id: &str) -> Option<TaskNode> {
502 let node_idx = self.task_map.remove(task_id)?;
503
504 let last_idx = NodeIndex::new(self.graph.node_count() - 1);
508
509 let removed = self.graph.remove_node(node_idx)?;
510
511 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 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 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 for task_id in task_ids {
548 if let Some(task) = self.get_task(task_id) {
549 let _ = sub.add_task(task.clone());
551 }
552 }
553
554 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 pub fn transitive_dependencies(&self, task_id: &str) -> Vec<String> {
569 let mut visited = HashSet::new();
570 let mut queue = VecDeque::new();
571
572 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(¤t) {
581 if visited.insert(dep.clone()) {
582 queue.push_back(dep);
583 }
584 }
585 }
586
587 visited.into_iter().collect()
588 }
589
590 pub fn transitive_dependents(&self, task_id: &str) -> Vec<String> {
594 let mut visited = HashSet::new();
595 let mut queue = VecDeque::new();
596
597 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(¤t) {
606 if visited.insert(dep.clone()) {
607 queue.push_back(dep);
608 }
609 }
610 }
611
612 visited.into_iter().collect()
613 }
614
615 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#[derive(Debug, Clone, Serialize, Deserialize)]
656pub struct DagSummary {
657 pub node_count: usize,
659 pub edge_count: usize,
661 pub root_count: usize,
663 pub leaf_count: usize,
665 pub max_in_degree: usize,
667 pub max_out_degree: usize,
669 pub data_edge_count: usize,
671 pub control_edge_count: usize,
673 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 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 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 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 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 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 let root_deps = dag.get_dependencies_with_edges("t1");
850 assert!(root_deps.is_empty());
851
852 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 assert!(dag.get_edge_between("t2", "t1").is_none());
900 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 assert!(!dag.has_dependencies("t1"));
929 assert!(dag.has_dependents("t1"));
930
931 assert!(dag.has_dependencies("t2"));
933 assert!(dag.has_dependents("t2"));
934
935 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 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 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 assert!(dag.remove_task("nonexistent").is_none());
993 }
994
995 #[test]
996 fn test_remove_task_keeps_task_map_consistent_after_swap() {
997 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 dag.add_dependency("t2", "t3", TaskEdge::default()).ok();
1006 dag.add_dependency("t1", "t3", TaskEdge::default()).ok();
1007
1008 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 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); assert!(dag.has_dependency("t2", "t3"));
1020 assert_eq!(dag.get_dependencies("t3"), vec!["t2".to_string()]);
1021
1022 let t2 = dag.get_task("t2").expect("t2 must still be reachable");
1024 assert_eq!(t2.name, "Task 2");
1025
1026 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 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 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 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); assert_eq!(summary.max_out_degree, 2); 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}