use std::collections::HashMap;
use super::identity::PhaseIdentity;
use super::storage::{Checkpoint, PhaseStatus};
#[derive(Clone, Debug)]
pub enum ResumeAction {
Skip,
CursorResume { cursor_state: serde_json::Value },
ReRun,
IdentityMismatch { reason: String },
}
#[derive(Clone, Debug, Default)]
pub struct ResumePlan {
actions: HashMap<String, ResumeAction>,
pub is_resume: bool,
}
impl ResumePlan {
pub fn fresh() -> Self {
Self::default()
}
pub fn from_checkpoint(
saved: &Checkpoint,
candidates: &[(PhaseIdentity, bool)],
current_params: &HashMap<String, String>,
) -> Self {
let saved_index: HashMap<String, &super::storage::PhaseEntry> = saved
.phases
.iter()
.map(|e| (identity_key(&e.identity), e))
.collect();
let mut actions = HashMap::with_capacity(candidates.len());
for (cand, declared_idempotent) in candidates {
let key = identity_key(cand);
let action = match saved_index.get(&key) {
None => ResumeAction::ReRun,
Some(saved_entry) => {
classify(cand, saved_entry, *declared_idempotent, current_params)
}
};
actions.insert(key, action);
}
Self {
actions,
is_resume: true,
}
}
pub fn action_for(&self, identity: &PhaseIdentity) -> ResumeAction {
let key = identity_key(identity);
self.actions
.get(&key)
.cloned()
.unwrap_or(ResumeAction::ReRun)
}
pub fn skip_count(&self) -> usize {
self.actions
.values()
.filter(|a| matches!(a, ResumeAction::Skip))
.count()
}
pub fn cursor_resume_count(&self) -> usize {
self.actions
.values()
.filter(|a| matches!(a, ResumeAction::CursorResume { .. }))
.count()
}
pub fn mismatch_count(&self) -> usize {
self.actions
.values()
.filter(|a| matches!(a, ResumeAction::IdentityMismatch { .. }))
.count()
}
}
fn classify(
candidate: &PhaseIdentity,
saved: &super::storage::PhaseEntry,
declared_idempotent: bool,
current_params: &HashMap<String, String>,
) -> ResumeAction {
if !declared_idempotent {
return ResumeAction::ReRun;
}
match saved.status {
PhaseStatus::Completed => {
if !candidate.matches_full(&saved.identity) {
return ResumeAction::IdentityMismatch {
reason: format!(
"phase '{}': scope or phase config changed since \
this phase last ran (base hash differs)",
phase_label(&candidate.yaml_path),
),
};
}
if let Err(mismatch) =
saved_params_still_valid(saved.params_consumed.as_deref(), current_params)
{
return ResumeAction::IdentityMismatch {
reason: format!(
"phase '{}': {mismatch} since this phase last ran",
phase_label(&candidate.yaml_path),
),
};
}
if !saved.skip_eligible {
return ResumeAction::ReRun;
}
ResumeAction::Skip
}
PhaseStatus::Running => match &saved.cursor_state {
Some(cs) => ResumeAction::CursorResume {
cursor_state: cs.clone(),
},
None => ResumeAction::ReRun,
},
PhaseStatus::Pending | PhaseStatus::Failed => ResumeAction::ReRun,
}
}
fn saved_params_still_valid(
stored_json: Option<&str>,
current_params: &HashMap<String, String>,
) -> Result<(), String> {
let Some(json) = stored_json else {
return Err("saved entry carries no consumed-params record".into());
};
let Ok(stored) = serde_json::from_str::<std::collections::BTreeMap<String, String>>(json)
else {
return Err("saved consumed-params record is unreadable".into());
};
for (name, stored_digest) in stored {
let current = current_params
.get(&name)
.map(|v| crate::checkpoint::params_scope::value_digest(v));
if current.as_deref() != Some(stored_digest.as_str()) {
return Err(format!("param '{name}' changed"));
}
}
Ok(())
}
fn identity_key(identity: &PhaseIdentity) -> String {
let path_json = serde_json::to_string(&identity.yaml_path).unwrap_or_else(|_| String::new());
format!("{path_json}\x1f{}", identity.coords)
}
fn phase_label(yaml_path: &[super::identity::PathSegment]) -> String {
use super::identity::PathSegment;
yaml_path
.iter()
.filter_map(|seg| match seg {
PathSegment::Phase(n) => Some(n.clone()),
_ => None,
})
.next_back()
.unwrap_or_else(|| "<unknown>".to_string())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::checkpoint::{
PathSegment,
storage::{OpCounts, PhaseEntry},
};
fn ident_with_hash(name: &str, coords: &str, hash: Option<[u8; 32]>) -> PhaseIdentity {
PhaseIdentity {
yaml_path: vec![
PathSegment::Scenario("s".into()),
PathSegment::Phase(name.into()),
],
coords: coords.into(),
phase_hash: hash,
}
}
fn entry(identity: PhaseIdentity, status: PhaseStatus, skip_eligible: bool) -> PhaseEntry {
PhaseEntry {
identity,
skip_eligible,
params_consumed: Some("{}".into()),
status,
duration_secs: Some(1.0),
op_counts: Some(OpCounts::default()),
cursor_state: None,
error: None,
}
}
fn checkpoint_with(phases: Vec<PhaseEntry>) -> Checkpoint {
Checkpoint {
version: 1,
session: "s".into(),
started_at: "t".into(),
checkpoint_at: "t".into(),
invocation: 1,
phases,
}
}
#[test]
fn fresh_plan_rerun_for_everything() {
let plan = ResumePlan::fresh();
assert!(matches!(
plan.action_for(&ident_with_hash("p", "", None)),
ResumeAction::ReRun
));
assert!(!plan.is_resume);
}
#[test]
fn consumed_param_change_invalidates_and_names_the_param() {
let h = [0xab; 32];
let id = ident_with_hash("load", "", Some(h));
let mut e = entry(id.clone(), PhaseStatus::Completed, true);
e.params_consumed = Some(format!(
r#"{{"dataset":"{}"}}"#,
crate::checkpoint::params_scope::value_digest("example"),
));
let saved = checkpoint_with(vec![e]);
let mut params = HashMap::new();
params.insert("dataset".to_string(), "sift10m".to_string());
let plan = ResumePlan::from_checkpoint(&saved, &[(id.clone(), true)], ¶ms);
match plan.action_for(&id) {
ResumeAction::IdentityMismatch { reason } => {
assert!(
reason.contains("param 'dataset' changed"),
"reason must name the param: {reason}"
);
}
other => panic!("expected IdentityMismatch, got {other:?}"),
}
}
#[test]
fn unrelated_param_change_still_skips() {
let h = [0xab; 32];
let id = ident_with_hash("load", "", Some(h));
let mut e = entry(id.clone(), PhaseStatus::Completed, true);
e.params_consumed = Some(format!(
r#"{{"dataset":"{}"}}"#,
crate::checkpoint::params_scope::value_digest("example"),
));
let saved = checkpoint_with(vec![e]);
let mut params = HashMap::new();
params.insert("dataset".to_string(), "example".to_string());
params.insert("suite_k".to_string(), "100".to_string());
let plan = ResumePlan::from_checkpoint(&saved, &[(id.clone(), true)], ¶ms);
assert!(
matches!(plan.action_for(&id), ResumeAction::Skip),
"an unconsumed param's change must not invalidate"
);
}
#[test]
fn completed_idempotent_with_matching_hash_skips() {
let h = [0xab; 32];
let id = ident_with_hash("schema", "", Some(h));
let saved = checkpoint_with(vec![entry(id.clone(), PhaseStatus::Completed, true)]);
let plan = ResumePlan::from_checkpoint(&saved, &[(id.clone(), true)], &HashMap::new());
assert!(matches!(plan.action_for(&id), ResumeAction::Skip));
assert_eq!(plan.skip_count(), 1);
}
#[test]
fn completed_with_hash_mismatch_invalidates() {
let h_old = [0x01; 32];
let h_new = [0x02; 32];
let id_saved = ident_with_hash("schema", "", Some(h_old));
let id_now = ident_with_hash("schema", "", Some(h_new));
let saved = checkpoint_with(vec![entry(id_saved.clone(), PhaseStatus::Completed, true)]);
let plan = ResumePlan::from_checkpoint(&saved, &[(id_now.clone(), true)], &HashMap::new());
match plan.action_for(&id_now) {
ResumeAction::IdentityMismatch { reason } => {
assert!(reason.contains("schema"), "reason: {reason}");
}
other => panic!("expected IdentityMismatch, got {other:?}"),
}
assert_eq!(plan.mismatch_count(), 1);
}
#[test]
fn declared_none_always_reruns_even_if_completed() {
let id = ident_with_hash("schema", "", None);
let saved = checkpoint_with(vec![entry(id.clone(), PhaseStatus::Completed, false)]);
let plan = ResumePlan::from_checkpoint(&saved, &[(id.clone(), false)], &HashMap::new());
assert!(matches!(plan.action_for(&id), ResumeAction::ReRun));
}
#[test]
fn running_with_cursor_state_yields_cursor_resume() {
let id = ident_with_hash("rampup", "", None);
let mut e = entry(id.clone(), PhaseStatus::Running, true);
e.cursor_state = Some(serde_json::json!({"next_cycle": 12345}));
let saved = checkpoint_with(vec![e]);
let plan = ResumePlan::from_checkpoint(&saved, &[(id.clone(), true)], &HashMap::new());
match plan.action_for(&id) {
ResumeAction::CursorResume { cursor_state } => {
assert_eq!(cursor_state["next_cycle"], 12345);
}
other => panic!("expected CursorResume, got {other:?}"),
}
assert_eq!(plan.cursor_resume_count(), 1);
}
#[test]
fn running_without_cursor_state_reruns() {
let id = ident_with_hash("rampup", "", None);
let saved = checkpoint_with(vec![entry(id.clone(), PhaseStatus::Running, true)]);
let plan = ResumePlan::from_checkpoint(&saved, &[(id.clone(), true)], &HashMap::new());
assert!(matches!(plan.action_for(&id), ResumeAction::ReRun));
}
#[test]
fn unknown_candidate_reruns() {
let saved = checkpoint_with(vec![]);
let id = ident_with_hash("brand_new", "", None);
let plan = ResumePlan::from_checkpoint(&saved, &[(id.clone(), true)], &HashMap::new());
assert!(matches!(plan.action_for(&id), ResumeAction::ReRun));
}
#[test]
fn failed_phase_reruns() {
let id = ident_with_hash("flaky", "", None);
let mut e = entry(id.clone(), PhaseStatus::Failed, true);
e.error = Some("boom".into());
let saved = checkpoint_with(vec![e]);
let plan = ResumePlan::from_checkpoint(&saved, &[(id.clone(), true)], &HashMap::new());
assert!(matches!(plan.action_for(&id), ResumeAction::ReRun));
}
}