use std::collections::BTreeMap;
use super::run::{deps_passed, ensure_run_matches_dag};
use super::{
DagValidationError, TaskContextItem, TaskDag, TaskEventError, TaskOperator, TaskRisk, TaskRun,
TaskState, validate_dag,
};
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct TaskContextSlice {
pub dag_id: String,
pub task_id: String,
pub title: String,
pub operator: TaskOperator,
pub risk: TaskRisk,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub run_state: Option<TaskState>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub ready: Option<bool>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub blocked_by: Vec<String>,
pub depends_on: Vec<String>,
pub inputs: Vec<String>,
pub artifact_inputs: Vec<String>,
pub context: Vec<TaskContextItem>,
pub outputs: Vec<String>,
pub required_checks: Vec<String>,
pub tools: Vec<String>,
pub receipts: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TaskContextError {
InvalidDag(DagValidationError),
InvalidRun(TaskEventError),
UnknownTask(String),
}
pub fn task_context_slice(
dag: &TaskDag,
task_id: &str,
available_context: &[TaskContextItem],
) -> Result<TaskContextSlice, TaskContextError> {
validate_dag(dag).map_err(TaskContextError::InvalidDag)?;
let Some(node) = dag.nodes.iter().find(|node| node.id == task_id) else {
return Err(TaskContextError::UnknownTask(task_id.to_string()));
};
let context_by_id = available_context
.iter()
.map(|item| (item.id.as_str(), item))
.collect::<BTreeMap<_, _>>();
let mut context = Vec::new();
let mut artifact_inputs = Vec::new();
for input in &node.inputs {
if let Some(item) = context_by_id.get(input.as_str()) {
context.push((*item).clone());
} else {
artifact_inputs.push(input.clone());
}
}
if node.inputs.is_empty() {
artifact_inputs.extend(dependency_outputs(dag, node));
}
artifact_inputs.sort();
artifact_inputs.dedup();
context.sort_by(|a, b| a.id.cmp(&b.id));
let mut receipts = node.receipts.clone();
receipts.extend(
context
.iter()
.filter_map(|item| (!item.source.trim().is_empty()).then_some(item.source.clone())),
);
receipts.sort();
receipts.dedup();
Ok(TaskContextSlice {
dag_id: dag.id.clone(),
task_id: node.id.clone(),
title: node.title.clone(),
operator: node.operator,
risk: node.risk,
run_state: None,
ready: None,
blocked_by: Vec::new(),
depends_on: node.depends_on.clone(),
inputs: node.inputs.clone(),
artifact_inputs,
context,
outputs: node.outputs.clone(),
required_checks: node.required_checks.clone(),
tools: node.tools.clone(),
receipts,
})
}
fn dependency_outputs(dag: &TaskDag, node: &super::TaskNode) -> Vec<String> {
let nodes = dag
.nodes
.iter()
.map(|node| (node.id.as_str(), node))
.collect::<BTreeMap<_, _>>();
node.depends_on
.iter()
.filter_map(|dependency| nodes.get(dependency.as_str()))
.flat_map(|dependency| dependency.outputs.iter().cloned())
.filter(|output| !output.trim().is_empty())
.collect()
}
pub fn task_context_slice_for_run(
dag: &TaskDag,
run: &TaskRun,
task_id: &str,
available_context: &[TaskContextItem],
) -> Result<TaskContextSlice, TaskContextError> {
let mut slice = task_context_slice(dag, task_id, available_context)?;
ensure_run_matches_dag(dag, run).map_err(TaskContextError::InvalidRun)?;
let Some(state) = run.states.get(task_id).copied() else {
return Err(TaskContextError::UnknownTask(task_id.to_string()));
};
let blocked_by = deps_passed(dag, run, task_id);
let ready = state == TaskState::Ready && blocked_by.is_empty();
slice.run_state = Some(state);
slice.ready = Some(ready);
slice.blocked_by = blocked_by;
Ok(slice)
}