use std::collections::{BTreeMap, BTreeSet};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct RetryPolicy {
#[serde(default = "one")]
pub max_attempts: u32,
}
fn one() -> u32 {
1
}
impl Default for RetryPolicy {
fn default() -> Self {
Self { max_attempts: 1 }
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Step {
pub id: String,
pub function: String,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub depends_on: Vec<String>,
#[serde(default)]
pub retry: RetryPolicy,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub compensate: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Workflow {
pub name: String,
pub steps: Vec<Step>,
}
impl Workflow {
pub fn validate(&self) -> Result<(), String> {
if self.steps.is_empty() {
return Err("workflow has no steps".to_string());
}
let mut ids = BTreeSet::new();
for step in &self.steps {
if !ids.insert(step.id.as_str()) {
return Err(format!("duplicate step id {:?}", step.id));
}
}
for step in &self.steps {
for dep in &step.depends_on {
if !ids.contains(dep.as_str()) {
return Err(format!("step {:?} depends on unknown {dep:?}", step.id));
}
if dep == &step.id {
return Err(format!("step {:?} depends on itself", step.id));
}
}
}
self.check_acyclic()
}
pub fn step(&self, id: &str) -> Option<&Step> {
self.steps.iter().find(|s| s.id == id)
}
fn check_acyclic(&self) -> Result<(), String> {
let mut state: BTreeMap<&str, u8> =
self.steps.iter().map(|s| (s.id.as_str(), 0u8)).collect();
for root in &self.steps {
if state[root.id.as_str()] != 0 {
continue;
}
let mut stack: Vec<(&str, usize)> = vec![(root.id.as_str(), 0)];
while let Some((id, idx)) = stack.pop() {
let step = self.step(id).expect("id came from steps");
if idx == 0 {
state.insert(id, 1);
}
if idx < step.depends_on.len() {
stack.push((id, idx + 1));
let dep = step.depends_on[idx].as_str();
match state[dep] {
1 => return Err(format!("cycle through step {dep:?}")),
0 => stack.push((dep, 0)),
_ => {}
}
} else {
state.insert(id, 2);
}
}
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum StepStatus {
#[default]
Pending,
Running,
Succeeded,
Failed,
Compensated,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum WorkflowStatus {
#[default]
Running,
Succeeded,
Failed,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct StepRun {
pub id: String,
pub status: StepStatus,
#[serde(default)]
pub attempts: u32,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_b64: Option<String>,
pub updated: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct WorkflowRun {
pub id: String,
pub workflow: String,
pub status: WorkflowStatus,
pub steps: BTreeMap<String, StepRun>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub completed_order: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_b64: Option<String>,
pub created: u64,
pub updated: u64,
}
impl WorkflowRun {
pub fn start(
workflow: &Workflow,
id: impl Into<String>,
input_b64: Option<String>,
now: u64,
) -> Self {
let steps = workflow
.steps
.iter()
.map(|s| {
(
s.id.clone(),
StepRun {
id: s.id.clone(),
status: StepStatus::Pending,
attempts: 0,
output_b64: None,
updated: now,
},
)
})
.collect();
Self {
id: id.into(),
workflow: workflow.name.clone(),
status: WorkflowStatus::Running,
steps,
completed_order: Vec::new(),
input_b64,
created: now,
updated: now,
}
}
pub fn is_terminal(&self) -> bool {
matches!(
self.status,
WorkflowStatus::Succeeded | WorkflowStatus::Failed
)
}
pub fn ready_steps(&self, workflow: &Workflow) -> Vec<String> {
workflow
.steps
.iter()
.filter(|step| {
self.steps
.get(&step.id)
.is_some_and(|r| r.status == StepStatus::Pending)
&& step.depends_on.iter().all(|dep| {
self.steps
.get(dep)
.is_some_and(|r| r.status == StepStatus::Succeeded)
})
})
.map(|s| s.id.clone())
.collect()
}
pub fn all_succeeded(&self) -> bool {
self.steps
.values()
.all(|r| r.status == StepStatus::Succeeded)
}
}
pub mod keys {
pub fn definition(project: &str, name: &str) -> String {
format!("project/{project}/workflows/{name}")
}
pub fn run(project: &str, name: &str, id: &str) -> String {
format!("project/{project}/workflows/{name}/runs/{id}")
}
pub fn runs_prefix(project: &str, name: &str) -> String {
format!("project/{project}/workflows/{name}/runs/")
}
pub fn definitions_prefix(project: &str) -> String {
format!("project/{project}/workflows/")
}
}
#[cfg(test)]
mod tests {
use super::*;
fn step(id: &str, deps: &[&str]) -> Step {
Step {
id: id.to_string(),
function: format!("fn-{id}"),
depends_on: deps.iter().map(std::string::ToString::to_string).collect(),
retry: RetryPolicy::default(),
compensate: None,
}
}
#[test]
fn validate_catches_dupes_missing_deps_and_cycles() {
let ok = Workflow {
name: "w".into(),
steps: vec![step("a", &[]), step("b", &["a"]), step("c", &["b"])],
};
assert!(ok.validate().is_ok());
assert!(Workflow {
name: "w".into(),
steps: vec![]
}
.validate()
.is_err());
assert!(Workflow {
name: "w".into(),
steps: vec![step("a", &[]), step("a", &[])],
}
.validate()
.is_err());
assert!(Workflow {
name: "w".into(),
steps: vec![step("a", &["ghost"])],
}
.validate()
.is_err());
assert!(Workflow {
name: "w".into(),
steps: vec![step("a", &["b"]), step("b", &["a"])],
}
.validate()
.is_err());
}
#[test]
fn ready_steps_respects_the_dag_barrier() {
let wf = Workflow {
name: "w".into(),
steps: vec![
step("root", &[]),
step("x", &["root"]),
step("y", &["root"]),
step("join", &["x", "y"]),
],
};
assert!(wf.validate().is_ok());
let mut run = WorkflowRun::start(&wf, "r1", None, 0);
assert_eq!(run.ready_steps(&wf), vec!["root".to_string()]);
run.steps.get_mut("root").unwrap().status = StepStatus::Succeeded;
let ready: BTreeSet<String> = run.ready_steps(&wf).into_iter().collect();
assert_eq!(ready, BTreeSet::from(["x".to_string(), "y".to_string()]));
run.steps.get_mut("x").unwrap().status = StepStatus::Succeeded;
assert!(!run.ready_steps(&wf).contains(&"join".to_string()));
run.steps.get_mut("y").unwrap().status = StepStatus::Succeeded;
assert_eq!(run.ready_steps(&wf), vec!["join".to_string()]);
run.steps.get_mut("join").unwrap().status = StepStatus::Succeeded;
assert!(run.all_succeeded());
}
#[test]
fn run_serde_round_trips_and_reports_terminal() {
let wf = Workflow {
name: "w".into(),
steps: vec![step("a", &[])],
};
let run = WorkflowRun::start(&wf, "r1", Some("aGk=".into()), 5);
assert!(!run.is_terminal());
let json = serde_json::to_string(&run).unwrap();
let back: WorkflowRun = serde_json::from_str(&json).unwrap();
assert_eq!(back, run);
assert!(json.contains("\"status\":\"running\""));
}
#[test]
fn keyspace_is_stable() {
assert_eq!(
keys::definition("default", "etl"),
"project/default/workflows/etl"
);
assert_eq!(
keys::run("default", "etl", "r1"),
"project/default/workflows/etl/runs/r1"
);
assert_eq!(
keys::runs_prefix("default", "etl"),
"project/default/workflows/etl/runs/"
);
}
}