use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use uuid::Uuid;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum NodeStatus {
Pending,
Running,
Done,
Failed,
Yielded,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DAGNode {
pub id: Uuid,
pub step_type: String,
pub depends_on: Vec<Uuid>,
pub status: NodeStatus,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DAG {
pub id: Uuid,
pub nodes: Vec<DAGNode>,
pub created_at: DateTime<Utc>,
}
impl DAG {
pub fn new() -> Self {
Self {
id: Uuid::new_v4(),
nodes: Vec::new(),
created_at: Utc::now(),
}
}
pub fn add_node(&mut self, step_type: impl Into<String>, depends_on: Vec<Uuid>) -> Uuid {
let id = Uuid::new_v4();
self.nodes.push(DAGNode {
id,
step_type: step_type.into(),
depends_on,
status: NodeStatus::Pending,
});
id
}
pub fn restore_statuses(&mut self, checkpointed: &DAG) {
let status_map: std::collections::HashMap<Uuid, NodeStatus> =
checkpointed.nodes.iter().map(|n| (n.id, n.status.clone())).collect();
for node in &mut self.nodes {
if let Some(status) = status_map.get(&node.id) {
node.status = status.clone();
}
}
}
}
impl Default for DAG {
fn default() -> Self {
Self::new()
}
}
pub struct DAGBuilder {
entries: Vec<(String, Vec<String>)>,
}
impl DAGBuilder {
pub fn node(mut self, step_type: &str, depends_on: &[&str]) -> Self {
self.entries.push((
step_type.to_string(),
depends_on.iter().map(|s| s.to_string()).collect(),
));
self
}
pub fn build(self) -> Result<DAG, String> {
let mut dag = DAG::new();
let mut name_to_id: std::collections::HashMap<String, Uuid> = std::collections::HashMap::new();
for (step_type, _) in &self.entries {
let id = Uuid::new_v4();
name_to_id.insert(step_type.clone(), id);
}
for (step_type, deps) in &self.entries {
let mut resolved_deps = Vec::new();
for dep in deps {
match name_to_id.get(dep) {
Some(id) => resolved_deps.push(*id),
None => return Err(format!("Dependency '{}' not found for node '{}'", dep, step_type)),
}
}
let id = name_to_id[step_type];
dag.nodes.push(DAGNode {
id,
step_type: step_type.clone(),
depends_on: resolved_deps,
status: NodeStatus::Pending,
});
}
let n = dag.nodes.len();
let _id_to_idx: std::collections::HashMap<Uuid, usize> =
dag.nodes.iter().enumerate().map(|(i, node)| (node.id, i)).collect();
let mut in_degree = vec![0usize; n];
for (i, node) in dag.nodes.iter().enumerate() {
in_degree[i] = node.depends_on.len();
}
let mut queue: std::collections::VecDeque<usize> = in_degree
.iter()
.enumerate()
.filter(|(_, &d)| d == 0)
.map(|(i, _)| i)
.collect();
let mut visited = 0usize;
while let Some(idx) = queue.pop_front() {
visited += 1;
let node_id = dag.nodes[idx].id;
for (i, node) in dag.nodes.iter().enumerate() {
if node.depends_on.contains(&node_id) {
in_degree[i] -= 1;
if in_degree[i] == 0 {
queue.push_back(i);
}
}
}
}
if visited != n {
return Err("Cycle detected in DAG".to_string());
}
Ok(dag)
}
}
impl DAG {
pub fn builder() -> DAGBuilder {
DAGBuilder { entries: Vec::new() }
}
}