use async_trait::async_trait;
use std::fmt;
use std::ops::Add;
use std::sync::{Arc, RwLock};
use crate::error::AggregateError;
use crate::sync::{Interrupt, Sender};
type Link<T> = Option<Arc<RwLock<Node<T>>>>;
enum Node<T> {
Item { value: T, next: Link<T> },
Fork { next: Vec<Link<T>> },
Join { next: Link<T> },
}
impl<T> Node<T> {
pub fn item(value: T, next: Link<T>) -> Self {
Node::Item { value, next }
}
pub fn join(next: Link<T>) -> Self {
Node::Join { next }
}
pub fn fork(next: Vec<Link<T>>) -> Self {
Node::Fork { next }
}
pub fn into_link(self) -> Link<T> {
Some(Arc::new(RwLock::new(self)))
}
}
struct Iter<T> {
stack: Vec<(Link<T>, Vec<usize>)>,
}
impl<T> Iterator for Iter<T> {
type Item = Arc<RwLock<Node<T>>>;
fn next(&mut self) -> Option<Self::Item> {
while let Some((link, branching)) = self.stack.pop() {
if let Some(node_rc) = link {
let node_ref = node_rc.read().unwrap();
match &*node_ref {
Node::Item { next, .. } => {
self.stack.push((next.clone(), branching));
return Some(node_rc.clone());
}
Node::Fork { next } => {
for (i, branch_head) in next.iter().rev().enumerate() {
let mut branching = branching.clone();
branching.push(i);
self.stack.push((branch_head.clone(), branching));
}
return Some(node_rc.clone());
}
Node::Join { next } => {
let mut branching = branching;
if let Some(branch) = branching.pop() {
if branch == 0 {
self.stack.push((next.clone(), branching));
return Some(node_rc.clone());
}
}
}
}
}
}
None }
}
pub struct Dag<T> {
head: Link<T>,
tail: Link<T>,
}
impl<T: Clone> Clone for Dag<T> {
fn clone(&self) -> Self {
fn deep_clone<T: Clone>(head: Link<T>) -> (Link<T>, Link<T>) {
if let Some(node_rc) = head {
let node_ref = node_rc.read().unwrap();
match &*node_ref {
Node::Item { value, next } => {
let (next, tail) = deep_clone(next.clone());
let node = Node::Item {
value: (*value).clone(),
next,
}
.into_link();
let tail = tail.or(node.clone());
(node, tail)
}
Node::Fork { next } => {
let mut heads: Vec<Link<T>> = Vec::new();
let mut tails: Vec<Link<T>> = Vec::new();
for branch in next {
let (h, t) = deep_clone(branch.clone());
heads.push(h);
tails.push(t);
}
if let Some(tail) = tails.last() {
let (join, tail) = if let Some(tail_rc) = tail {
let next = match &*tail_rc.read().unwrap() {
Node::Item { next, .. } => next.clone(),
Node::Join { next, .. } => next.clone(),
_ => unreachable!("tail cannot be a fork"),
};
let (next, tail) = deep_clone(next);
let join = Node::Join { next }.into_link();
let tail = tail.or(join.clone());
(join, tail)
} else {
let join = Node::Join { next: None }.into_link();
(join.clone(), join)
};
for t_rc in tails.into_iter().flatten() {
match &mut *t_rc.write().unwrap() {
Node::Item { ref mut next, .. } => *next = join.clone(),
Node::Join { ref mut next, .. } => *next = join.clone(),
_ => unreachable!("tail cannot be a fork"),
}
}
(Node::Fork { next: heads }.into_link(), tail)
} else {
(None, None)
}
}
Node::Join { next } => {
(next.clone(), None)
}
}
} else {
(None, None)
}
}
let (head, tail) = deep_clone(self.head.clone());
Self { head, tail }
}
}
impl<T> Default for Dag<T> {
fn default() -> Self {
Dag {
head: None,
tail: None,
}
}
}
impl<T: PartialEq> PartialEq for Dag<T> {
fn eq(&self, other: &Self) -> bool {
for (left, rght) in self.iter().zip(other.iter()) {
if let (
Node::Item {
value: left_value, ..
},
Node::Item {
value: rght_value, ..
},
) = (&*left.read().unwrap(), &*rght.read().unwrap())
{
if left_value != rght_value {
return false;
}
} else {
return false;
}
}
true
}
}
impl<T: Eq> Eq for Dag<T> {}
impl<T> From<T> for Dag<T> {
fn from(value: T) -> Self {
Dag::seq([value])
}
}
impl<T> Dag<T> {
pub fn new(branches: impl IntoIterator<Item = Dag<T>>) -> Dag<T> {
let mut branches: Vec<Dag<T>> = branches
.into_iter()
.filter(|branch| branch.head.is_some())
.collect();
if branches.len() == 1 {
return branches.pop().unwrap();
}
let mut next: Vec<Link<T>> = Vec::new();
let tail = Node::<T>::join(None).into_link();
for branch in branches {
next.push(branch.head);
debug_assert!(branch.tail.is_some());
if let Some(tail_rc) = branch.tail {
match *tail_rc.write().unwrap() {
Node::Item { ref mut next, .. } => {
*next = tail.clone();
}
Node::Join { ref mut next } => {
*next = tail.clone();
}
Node::Fork { .. } => unreachable!(),
}
}
}
if next.is_empty() {
return Dag::default();
}
Dag {
head: Node::fork(next).into_link(),
tail,
}
}
pub fn seq(elems: impl IntoIterator<Item = impl Into<T>>) -> Dag<T> {
let mut iter = elems.into_iter();
let mut head: Link<T> = None;
let mut tail: Link<T> = None;
if let Some(value) = iter.next() {
head = Node::item(value.into(), None).into_link();
tail = head.clone();
for value in iter {
let new_node = Node::item(value.into(), None).into_link();
if let Some(tail_node) = tail {
if let Node::Item { ref mut next, .. } = *tail_node.write().unwrap() {
*next = new_node.clone();
}
}
tail = new_node;
}
}
Dag { head, tail }
}
pub fn is_empty(&self) -> bool {
self.tail.is_none()
}
pub fn concat(self, other: impl Into<Dag<T>>) -> Self {
let other = other.into();
if let Some(tail_node) = &self.tail {
match *tail_node.write().unwrap() {
Node::Item { ref mut next, .. } => {
*next = other.head;
}
Node::Join { ref mut next } => {
*next = other.head;
}
_ => unreachable!("tail cannot be a fork"),
}
} else {
return other;
}
Dag {
head: self.head,
tail: other.tail.or(self.tail),
}
}
pub fn prepend(self, other: impl Into<Dag<T>>) -> Dag<T> {
other.into().concat(self)
}
fn iter(&self) -> Iter<T> {
Iter {
stack: vec![(self.head.clone(), Vec::new())],
}
}
pub fn any(&self, condition: impl Fn(&T) -> bool) -> bool {
for node in self.iter() {
if let Node::Item { value, .. } = &*node.read().unwrap() {
if condition(value) {
return true;
}
}
}
false
}
pub fn all(&self, condition: impl Fn(&T) -> bool) -> bool {
for node in self.iter() {
if let Node::Item { value, .. } = &*node.read().unwrap() {
if !condition(value) {
return false;
}
}
}
true
}
pub fn shallow_clone(&self) -> Self {
Self {
head: self.head.clone(),
tail: self.tail.clone(),
}
}
pub fn reverse(self) -> Dag<T> {
let Dag { head, .. } = self;
let tail = head.clone();
let mut stack = vec![(head, None as Link<T>, Vec::<usize>::new())];
let mut results: Vec<Link<T>> = Vec::new();
while let Some((head, prev, branching)) = stack.pop() {
if let Some(node_rc) = head.clone() {
match *node_rc.write().unwrap() {
Node::Item { ref mut next, .. } => {
let newhead = next.clone();
*next = prev;
stack.push((newhead, head, branching));
}
Node::Fork { ref next } => {
let prev = Node::join(prev).into_link();
for (i, br_head) in next.iter().rev().enumerate() {
let mut branching = branching.clone();
branching.push(i);
stack.push((br_head.clone(), prev.clone(), branching));
}
}
Node::Join { ref next } => {
results.push(prev);
let mut branching = branching;
if let Some(branch) = branching.pop() {
if branch == 0 {
let head = Node::fork(results).into_link();
stack.push((next.clone(), head, branching));
results = Vec::new();
}
}
}
}
} else {
results.push(prev);
}
}
if results.is_empty() {
Dag::default()
} else {
debug_assert!(
results.len() == 1,
"Expected exactly one result after reversal"
);
let head = results.pop().unwrap();
Dag { head, tail }
}
}
}
impl<T, R> Add<R> for Dag<T>
where
R: Into<Dag<T>>,
{
type Output = Self;
fn add(self, other: R) -> Self {
self.concat(other)
}
}
impl<T: fmt::Display> fmt::Display for Dag<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fn fmt_node<T: fmt::Display>(
f: &mut fmt::Formatter<'_>,
node: &Node<T>,
indent: usize,
index: usize,
branching: Vec<(usize, bool)>,
) -> fmt::Result {
let fmt_newline =
|f: &mut fmt::Formatter, level: usize, condition: bool| -> fmt::Result {
if condition {
writeln!(f)?;
write!(f, "{}", " ".repeat(level))?;
}
Ok(())
};
match node {
Node::Item { value, next } => {
fmt_newline(f, indent, index > 0)?;
write!(f, "- {value}")?;
if let Some(next_rc) = next {
fmt_node(f, &*next_rc.read().unwrap(), indent, index + 1, branching)?;
}
}
Node::Fork { next } => {
fmt_newline(f, indent, index > 0)?;
write!(f, "+ ")?;
for (br_idx, branch) in next.iter().enumerate() {
if let Some(branch_head) = branch {
fmt_newline(f, indent + 1, br_idx > 0)?;
write!(f, "~ ")?;
let mut updated_branching = branching.clone();
updated_branching.push((index, br_idx == next.len() - 1));
fmt_node(
f,
&*branch_head.read().unwrap(),
indent + 2,
0,
updated_branching,
)?;
}
}
}
Node::Join { next } => {
let mut branching = branching;
if let Some((index, is_last)) = branching.pop() {
if is_last {
if let Some(next_rc) = next {
fmt_node(
f,
&*next_rc.read().unwrap(),
indent - 2,
index + 1,
branching,
)?;
}
}
}
}
}
Ok(())
}
if let Some(root) = &self.head {
fmt_node(
f,
&*root.read().unwrap(),
0, 0, Vec::new(), )?
}
Ok(())
}
}
#[macro_export]
macro_rules! seq {
($($value:expr),* $(,)?) => {
Dag::seq([$($value),*])
};
}
#[macro_export]
macro_rules! dag {
($($branch:expr),* $(,)?) => {
Dag::new([$($branch),*])
};
}
#[macro_export]
macro_rules! par {
($($value:expr),* $(,)?) => {
Dag::new([
$(Dag::seq([$value])),*
])
}
}
pub enum ExecutionStatus {
Completed,
Interrupted,
}
#[async_trait]
pub trait Task {
type Input;
type Changes;
type Error;
async fn run(&self, input: &Self::Input) -> Result<Self::Changes, Self::Error>;
}
impl<T> Dag<T>
where
T: Task + Clone,
T::Input: Clone,
{
pub async fn execute(
self,
input: &Arc<crate::sync::RwLock<T::Input>>,
channel: Sender<T::Changes>,
interrupt: Interrupt,
) -> Result<ExecutionStatus, AggregateError<T::Error>> {
enum InnerNode<T> {
Item { task: T, next: Link<T> },
Fork { branches: Vec<Link<T>> },
Join { next: Link<T> },
}
enum InnerError<E> {
Failure(Vec<E>),
Interrupted,
}
async fn run_task<T: Task>(
task: T,
value: &T::Input,
interrupt: &Interrupt,
) -> Result<T::Changes, InnerError<T::Error>> {
let future = task.run(value);
tokio::select! {
_ = interrupt.wait() => {
Err(InnerError::Interrupted)
}
result = future => {
result.map_err(|e| InnerError::Failure(vec![e]))
}
}
}
async fn exec_node<T>(
node: Link<T>,
input: &Arc<crate::sync::RwLock<T::Input>>,
channel: &Sender<T::Changes>,
interrupt: &Interrupt,
) -> Result<Link<T>, InnerError<T::Error>>
where
T: Task + Clone,
T::Input: Clone,
{
let mut current = node;
let mut errors = Vec::new();
while let Some(node_rc) = current {
if interrupt.is_set() {
return Err(InnerError::Interrupted);
}
let node = match &*node_rc.read().unwrap() {
Node::Item { value, next } => InnerNode::Item {
task: value.clone(),
next: next.clone(),
},
Node::Fork { next } => InnerNode::Fork {
branches: next.clone(),
},
Node::Join { next } => InnerNode::Join { next: next.clone() },
};
match node {
InnerNode::Item { task, next } => {
let value = {
let guard = input.read().await;
guard.clone()
};
match run_task(task, &value, interrupt).await {
Ok(changes) => {
if channel.send(changes).await.is_err() {
return Err(InnerError::Interrupted);
}
}
Err(InnerError::Interrupted) => return Err(InnerError::Interrupted),
Err(InnerError::Failure(mut err)) => {
errors.append(&mut err);
break;
}
};
current = next;
}
InnerNode::Fork { branches } => {
let mut futures = Vec::new();
for branch in branches.into_iter().filter(|b| b.is_some()) {
futures.push(exec_node(branch, input, channel, interrupt));
}
let results = futures::future::join_all(futures).await;
let mut join_next: Link<T> = None;
for res in results {
match res {
Ok(next) => {
join_next = next;
}
Err(e) => match e {
InnerError::Interrupted => return Err(InnerError::Interrupted),
InnerError::Failure(mut err) => errors.append(&mut err),
},
}
}
if !errors.is_empty() {
return Err(InnerError::Failure(errors));
}
current = join_next;
}
InnerNode::Join { next } => {
return Ok(next);
}
}
}
if errors.is_empty() {
Ok(None)
} else {
Err(InnerError::Failure(errors))
}
}
let mut next = self.head;
while next.is_some() {
next = match exec_node(next, input, &channel, &interrupt).await {
Ok(next) => next,
Err(InnerError::Interrupted) => return Ok(ExecutionStatus::Interrupted),
Err(InnerError::Failure(err)) => return Err(AggregateError(err)),
}
}
Ok(ExecutionStatus::Completed)
}
}
#[cfg(test)]
mod tests {
use async_trait::async_trait;
use dedent::dedent;
use pretty_assertions::{assert_eq, assert_str_eq};
use std::{
sync::atomic::{AtomicUsize, Ordering},
time::Instant,
};
use super::*;
use crate::sync::channel;
fn is_item<T>(node: &Arc<RwLock<Node<T>>>) -> bool {
if let Node::Item { .. } = &*node.read().unwrap() {
return true;
}
false
}
#[test]
fn test_empty_dag() {
let dag: Dag<i32> = Dag::default();
assert!(dag.head.is_none());
}
#[test]
fn test_dag_from_list() {
let elements = vec![1, 2, 3, 4];
let dag = Dag::<i32>::seq(elements.clone());
let mut head = dag.head;
for &value in &elements {
assert!(head.is_some());
if let Some(head_rc) = head {
if let Node::Item {
value: node_value,
next,
} = &*head_rc.read().unwrap()
{
assert_eq!(*node_value, value);
head = next.clone();
} else {
panic!("expected an item node");
}
}
}
assert!(head.is_none());
}
#[test]
fn test_dag_from_empty_list() {
let dag: Dag<i32> = Dag::seq(Vec::<i32>::new());
assert!(dag.is_empty());
let dag: Dag<i32> = Dag::new(vec![Dag::seq(Vec::<i32>::new())]);
assert!(dag.is_empty());
}
#[test]
fn test_dag_from_single_branch() {
let dag: Dag<i32> = dag!(seq!(1, 2, 3));
assert!(dag.head.is_some());
if let Some(head_rc) = dag.head {
let node = &*head_rc.read().unwrap();
assert!(matches!(node, Node::Item { value: 1, .. }));
}
}
#[test]
fn test_dag_construction() {
let dag: Dag<i32> = seq!(1, 2, 3, 4);
assert!(dag.head.is_some());
if let Some(head_rc) = dag.head {
let node = &*head_rc.read().unwrap();
assert!(matches!(node, Node::Item { value: 1, .. }));
}
assert!(dag.tail.is_some());
if let Some(tail_rc) = dag.tail {
let node = &*tail_rc.read().unwrap();
assert!(matches!(node, Node::Item { value: 4, .. }));
}
}
#[test]
fn test_clone_sequence() {
let dag: Dag<i32> = seq!(1, 2, 3);
let clone = dag.clone();
assert_eq!(clone.to_string(), "- 1\n- 2\n- 3");
}
#[test]
fn test_clone_fork() {
let dag: Dag<i32> = seq!(1) + par!(2, 3, 4) + seq!(5);
let clone = dag.clone();
assert_eq!(clone.to_string(), "- 1\n+ ~ - 2\n ~ - 3\n ~ - 4\n- 5");
}
#[test]
fn test_clone_deep_nested_dag() {
let dag: Dag<char> = seq!('A')
+ dag!(
seq!('B', 'C') + dag!(seq!('D', 'E'), seq!('F')),
seq!('G', 'H', 'I')
)
+ seq!('J', 'K');
let dag = dag.clone();
assert_str_eq!(
dag.to_string(),
dedent!(
r#"
- A
+ ~ - B
- C
+ ~ - D
- E
~ - F
~ - G
- H
- I
- J
- K
"#
)
);
}
#[test]
fn test_iterate_linear_graph() {
let elements = vec![1, 2, 3];
let dag = Dag::<i32>::seq(elements.clone());
let mut result = Vec::new();
for node in dag.iter() {
let node_ref = node.read().unwrap();
match &*node_ref {
Node::Item { value, .. } => result.push(*value), Node::Fork { .. } => panic!("unexpected fork node in a linear graph"),
Node::Join { .. } => panic!("unexpected join node in a linear graph"),
}
}
assert_eq!(result, elements);
}
#[test]
fn test_iterate_forked_graph() {
let dag: Dag<i32> = seq!(1, 2)
+ dag!(
seq!(3) + dag!(seq!(4, 5), dag!(seq!(6), seq!(7)) + seq!(8)) + seq!(9),
seq!(10) + dag!(seq!(11), seq!(12)),
)
+ seq!(13);
let elems: Vec<i32> = dag
.iter()
.filter(is_item)
.map(|node| match &*node.read().unwrap() {
Node::Item { value, .. } => *value,
_ => unreachable!(),
})
.collect();
assert_eq!(elems, vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13])
}
#[test]
fn test_empty_dag_string_representation() {
let dag: Dag<char> = Dag::default();
assert_eq!(dag.to_string(), "");
}
#[test]
fn converts_linked_list_to_string() {
let dag: Dag<char> = seq!('A', 'B', 'C', 'D');
assert_str_eq!(
dag.to_string(),
dedent!(
r#"
- A
- B
- C
- D
"#
)
);
}
#[test]
fn modifying_a_clone_should_not_affect_the_original() {
let dag: Dag<char> = seq!('A') + par!('B', 'C', 'D');
let new_dag = dag.clone() + seq!('E');
assert_str_eq!(
new_dag.to_string(),
dedent!(
r#"
- A
+ ~ - B
~ - C
~ - D
- E
"#
),
"new dag should contain the new element"
);
assert_str_eq!(
dag.to_string(),
dedent!(
r#"
- A
+ ~ - B
~ - C
~ - D
"#
),
"old dag should remain the same"
);
}
#[test]
fn test_concatenation_with_empty_dag() {
let non_empty: Dag<i32> = seq!(1, 2, 3);
let empty: Dag<i32> = Dag::default();
assert!(!non_empty.is_empty());
assert!(empty.is_empty());
let result = non_empty.clone() + empty.clone();
assert!(!result.is_empty());
assert_eq!(result.to_string(), "- 1\n- 2\n- 3");
let result2 = empty + non_empty;
assert!(!result2.is_empty());
assert_eq!(result2.to_string(), "- 1\n- 2\n- 3");
}
#[test]
fn test_concatenation_with_forked_empty_dag() {
let non_empty: Dag<i32> = seq!(1, 2);
let forked_with_empty: Dag<i32> = dag!(seq!(3), Dag::default());
let result = non_empty + forked_with_empty;
assert!(!result.is_empty());
assert_eq!(result.to_string(), "- 1\n- 2\n- 3");
}
#[test]
fn test_empty_dag_concatenation_preserves_tail() {
let first: Dag<i32> = seq!(1);
let second: Dag<i32> = Dag::default(); let third: Dag<i32> = seq!(2);
let result = first + second + third;
assert!(!result.is_empty());
assert_eq!(result.to_string(), "- 1\n- 2");
}
#[test]
fn test_basic_concatenation_of_sequences() {
let first: Dag<i32> = seq!(1, 2);
let second: Dag<i32> = Dag::default(); let third: Dag<i32> = seq!(3);
let result = first + second + third;
assert!(!result.is_empty());
assert_eq!(result.to_string(), "- 1\n- 2\n- 3");
}
#[test]
fn test_basic_prepend() {
let first: Dag<i32> = seq!(1, 2);
let second: Dag<i32> = Dag::default(); let third: Dag<i32> = seq!(3);
let result = third.prepend(second).prepend(first);
assert!(!result.is_empty());
assert_eq!(result.to_string(), "- 1\n- 2\n- 3");
}
#[test]
fn test_prepend_with_forked_empty_dag() {
let non_empty: Dag<i32> = seq!(1, 2);
let forked_with_empty: Dag<i32> = dag!(seq!(3), Dag::default());
let result = forked_with_empty.prepend(non_empty);
assert!(!result.is_empty());
assert_eq!(result.to_string(), "- 1\n- 2\n- 3");
}
#[test]
fn test_dag_new_with_shared_nodes() {
let single_element: Dag<i32> = seq!(42);
let branch1 = single_element.clone();
let branch2 = seq!(1, 2);
let forked_dag = dag!(branch1, branch2);
assert_eq!(single_element.to_string(), "- 42");
assert_eq!(forked_dag.to_string(), "+ ~ - 42\n ~ - 1\n - 2");
}
#[test]
fn test_dag_new_with_multiple_single_elements() {
let elem1: Dag<i32> = seq!(1);
let elem2: Dag<i32> = seq!(2);
let elem3: Dag<i32> = seq!(3);
let forked = dag!(elem1.clone(), elem2.clone(), elem3.clone());
assert_eq!(elem1.to_string(), "- 1");
assert_eq!(elem2.to_string(), "- 2");
assert_eq!(elem3.to_string(), "- 3");
assert_eq!(forked.to_string(), "+ ~ - 1\n ~ - 2\n ~ - 3");
}
#[test]
fn test_dag_new_edge_cases() {
let empty1: Dag<i32> = Dag::default();
let empty2: Dag<i32> = Dag::default();
let non_empty: Dag<i32> = seq!(42);
let all_empty = dag!(empty1.clone(), empty2.clone());
assert!(all_empty.is_empty());
let mixed = dag!(empty1, non_empty.clone(), empty2);
assert_eq!(mixed.to_string(), "- 42");
let single_branch = dag!(non_empty);
assert_eq!(single_branch.to_string(), "- 42");
}
#[test]
fn test_dag_seq_edge_cases() {
let empty_seq: Dag<i32> = Dag::seq(Vec::<i32>::new());
assert!(empty_seq.is_empty());
assert_eq!(empty_seq.to_string(), "");
let single: Dag<i32> = seq!(42);
assert!(!single.is_empty());
assert_eq!(single.to_string(), "- 42");
}
#[test]
fn converts_branching_dag_to_string() {
let dag: Dag<char> = dag!(seq!('A', 'B'), seq!('C', 'D', 'E')) + seq!('F');
assert_str_eq!(
dag.to_string(),
dedent!(
r#"
+ ~ - A
- B
~ - C
- D
- E
- F
"#
)
);
}
#[test]
fn converts_complex_dag_to_string() {
let dag: Dag<char> = seq!('A')
+ dag!(
seq!('B', 'C') + dag!(seq!('D', 'E'), seq!('F')),
seq!('G', 'H', 'I')
)
+ seq!('J', 'K');
assert_str_eq!(
dag.to_string(),
dedent!(
r#"
- A
+ ~ - B
- C
+ ~ - D
- E
~ - F
~ - G
- H
- I
- J
- K
"#
)
);
}
#[test]
fn converts_numeric_dag_to_string() {
let dag: Dag<i32> = seq!(1, 2)
+ dag!(
seq!(3) + dag!(seq!(4, 5), dag!(seq!(6), seq!(7)) + seq!(8)) + seq!(9),
seq!(10) + par!(11, 12),
)
+ seq!(13);
assert_str_eq!(
dag.to_string(),
dedent!(
r#"
- 1
- 2
+ ~ - 3
+ ~ - 4
- 5
~ + ~ - 6
~ - 7
- 8
- 9
~ - 10
+ ~ - 11
~ - 12
- 13
"#
)
)
}
#[tokio::test]
async fn it_executes_simple_dag() {
#[derive(Clone)]
struct DummyTask;
#[async_trait]
impl Task for DummyTask {
type Input = ();
type Changes = ();
type Error = ();
async fn run(&self, _input: &Self::Input) -> Result<Self::Changes, Self::Error> {
Ok(())
}
}
let dag: Dag<DummyTask> = seq!(DummyTask, DummyTask, DummyTask);
let reader = Arc::new(tokio::sync::RwLock::new(()));
let (tx, mut rx) = channel(10);
let sigint = Interrupt::new();
let count_atomic = Arc::new(AtomicUsize::new(0));
let counter = count_atomic.clone();
tokio::spawn(async move {
while let Some(msg) = rx.recv().await {
let c = counter.load(Ordering::Relaxed);
counter.store(c + 1, Ordering::Relaxed);
msg.ack();
}
});
let result = dag.execute(&reader, tx, sigint).await;
assert!(matches!(result, Ok(ExecutionStatus::Completed)));
assert_eq!(count_atomic.load(Ordering::Relaxed), 3);
}
#[derive(Clone)]
struct SleepyTask {
pub name: &'static str,
pub delay_ms: u64,
}
#[async_trait]
impl Task for SleepyTask {
type Input = ();
type Changes = &'static str;
type Error = ();
async fn run(&self, _input: &Self::Input) -> Result<Self::Changes, Self::Error> {
tokio::time::sleep(std::time::Duration::from_millis(self.delay_ms)).await;
Ok(self.name)
}
}
#[tokio::test]
async fn test_concurrent_execution() {
let task_a = SleepyTask {
name: "A",
delay_ms: 100,
};
let task_b = SleepyTask {
name: "B",
delay_ms: 100,
};
let task_c = SleepyTask {
name: "C",
delay_ms: 0,
};
let dag: Dag<SleepyTask> = dag!(seq!(task_a), seq!(task_b)) + seq!(task_c);
let input = Arc::new(tokio::sync::RwLock::new(()));
let (tx, mut rx) = channel::<&'static str>(10);
let sigint = Interrupt::new();
let start = Instant::now();
let results = Arc::new(tokio::sync::RwLock::new(Vec::new()));
{
let results = Arc::clone(&results);
tokio::spawn(async move {
while let Some(msg) = rx.recv().await {
let mut res = results.write().await;
res.push(msg.data);
msg.ack();
}
});
}
let exec_result = dag.execute(&input, tx, sigint).await;
let elapsed = start.elapsed();
assert!(matches!(exec_result, Ok(ExecutionStatus::Completed)));
let results = results.read().await;
assert_eq!(*results, vec!["A", "B", "C"]);
assert!(
elapsed.as_millis() < 200,
"Execution took too long, not concurrent!"
);
}
#[tokio::test]
async fn test_interrupt_during_execution() {
let dag: Dag<SleepyTask> = seq!(
SleepyTask {
name: "A",
delay_ms: 100
},
SleepyTask {
name: "B",
delay_ms: 100
},
SleepyTask {
name: "C",
delay_ms: 100
}
);
let input = Arc::new(tokio::sync::RwLock::new(()));
let (tx, mut rx) = channel::<&'static str>(10);
let interrupt = Interrupt::new();
let interrupt_clone = interrupt.clone();
let results = Arc::new(tokio::sync::RwLock::new(Vec::new()));
{
let results = Arc::clone(&results);
tokio::spawn(async move {
while let Some(msg) = rx.recv().await {
let mut res = results.write().await;
res.push(msg.data);
msg.ack();
}
});
}
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
interrupt_clone.trigger();
});
let exec_result = dag.execute(&input, tx, interrupt).await;
assert!(matches!(exec_result, Ok(ExecutionStatus::Interrupted)));
let results = results.read().await;
assert!(
results.len() < 3,
"Expected partial execution but got all results"
);
}
#[derive(Clone)]
struct MaybeFailTask {
name: &'static str,
fail: bool,
}
#[async_trait]
impl Task for MaybeFailTask {
type Input = ();
type Changes = &'static str;
type Error = &'static str; async fn run(&self, _input: &Self::Input) -> Result<Self::Changes, Self::Error> {
if self.fail {
Err("task failed")
} else {
Ok(self.name)
}
}
}
#[tokio::test]
async fn test_error_interrupts_execution() {
let dag: Dag<MaybeFailTask> = dag!(
dag!(
seq!(
MaybeFailTask {
name: "A",
fail: false
},
MaybeFailTask {
name: "B",
fail: false
}
),
seq!(
MaybeFailTask {
name: "C",
fail: true
},
MaybeFailTask {
name: "D",
fail: false
}
)
),
seq!(MaybeFailTask {
name: "E",
fail: false
})
) + seq!(MaybeFailTask {
name: "F",
fail: false
});
let input = Arc::new(tokio::sync::RwLock::new(()));
let (tx, mut rx) = channel::<&'static str>(10);
let interrupt = Interrupt::new();
let results = Arc::new(tokio::sync::RwLock::new(Vec::new()));
{
let results = Arc::clone(&results);
tokio::spawn(async move {
while let Some(msg) = rx.recv().await {
let mut res = results.write().await;
res.push(msg.data);
msg.ack();
}
});
}
let exec_result = dag.execute(&input, tx, interrupt).await;
assert!(exec_result.is_err(), "Expected execution to fail on error");
let results = results.read().await;
assert_eq!(*results, vec!["A", "E", "B"]);
}
#[test]
fn test_contructing_linear_inverted_dag() {
let dag: Dag<i32> = Dag::default().prepend(1).prepend(2).prepend(3);
let elems: Vec<i32> = dag
.iter()
.filter(is_item)
.map(|node| match &*node.read().unwrap() {
Node::Item { value, .. } => *value,
_ => unreachable!(),
})
.collect();
assert_eq!(elems, vec![3, 2, 1])
}
#[test]
fn test_contructing_forking_inverted_dag() {
let dag: Dag<i32> = Dag::default()
.prepend(1)
.prepend(2)
.prepend(3)
.prepend(par!(5, 4))
.prepend(6);
let elems: Vec<i32> = dag
.iter()
.filter(is_item)
.map(|node| match &*node.read().unwrap() {
Node::Item { value, .. } => *value,
_ => unreachable!(),
})
.collect();
assert_eq!(elems, vec![6, 5, 4, 3, 2, 1])
}
#[test]
fn test_reverse_dag() {
let dag: Dag<i32> = Dag::default()
.prepend(1)
.prepend(2)
.prepend(3)
.prepend(par!(4, 5))
.prepend(6)
.reverse();
let elems: Vec<i32> = dag
.iter()
.filter(is_item)
.map(|node| match &*node.read().unwrap() {
Node::Item { value, .. } => *value,
_ => unreachable!(),
})
.collect();
assert_eq!(elems, vec![1, 2, 3, 4, 5, 6])
}
#[test]
fn test_reverse_empty_dag() {
let dag: Dag<i32> = Dag::default().reverse();
assert!(dag.head.is_none());
assert!(dag.tail.is_none());
}
#[test]
fn test_reverse_dag_with_forks() {
let dag: Dag<char> = seq!('A') + dag!(seq!('B', 'C'), seq!('D')) + seq!('E');
let reversed = dag.reverse();
assert_str_eq!(
reversed.to_string(),
dedent!(
r#"
- E
+ ~ - C
- B
~ - D
- A
"#
)
);
}
#[test]
fn test_reverse_single_item() {
let dag: Dag<i32> = seq!(42);
let reversed = dag.reverse();
let elems: Vec<i32> = reversed
.iter()
.filter(is_item)
.map(|node| match &*node.read().unwrap() {
Node::Item { value, .. } => *value,
_ => unreachable!(),
})
.collect();
assert_eq!(elems, vec![42]);
}
#[test]
fn test_reverse_nested_forks() {
let inner_fork = dag!(seq!('C'), seq!('D'));
let branch1 = seq!('B') + inner_fork;
let branch2 = seq!('E');
let dag: Dag<char> = seq!('A') + dag!(branch1, branch2) + seq!('F');
let reversed = dag.reverse();
assert_str_eq!(
reversed.to_string(),
dedent!(
r#"
- F
+ ~ + ~ - C
~ - D
- B
~ - E
- A
"#
)
);
}
#[test]
fn test_reverse_multiple_sequential_sections() {
let dag: Dag<i32> = seq!(1, 2) + dag!(seq!(3, 4), seq!(5, 6)) + seq!(7, 8);
let reversed = dag.reverse();
let elems: Vec<i32> = reversed
.iter()
.filter(is_item)
.map(|node| match &*node.read().unwrap() {
Node::Item { value, .. } => *value,
_ => unreachable!(),
})
.collect();
assert_eq!(elems, vec![8, 7, 4, 3, 6, 5, 2, 1]);
}
#[test]
fn test_reverse_three_way_fork() {
let dag: Dag<char> = seq!('A') + dag!(seq!('B'), seq!('C'), seq!('D')) + seq!('E');
let reversed = dag.reverse();
assert_str_eq!(
reversed.to_string(),
dedent!(
r#"
- E
+ ~ - B
~ - C
~ - D
- A
"#
)
);
}
#[test]
fn test_reverse_preserves_execution_semantics() {
let original: Dag<i32> = seq!(1, 2) + dag!(seq!(3, 4), seq!(5)) + seq!(6);
let double_reversed = original.shallow_clone().reverse().reverse();
let original_elems: Vec<i32> = original
.iter()
.filter(is_item)
.map(|node| match &*node.read().unwrap() {
Node::Item { value, .. } => *value,
_ => unreachable!(),
})
.collect();
let double_reversed_elems: Vec<i32> = double_reversed
.iter()
.filter(is_item)
.map(|node| match &*node.read().unwrap() {
Node::Item { value, .. } => *value,
_ => unreachable!(),
})
.collect();
assert_eq!(original_elems, double_reversed_elems);
}
#[test]
fn test_reverse_empty_branches() {
let dag: Dag<i32> = dag!(seq!(1, 2), Dag::default(), seq!(3));
let reversed = dag.reverse();
let elems: Vec<i32> = reversed
.iter()
.filter(is_item)
.map(|node| match &*node.read().unwrap() {
Node::Item { value, .. } => *value,
_ => unreachable!(),
})
.collect();
assert_eq!(elems, vec![2, 1, 3]);
}
}