eidos_kernel/workflow/
run.rs1use std::collections::{BTreeMap, BTreeSet};
2
3use super::{DagValidationError, TaskDag, TaskEvent, TaskRun, TaskState, validate_dag};
4
5#[derive(Debug, Clone, PartialEq, Eq)]
7pub enum TaskEventError {
8 InvalidDag(DagValidationError),
9 UnknownTask(String),
10 DependenciesNotPassed {
11 task_id: String,
12 blocked_by: Vec<String>,
13 },
14 TerminalTask {
15 task_id: String,
16 state: TaskState,
17 },
18 InvalidTransition {
19 task_id: String,
20 from: TaskState,
21 to: TaskState,
22 },
23}
24
25pub fn initial_run(dag: &TaskDag) -> Result<TaskRun, DagValidationError> {
27 validate_dag(dag)?;
28 let mut states = dag
29 .nodes
30 .iter()
31 .map(|node| (node.id.clone(), TaskState::Pending))
32 .collect::<BTreeMap<_, _>>();
33 mark_ready(dag, &mut states);
34 Ok(TaskRun {
35 dag_id: dag.id.clone(),
36 states,
37 events: Vec::new(),
38 })
39}
40
41pub fn ready_tasks(dag: &TaskDag, run: &TaskRun) -> Result<Vec<String>, TaskEventError> {
43 validate_dag(dag).map_err(TaskEventError::InvalidDag)?;
44 ensure_run_matches_dag(dag, run)?;
45 Ok(dag
46 .nodes
47 .iter()
48 .filter(|node| run.states.get(&node.id) == Some(&TaskState::Ready))
49 .map(|node| node.id.clone())
50 .collect::<BTreeSet<_>>()
51 .into_iter()
52 .collect())
53}
54
55pub fn apply_task_event(
60 dag: &TaskDag,
61 run: &TaskRun,
62 event: TaskEvent,
63) -> Result<TaskRun, TaskEventError> {
64 validate_dag(dag).map_err(TaskEventError::InvalidDag)?;
65 ensure_run_matches_dag(dag, run)?;
66
67 let Some(current) = run.states.get(&event.task_id).copied() else {
68 return Err(TaskEventError::UnknownTask(event.task_id));
69 };
70 if matches!(
71 current,
72 TaskState::Passed | TaskState::Failed | TaskState::Skipped
73 ) {
74 return Err(TaskEventError::TerminalTask {
75 task_id: event.task_id,
76 state: current,
77 });
78 }
79
80 validate_transition(dag, run, &event.task_id, current, event.to)?;
81
82 let mut next = run.clone();
83 next.states.insert(event.task_id.clone(), event.to);
84 next.events.push(event);
85 mark_ready(dag, &mut next.states);
86 Ok(next)
87}
88
89pub(super) fn ensure_run_matches_dag(dag: &TaskDag, run: &TaskRun) -> Result<(), TaskEventError> {
90 for node in &dag.nodes {
91 if !run.states.contains_key(&node.id) {
92 return Err(TaskEventError::UnknownTask(node.id.clone()));
93 }
94 }
95 Ok(())
96}
97
98fn validate_transition(
99 dag: &TaskDag,
100 run: &TaskRun,
101 task_id: &str,
102 from: TaskState,
103 to: TaskState,
104) -> Result<(), TaskEventError> {
105 match to {
106 TaskState::Ready | TaskState::Pending => Err(TaskEventError::InvalidTransition {
107 task_id: task_id.to_string(),
108 from,
109 to,
110 }),
111 TaskState::Running => {
112 let blocked_by = deps_passed(dag, run, task_id);
113 if !blocked_by.is_empty() {
114 return Err(TaskEventError::DependenciesNotPassed {
115 task_id: task_id.to_string(),
116 blocked_by,
117 });
118 }
119 if !matches!(from, TaskState::Ready) {
120 return Err(TaskEventError::InvalidTransition {
121 task_id: task_id.to_string(),
122 from,
123 to,
124 });
125 }
126 Ok(())
127 }
128 TaskState::Passed | TaskState::Failed => {
129 if from == TaskState::Running {
130 Ok(())
131 } else {
132 Err(TaskEventError::InvalidTransition {
133 task_id: task_id.to_string(),
134 from,
135 to,
136 })
137 }
138 }
139 TaskState::Blocked => {
140 if matches!(
141 from,
142 TaskState::Pending | TaskState::Ready | TaskState::Running
143 ) {
144 Ok(())
145 } else {
146 Err(TaskEventError::InvalidTransition {
147 task_id: task_id.to_string(),
148 from,
149 to,
150 })
151 }
152 }
153 TaskState::Skipped => {
154 if matches!(
155 from,
156 TaskState::Pending | TaskState::Ready | TaskState::Blocked
157 ) {
158 Ok(())
159 } else {
160 Err(TaskEventError::InvalidTransition {
161 task_id: task_id.to_string(),
162 from,
163 to,
164 })
165 }
166 }
167 }
168}
169
170fn mark_ready(dag: &TaskDag, states: &mut BTreeMap<String, TaskState>) {
171 for node in &dag.nodes {
172 if states.get(&node.id) == Some(&TaskState::Pending)
173 && node
174 .depends_on
175 .iter()
176 .all(|dep| states.get(dep) == Some(&TaskState::Passed))
177 {
178 states.insert(node.id.clone(), TaskState::Ready);
179 }
180 }
181}
182
183pub(super) fn deps_passed(dag: &TaskDag, run: &TaskRun, task_id: &str) -> Vec<String> {
184 dag.nodes
185 .iter()
186 .find(|node| node.id == task_id)
187 .map(|node| {
188 node.depends_on
189 .iter()
190 .filter(|dep| run.states.get(*dep) != Some(&TaskState::Passed))
191 .cloned()
192 .collect()
193 })
194 .unwrap_or_default()
195}