use std::collections::{HashMap, HashSet};
use serde::{Deserialize, Serialize};
use super::types::{TaskId, TaskState, TaskStatus};
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct StatusCounts {
pub submitted: usize,
pub running: usize,
pub waiting: usize,
pub blocked: usize,
pub done: usize,
pub failed: usize,
}
impl StatusCounts {
pub fn total(&self) -> usize {
self.submitted + self.running + self.waiting + self.blocked + self.done + self.failed
}
}
#[derive(Debug, Default, Clone, Serialize, Deserialize)]
pub struct WorkflowState {
pub(super) tasks: HashMap<TaskId, TaskState>,
pub(super) cancelled: HashSet<TaskId>,
pub(super) children: HashMap<TaskId, Vec<TaskId>>,
pub(super) parents: HashMap<TaskId, TaskId>,
}
impl WorkflowState {
pub fn new() -> Self {
Self::default()
}
pub fn get(&self, id: TaskId) -> Option<TaskState> {
self.tasks.get(&id).copied()
}
pub fn len(&self) -> usize {
self.tasks.len()
}
pub fn is_empty(&self) -> bool {
self.tasks.is_empty()
}
pub fn contains(&self, id: TaskId) -> bool {
self.tasks.contains_key(&id)
}
pub fn all(&self) -> impl Iterator<Item = (TaskId, TaskState)> + '_ {
self.tasks.iter().map(|(id, st)| (*id, *st))
}
pub fn in_status(&self, status: TaskStatus) -> impl Iterator<Item = TaskId> + '_ {
self.tasks
.iter()
.filter(move |(_, st)| st.status == status)
.map(|(id, _)| *id)
}
pub fn is_cancel_requested(&self, id: TaskId) -> bool {
self.cancelled.contains(&id)
}
pub fn cancel_requested_count(&self) -> usize {
self.cancelled.len()
}
pub fn children_of(&self, id: TaskId) -> &[TaskId] {
self.children.get(&id).map(Vec::as_slice).unwrap_or(&[])
}
pub fn parent_of(&self, id: TaskId) -> Option<TaskId> {
self.parents.get(&id).copied()
}
pub fn descendants(&self, id: TaskId) -> Vec<TaskId> {
let mut out = Vec::new();
let mut seen = std::collections::HashSet::new();
let mut queue = std::collections::VecDeque::from(self.children_of(id).to_vec());
while let Some(t) = queue.pop_front() {
if !seen.insert(t) {
continue; }
out.push(t);
queue.extend(self.children_of(t).iter().copied());
}
out
}
pub fn subtree(&self, id: TaskId) -> Vec<TaskId> {
let mut all = vec![id];
all.extend(self.descendants(id));
all
}
pub fn status_counts(&self) -> StatusCounts {
let mut c = StatusCounts::default();
for st in self.tasks.values() {
match st.status {
TaskStatus::Submitted => c.submitted += 1,
TaskStatus::Running => c.running += 1,
TaskStatus::Waiting => c.waiting += 1,
TaskStatus::Blocked => c.blocked += 1,
TaskStatus::Done => c.done += 1,
TaskStatus::Failed => c.failed += 1,
}
}
c
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_state() {
let s = WorkflowState::new();
assert!(s.is_empty());
assert_eq!(s.len(), 0);
assert!(s.get(1).is_none());
assert!(!s.contains(1));
}
#[test]
fn in_status_filters() {
let mut s = WorkflowState::new();
s.tasks.insert(1, TaskState::submitted());
s.tasks.insert(
2,
TaskState {
step: 3,
status: TaskStatus::Running,
attempts: 0,
},
);
s.tasks.insert(
3,
TaskState {
step: 9,
status: TaskStatus::Done,
attempts: 1,
},
);
assert_eq!(s.len(), 3);
let running: Vec<TaskId> = s.in_status(TaskStatus::Running).collect();
assert_eq!(running, vec![2]);
assert_eq!(s.in_status(TaskStatus::Submitted).count(), 1);
assert_eq!(s.in_status(TaskStatus::Done).count(), 1);
assert_eq!(s.get(2).unwrap().step, 3);
}
}