Skip to main content

eidos_kernel/workflow/
run.rs

1use std::collections::{BTreeMap, BTreeSet};
2
3use super::{DagValidationError, TaskDag, TaskEvent, TaskRun, TaskState, validate_dag};
4
5/// Invalid state-transition attempts.
6#[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
25/// Create an initial run projection and mark dependency-free tasks as ready.
26pub 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
41/// Return ready task ids in deterministic order for the current run projection.
42pub 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
55/// Apply a state event and return the next run projection.
56///
57/// This reducer is deliberately conservative: only ready tasks can start, only running tasks can
58/// pass/fail, terminal tasks cannot be changed, and newly unblocked pending tasks become ready.
59pub 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}