use std::collections::{BTreeMap, BTreeSet};
use super::{DagValidationError, TaskDag, TaskEvent, TaskRun, TaskState, validate_dag};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TaskEventError {
InvalidDag(DagValidationError),
UnknownTask(String),
DependenciesNotPassed {
task_id: String,
blocked_by: Vec<String>,
},
TerminalTask {
task_id: String,
state: TaskState,
},
InvalidTransition {
task_id: String,
from: TaskState,
to: TaskState,
},
}
pub fn initial_run(dag: &TaskDag) -> Result<TaskRun, DagValidationError> {
validate_dag(dag)?;
let mut states = dag
.nodes
.iter()
.map(|node| (node.id.clone(), TaskState::Pending))
.collect::<BTreeMap<_, _>>();
mark_ready(dag, &mut states);
Ok(TaskRun {
dag_id: dag.id.clone(),
states,
events: Vec::new(),
})
}
pub fn ready_tasks(dag: &TaskDag, run: &TaskRun) -> Result<Vec<String>, TaskEventError> {
validate_dag(dag).map_err(TaskEventError::InvalidDag)?;
ensure_run_matches_dag(dag, run)?;
Ok(dag
.nodes
.iter()
.filter(|node| run.states.get(&node.id) == Some(&TaskState::Ready))
.map(|node| node.id.clone())
.collect::<BTreeSet<_>>()
.into_iter()
.collect())
}
pub fn apply_task_event(
dag: &TaskDag,
run: &TaskRun,
event: TaskEvent,
) -> Result<TaskRun, TaskEventError> {
validate_dag(dag).map_err(TaskEventError::InvalidDag)?;
ensure_run_matches_dag(dag, run)?;
let Some(current) = run.states.get(&event.task_id).copied() else {
return Err(TaskEventError::UnknownTask(event.task_id));
};
if matches!(
current,
TaskState::Passed | TaskState::Failed | TaskState::Skipped
) {
return Err(TaskEventError::TerminalTask {
task_id: event.task_id,
state: current,
});
}
validate_transition(dag, run, &event.task_id, current, event.to)?;
let mut next = run.clone();
next.states.insert(event.task_id.clone(), event.to);
next.events.push(event);
mark_ready(dag, &mut next.states);
Ok(next)
}
pub(super) fn ensure_run_matches_dag(dag: &TaskDag, run: &TaskRun) -> Result<(), TaskEventError> {
for node in &dag.nodes {
if !run.states.contains_key(&node.id) {
return Err(TaskEventError::UnknownTask(node.id.clone()));
}
}
Ok(())
}
fn validate_transition(
dag: &TaskDag,
run: &TaskRun,
task_id: &str,
from: TaskState,
to: TaskState,
) -> Result<(), TaskEventError> {
match to {
TaskState::Ready | TaskState::Pending => Err(TaskEventError::InvalidTransition {
task_id: task_id.to_string(),
from,
to,
}),
TaskState::Running => {
let blocked_by = deps_passed(dag, run, task_id);
if !blocked_by.is_empty() {
return Err(TaskEventError::DependenciesNotPassed {
task_id: task_id.to_string(),
blocked_by,
});
}
if !matches!(from, TaskState::Ready) {
return Err(TaskEventError::InvalidTransition {
task_id: task_id.to_string(),
from,
to,
});
}
Ok(())
}
TaskState::Passed | TaskState::Failed => {
if from == TaskState::Running {
Ok(())
} else {
Err(TaskEventError::InvalidTransition {
task_id: task_id.to_string(),
from,
to,
})
}
}
TaskState::Blocked => {
if matches!(
from,
TaskState::Pending | TaskState::Ready | TaskState::Running
) {
Ok(())
} else {
Err(TaskEventError::InvalidTransition {
task_id: task_id.to_string(),
from,
to,
})
}
}
TaskState::Skipped => {
if matches!(
from,
TaskState::Pending | TaskState::Ready | TaskState::Blocked
) {
Ok(())
} else {
Err(TaskEventError::InvalidTransition {
task_id: task_id.to_string(),
from,
to,
})
}
}
}
}
fn mark_ready(dag: &TaskDag, states: &mut BTreeMap<String, TaskState>) {
for node in &dag.nodes {
if states.get(&node.id) == Some(&TaskState::Pending)
&& node
.depends_on
.iter()
.all(|dep| states.get(dep) == Some(&TaskState::Passed))
{
states.insert(node.id.clone(), TaskState::Ready);
}
}
}
pub(super) fn deps_passed(dag: &TaskDag, run: &TaskRun, task_id: &str) -> Vec<String> {
dag.nodes
.iter()
.find(|node| node.id == task_id)
.map(|node| {
node.depends_on
.iter()
.filter(|dep| run.states.get(*dep) != Some(&TaskState::Passed))
.cloned()
.collect()
})
.unwrap_or_default()
}