use oxo_flow_core::config::WorkflowConfig;
use oxo_flow_core::dag::WorkflowDag;
use oxo_flow_core::executor::checkpoint::CheckpointState;
use oxo_flow_core::rule::Rule;
use std::collections::{HashMap, HashSet};
use std::path::Path;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RuleStatus {
NeverCompleted,
ConfigInvalidated,
InputInvalidated,
OutputsMissing,
Cascaded { from: String },
Skipped,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PreviewRule {
pub name: String,
pub status: RuleStatus,
}
#[derive(Debug, Clone)]
pub struct RunPreview {
pub checkpoint_path: std::path::PathBuf,
pub checkpoint_modified: Option<std::time::SystemTime>,
pub completed_total: usize,
pub plan: Vec<PreviewRule>,
pub will_skip: usize,
pub protected_outside: usize,
pub cascade_chains: Vec<Vec<String>>,
}
#[allow(clippy::too_many_arguments)]
pub fn preview_run_plan(
ck: &CheckpointState,
config: &WorkflowConfig,
dag: &WorkflowDag,
order: &[String],
workdir: &Path,
wildcard_values: &HashMap<String, String>,
sensitive_keys: &HashSet<String>,
interpreter_map: &HashMap<String, String>,
checkpoint_path: &Path,
) -> RunPreview {
let completed_original: HashSet<String> = ck.completed_rules.clone();
let mut clone = ck.clone();
let config_report = oxo_flow_core::config_impact::detect_config_changes(
&mut clone,
&config.rules,
dag,
&config.config,
sensitive_keys,
interpreter_map,
);
let config_invalidated: HashSet<String> = config_report.invalidated.iter().cloned().collect();
let (manifest_invalidated, _baselined) =
detect_input_manifest_invalidations(&mut clone, config, order, workdir, wildcard_values);
let seeds: HashSet<String> = manifest_invalidated.iter().cloned().collect();
invalidate_with_downstream(&mut clone, dag, &seeds);
let seeds_for_cascade: HashSet<String> = config_invalidated
.union(&manifest_invalidated)
.cloned()
.collect();
let order_set: HashSet<&str> = order.iter().map(String::as_str).collect();
let mut plan = Vec::with_capacity(order.len());
let mut will_skip = 0usize;
for name in order {
let status = if !completed_original.contains(name) {
RuleStatus::NeverCompleted
} else if config_invalidated.contains(name) {
RuleStatus::ConfigInvalidated
} else if manifest_invalidated.contains(name) {
RuleStatus::InputInvalidated
} else if let Some(rule) = config.get_rule(name)
&& !rule_outputs_exist(rule, workdir, wildcard_values)
{
RuleStatus::OutputsMissing
} else if !clone.completed_rules.contains(name) {
RuleStatus::Cascaded {
from: nearest_seed(name, dag, &seeds_for_cascade),
}
} else {
RuleStatus::Skipped
};
if status == RuleStatus::Skipped {
will_skip += 1;
}
plan.push(PreviewRule {
name: name.clone(),
status,
});
}
let mut cascade_chains: Vec<Vec<String>> = Vec::new();
for seed in seeds_for_cascade.iter() {
if !completed_original.contains(seed) {
continue; }
let chain = cascade_chain(seed, dag, order_set.clone());
if chain.len() > 1 {
cascade_chains.push(chain);
}
}
cascade_chains.sort_by(|a, b| a.first().cmp(&b.first()));
let protected_outside = completed_original
.iter()
.filter(|name| !order_set.contains(name.as_str()))
.count();
RunPreview {
checkpoint_path: checkpoint_path.to_path_buf(),
checkpoint_modified: std::fs::metadata(checkpoint_path)
.and_then(|m| m.modified())
.ok(),
completed_total: completed_original.len(),
plan,
will_skip,
protected_outside,
cascade_chains,
}
}
fn cascade_chain(seed: &str, dag: &WorkflowDag, order_set: HashSet<&str>) -> Vec<String> {
let mut chain = vec![seed.to_string()];
let mut frontier = vec![seed.to_string()];
let mut visited: HashSet<String> = frontier.iter().cloned().collect();
while let Some(name) = frontier.pop() {
let Ok(dependents) = dag.dependents(&name) else {
continue;
};
let mut next: Vec<String> = dependents
.into_iter()
.filter(|d| order_set.contains(d.as_str()))
.filter(|d| visited.insert(d.clone()))
.collect();
next.sort();
chain.extend(next.iter().cloned());
frontier.extend(next);
}
chain
}
fn nearest_seed(rule: &str, dag: &WorkflowDag, seeds: &HashSet<String>) -> String {
if seeds.contains(rule) {
return rule.to_string();
}
let mut sorted_seeds: Vec<&String> = seeds.iter().collect();
sorted_seeds.sort();
for seed in sorted_seeds {
if let Some(path) = dag_path_exists(seed, rule, dag) {
let _ = path;
return seed.clone();
}
}
"<unknown>".to_string()
}
fn dag_path_exists(from: &str, to: &str, dag: &WorkflowDag) -> Option<usize> {
let mut frontier = vec![(from.to_string(), 0)];
let mut visited: HashSet<String> = frontier.iter().map(|(n, _)| n.clone()).collect();
while let Some((name, depth)) = frontier.pop() {
if name == to {
return Some(depth);
}
let Ok(dependents) = dag.dependents(&name) else {
continue;
};
for dependent in dependents {
if visited.insert(dependent.clone()) {
frontier.push((dependent, depth + 1));
}
}
}
None
}
pub fn rule_outputs_exist(
rule: &Rule,
workdir: &Path,
wildcard_values: &HashMap<String, String>,
) -> bool {
rule.output.iter().all(|output| {
let expanded =
oxo_flow_core::executor::checkpoint::expand_config_in_path(output, wildcard_values);
expanded.contains('{') || workdir.join(&expanded).exists()
})
}
pub fn detect_input_manifest_invalidations(
ck: &mut CheckpointState,
config: &WorkflowConfig,
order: &[String],
workdir: &Path,
wildcard_values: &HashMap<String, String>,
) -> (HashSet<String>, usize) {
let mut mismatched: HashSet<String> = HashSet::new();
let mut baselined = 0usize;
for name in order {
if !ck.completed_rules.contains(name) {
continue;
}
let Some(rule) = config.get_rule(name) else {
continue;
};
match oxo_flow_core::executor::checkpoint::snapshot_input_manifest(
rule,
workdir,
wildcard_values,
) {
Ok(Some(current)) => match ck.input_manifests.get(name) {
Some(recorded) if *recorded == current => {}
Some(_) => {
mismatched.insert(name.clone());
}
None => {
ck.record_input_manifest(name, current);
baselined += 1;
}
},
Ok(None) => {}
Err(_) => {
mismatched.insert(name.clone());
}
}
}
(mismatched, baselined)
}
pub(crate) fn invalidate_with_downstream(
ck: &mut CheckpointState,
dag: &WorkflowDag,
seeds: &HashSet<String>,
) -> Vec<String> {
let mut invalidated: HashSet<String> = seeds.clone();
let mut frontier: Vec<String> = seeds.iter().cloned().collect();
while let Some(name) = frontier.pop() {
if let Ok(dependents) = dag.dependents(&name) {
for dependent in dependents {
if invalidated.insert(dependent.clone()) {
frontier.push(dependent);
}
}
}
}
for name in &invalidated {
ck.completed_rules.remove(name);
}
let mut names: Vec<String> = invalidated.into_iter().collect();
names.sort();
names
}
#[cfg(test)]
mod tests {
use super::*;
use oxo_flow_core::config::WorkflowConfig;
use oxo_flow_core::executor::checkpoint::snapshot_input_manifest;
fn fixture(
dir: &std::path::Path,
) -> (
WorkflowConfig,
WorkflowDag,
Vec<String>,
HashMap<String, String>,
) {
let toml = r#"
[workflow]
name = "t"
version = "1.0"
[config]
ref = "ref.fa"
[[sample_groups]]
name = "cohort"
samples = ["S1", "S2"]
[[rules]]
name = "trim"
input = ["raw/{sample}.fq"]
output = ["trimmed/{sample}.fq"]
shell = "cp {input[0]} {output[0]}"
[[rules]]
name = "align"
input = ["trimmed/{sample}.fq"]
output = ["aligned/{sample}.bam"]
depends_on = ["trim"]
shell = "cp {input[0]} {output[0]}"
[[rules]]
name = "combine"
input = ["aligned/*.bam"]
output = ["combined.bam"]
depends_on = ["align"]
shell = "cat {input[0]} > {output[0]} && echo {config.ref}"
"#;
for f in [
"raw/S1.fq",
"raw/S2.fq",
"trimmed/S1.fq",
"trimmed/S2.fq",
"aligned/S1.bam",
"aligned/S2.bam",
"combined.bam",
] {
let p = dir.join(f);
std::fs::create_dir_all(p.parent().unwrap()).unwrap();
std::fs::write(&p, "data1").unwrap();
}
let mut config = WorkflowConfig::parse(toml).unwrap();
config.apply_defaults();
config.expand_wildcards().unwrap();
let dag = WorkflowDag::from_rules(&config.rules).unwrap();
let order = dag.execution_order().unwrap();
let mut wildcard_values = HashMap::new();
for (key, value) in &config.config {
wildcard_values.insert(
format!("config.{key}"),
value.as_str().unwrap_or_default().to_string(),
);
}
(config, dag, order, wildcard_values)
}
fn completed_checkpoint(
config: &WorkflowConfig,
order: &[String],
dir: &std::path::Path,
wildcard_values: &HashMap<String, String>,
) -> CheckpointState {
let mut ck = CheckpointState::new();
for name in order {
ck.completed_rules.insert(name.clone());
if let Some(rule) = config.get_rule(name)
&& let Ok(Some(manifest)) = snapshot_input_manifest(rule, dir, wildcard_values)
{
ck.record_input_manifest(name, manifest);
}
}
ck
}
fn sensitive() -> HashSet<String> {
HashSet::new()
}
fn run_preview(
ck: &CheckpointState,
config: &WorkflowConfig,
dag: &WorkflowDag,
order: &[String],
dir: &std::path::Path,
wildcard_values: &HashMap<String, String>,
) -> RunPreview {
preview_run_plan(
ck,
config,
dag,
order,
dir,
wildcard_values,
&sensitive(),
&config.workflow.interpreter_map,
&dir.join(".oxo-flow/checkpoint.json"),
)
}
fn status_of<'a>(preview: &'a RunPreview, name: &str) -> &'a RuleStatus {
&preview
.plan
.iter()
.find(|r| r.name == name)
.unwrap_or_else(|| panic!("{name} missing from plan"))
.status
}
#[test]
fn empty_checkpoint_means_everything_runs() {
let dir = tempfile::tempdir().unwrap();
let (config, dag, order, wildcard_values) = fixture(dir.path());
let ck = CheckpointState::new();
let preview = run_preview(&ck, &config, &dag, &order, dir.path(), &wildcard_values);
assert_eq!(preview.plan.len(), 5);
assert!(
preview
.plan
.iter()
.all(|r| r.status == RuleStatus::NeverCompleted)
);
assert_eq!(preview.will_skip, 0);
assert_eq!(preview.protected_outside, 0);
assert!(preview.cascade_chains.is_empty());
}
#[test]
fn up_to_date_workflow_skips_everything() {
let dir = tempfile::tempdir().unwrap();
let (config, dag, order, wildcard_values) = fixture(dir.path());
let ck = completed_checkpoint(&config, &order, dir.path(), &wildcard_values);
let preview = run_preview(&ck, &config, &dag, &order, dir.path(), &wildcard_values);
assert!(preview.plan.iter().all(|r| r.status == RuleStatus::Skipped));
assert_eq!(preview.will_skip, 5);
assert_eq!(preview.completed_total, 5);
assert!(preview.cascade_chains.is_empty());
}
#[test]
fn changed_input_invalidates_rule_and_cascades_downstream() {
let dir = tempfile::tempdir().unwrap();
let (config, dag, order, wildcard_values) = fixture(dir.path());
let ck = completed_checkpoint(&config, &order, dir.path(), &wildcard_values);
std::fs::write(dir.path().join("raw/S1.fq"), "completely new content").unwrap();
let preview = run_preview(&ck, &config, &dag, &order, dir.path(), &wildcard_values);
assert_eq!(
status_of(&preview, "trim_cohort_S1"),
&RuleStatus::InputInvalidated
);
assert_eq!(
status_of(&preview, "align_cohort_S1"),
&RuleStatus::Cascaded {
from: "trim_cohort_S1".to_string()
}
);
assert_eq!(
status_of(&preview, "align_cohort_S2"),
&RuleStatus::Cascaded {
from: "trim_cohort_S1".to_string()
}
);
assert_eq!(
status_of(&preview, "combine"),
&RuleStatus::Cascaded {
from: "trim_cohort_S1".to_string()
}
);
assert_eq!(status_of(&preview, "trim_cohort_S2"), &RuleStatus::Skipped);
assert_eq!(preview.will_skip, 1);
assert!(preview.cascade_chains.iter().any(|c| c
== &vec![
"trim_cohort_S1".to_string(),
"align_cohort_S1".to_string(),
"align_cohort_S2".to_string(),
"combine".to_string()
]));
}
#[test]
fn missing_output_marks_rerun_but_leaves_manifest_intact() {
let dir = tempfile::tempdir().unwrap();
let (config, dag, order, wildcard_values) = fixture(dir.path());
let ck = completed_checkpoint(&config, &order, dir.path(), &wildcard_values);
std::fs::remove_file(dir.path().join("combined.bam")).unwrap();
let preview = run_preview(&ck, &config, &dag, &order, dir.path(), &wildcard_values);
assert_eq!(status_of(&preview, "combine"), &RuleStatus::OutputsMissing);
assert_eq!(status_of(&preview, "trim_cohort_S1"), &RuleStatus::Skipped);
}
#[test]
fn target_subset_counts_protected_outside() {
let dir = tempfile::tempdir().unwrap();
let (config, dag, order, wildcard_values) = fixture(dir.path());
let ck = completed_checkpoint(&config, &order, dir.path(), &wildcard_values);
let targets = ["trim_cohort_S1".to_string()];
let subset: Vec<String> = dag
.execution_order_for_targets(&targets.iter().map(String::as_str).collect::<Vec<_>>())
.unwrap();
let preview = run_preview(&ck, &config, &dag, &subset, dir.path(), &wildcard_values);
assert_eq!(preview.plan.len(), 1);
assert_eq!(preview.protected_outside, 4);
assert_eq!(preview.will_skip, 1);
}
#[test]
fn config_change_invalidates_referencing_rules() {
let dir = tempfile::tempdir().unwrap();
let (config, dag, order, wildcard_values) = fixture(dir.path());
let mut ck = completed_checkpoint(&config, &order, dir.path(), &wildcard_values);
oxo_flow_core::config_impact::detect_config_changes(
&mut ck,
&config.rules,
&dag,
&config.config,
&sensitive(),
&config.workflow.interpreter_map,
);
let mut changed_config = config.clone();
changed_config
.config
.insert("ref".to_string(), toml::Value::String("other.fa".into()));
let mut changed_wildcards = wildcard_values.clone();
changed_wildcards.insert("config.ref".to_string(), "other.fa".into());
let preview = run_preview(
&ck,
&changed_config,
&dag,
&order,
dir.path(),
&changed_wildcards,
);
assert_eq!(
status_of(&preview, "combine"),
&RuleStatus::ConfigInvalidated
);
assert_eq!(status_of(&preview, "trim_cohort_S1"), &RuleStatus::Skipped);
}
#[test]
fn legacy_checkpoint_adopts_baseline_without_invalidating() {
let dir = tempfile::tempdir().unwrap();
let (config, dag, order, wildcard_values) = fixture(dir.path());
let mut ck = completed_checkpoint(&config, &order, dir.path(), &wildcard_values);
ck.input_manifests.clear();
ck.config_snapshot.clear();
ck.rule_fingerprints.clear();
let preview = run_preview(&ck, &config, &dag, &order, dir.path(), &wildcard_values);
assert!(preview.plan.iter().all(|r| r.status == RuleStatus::Skipped));
}
}