use super::*;
#[test]
fn checkpoint_time_excluded_from_wall_ms_accum() {
let world_size = 2usize;
let cfg = cfg_sync_cpu(world_size);
let mut coord = ClusterCoordinator::for_test(cfg);
coord.set_wall_ms_accum_for_test(0, 100.0);
coord.set_wall_ms_accum_for_test(1, 50.0);
coord.handle_checkpoint_result(0, 7, 30.0, None);
assert!(
(coord.wall_ms_accum_for_test(0) - 70.0).abs() < 1e-9,
"wall_ms_accum[0] = {} (expected 70.0)",
coord.wall_ms_accum_for_test(0),
);
assert!(
(coord.wall_ms_accum_for_test(1) - 50.0).abs() < 1e-9,
"wall_ms_accum[1] = {} (expected 50.0 untouched)",
coord.wall_ms_accum_for_test(1),
);
assert_eq!(coord.last_checkpoint_elapsed_ms_ewma(), Some(30.0));
assert_eq!(coord.checkpoint_role(), 0);
assert_eq!(coord.checkpoint_tried_count(7), 0);
}
#[test]
fn checkpoint_failure_records_tried_and_failovers_role() {
let world_size = 3usize;
let cfg = cfg_sync_cpu(world_size);
let mut coord = ClusterCoordinator::for_test(cfg);
assert_eq!(coord.checkpoint_role(), 0);
coord.handle_checkpoint_result(
0, 5, 12.0, Some("disk full".into()),
);
assert_eq!(coord.checkpoint_role(), 1, "role should fail over to rank 1");
assert_eq!(coord.checkpoint_tried_count(5), 1);
coord.handle_checkpoint_result(
1, 5, 8.0, Some("io error".into()),
);
assert_eq!(coord.checkpoint_role(), 2);
assert_eq!(coord.checkpoint_tried_count(5), 2);
coord.handle_checkpoint_result(
2, 5, 5.0, Some("permission denied".into()),
);
assert_eq!(
coord.checkpoint_tried_count(5),
0,
"exhaustion should clear tried_ranks[version]"
);
}
#[test]
fn checkpoint_success_after_failure_clears_tried() {
let world_size = 3usize;
let cfg = cfg_sync_cpu(world_size);
let mut coord = ClusterCoordinator::for_test(cfg);
coord.handle_checkpoint_result(0, 4, 10.0, Some("oom".into()));
assert_eq!(coord.checkpoint_tried_count(4), 1);
assert_eq!(coord.checkpoint_role(), 1);
coord.handle_checkpoint_result(1, 4, 7.0, None);
assert_eq!(coord.checkpoint_tried_count(4), 0);
assert_eq!(coord.checkpoint_role(), 1);
assert_eq!(coord.last_checkpoint_elapsed_ms_ewma(), Some(7.0));
}
#[test]
fn checkpoint_ewma_blends_successive_successes() {
let world_size = 2usize;
let cfg = cfg_sync_cpu(world_size);
let mut coord = ClusterCoordinator::for_test(cfg);
coord.handle_checkpoint_result(0, 1, 100.0, None);
assert_eq!(coord.last_checkpoint_elapsed_ms_ewma(), Some(100.0));
coord.handle_checkpoint_result(0, 2, 50.0, None);
let ewma = coord.last_checkpoint_elapsed_ms_ewma().unwrap();
assert!(
(ewma - 85.0).abs() < 1e-9,
"EWMA after 100 then 50 (alpha=0.3): got {ewma}, expected 85.0"
);
}
#[test]
fn eval_time_excluded_from_wall_ms_accum() {
let world_size = 2usize;
let cfg = cfg_sync_cpu(world_size);
let mut coord = ClusterCoordinator::for_test(cfg);
coord.set_wall_ms_accum_for_test(0, 100.0);
coord.set_wall_ms_accum_for_test(1, 50.0);
coord.handle_eval_result(0, 3, 0.42, 30.0, None);
assert!(
(coord.wall_ms_accum_for_test(0) - 70.0).abs() < 1e-9,
"wall_ms_accum[0] = {} (expected 70.0)",
coord.wall_ms_accum_for_test(0),
);
assert!(
(coord.wall_ms_accum_for_test(1) - 50.0).abs() < 1e-9,
"wall_ms_accum[1] = {} (expected 50.0 untouched)",
coord.wall_ms_accum_for_test(1),
);
assert_eq!(coord.last_eval_elapsed_ms_ewma(), Some(30.0));
}
#[test]
fn eval_ewma_blends_successive_results() {
let world_size = 2usize;
let cfg = cfg_sync_cpu(world_size);
let mut coord = ClusterCoordinator::for_test(cfg);
coord.handle_eval_result(0, 1, 0.1, 100.0, None);
assert_eq!(coord.last_eval_elapsed_ms_ewma(), Some(100.0));
coord.handle_eval_result(0, 2, 0.2, 50.0, None);
let ewma = coord.last_eval_elapsed_ms_ewma().unwrap();
assert!(
(ewma - 85.0).abs() < 1e-9,
"EWMA after 100 then 50 (alpha=0.3): got {ewma}, expected 85.0"
);
}
#[test]
fn eval_error_still_excludes_time_and_updates_ewma() {
let world_size = 2usize;
let cfg = cfg_sync_cpu(world_size);
let mut coord = ClusterCoordinator::for_test(cfg);
coord.set_wall_ms_accum_for_test(0, 100.0);
coord.handle_eval_result(0, 7, f64::NAN, 30.0, Some("boom".into()));
assert!(
(coord.wall_ms_accum_for_test(0) - 70.0).abs() < 1e-9,
"wall_ms_accum[0] = {} (expected 70.0 even on error)",
coord.wall_ms_accum_for_test(0),
);
assert_eq!(coord.last_eval_elapsed_ms_ewma(), Some(30.0));
}
#[test]
fn epoch_fn_time_excluded_from_wall_ms_accum() {
let world_size = 2usize;
let cfg = cfg_sync_cpu(world_size);
let mut coord = ClusterCoordinator::for_test(cfg);
coord.set_wall_ms_accum_for_test(0, 100.0);
coord.set_wall_ms_accum_for_test(1, 50.0);
coord.handle_epoch_fn_elapsed(0, 20.0);
assert!(
(coord.wall_ms_accum_for_test(0) - 80.0).abs() < 1e-9,
"wall_ms_accum[0] = {} (expected 80.0)",
coord.wall_ms_accum_for_test(0),
);
assert!(
(coord.wall_ms_accum_for_test(1) - 50.0).abs() < 1e-9,
"wall_ms_accum[1] = {} (expected 50.0 untouched)",
coord.wall_ms_accum_for_test(1),
);
assert_eq!(coord.last_epoch_fn_elapsed_ms_ewma(), Some(20.0));
}
#[test]
fn epoch_fn_ewma_blends_successive_reports() {
let world_size = 2usize;
let cfg = cfg_sync_cpu(world_size);
let mut coord = ClusterCoordinator::for_test(cfg);
coord.handle_epoch_fn_elapsed(0, 100.0);
assert_eq!(coord.last_epoch_fn_elapsed_ms_ewma(), Some(100.0));
coord.handle_epoch_fn_elapsed(0, 50.0);
let ewma = coord.last_epoch_fn_elapsed_ms_ewma().unwrap();
assert!(
(ewma - 85.0).abs() < 1e-9,
"EWMA after 100 then 50 (alpha=0.3): got {ewma}, expected 85.0"
);
}
fn build_coord_for_slack(remaining_batches: usize) -> ClusterCoordinator {
let world_size = 2;
let cfg = ClusterCoordinatorConfig::new(
ApplyPolicy::Sync,
AverageBackend::Cpu,
world_size,
ElChe::new(world_size, 10),
)
.no_divergence_guard();
let mut coord = ClusterCoordinator::for_test(cfg);
coord.el_che_mut_for_test().report_timing(
&[500.0, 1000.0],
&[10, 10],
10.0,
);
coord.set_callback_roles_for_test(0, 0, 0);
coord.install_chunk_pool_for_test(0, remaining_batches);
coord.set_rank_epoch_for_test(0, 0);
coord.set_rank_epoch_for_test(1, 0);
coord
}
#[test]
fn callback_slack_stages_on_firing_rank_for_last_cycle() {
let mut coord = build_coord_for_slack(25);
coord.handle_epoch_fn_elapsed(0, 1000.0);
coord.maybe_apply_callback_slack_for_test();
let slack = coord.el_che_for_test().pending_callback_slack_ms();
assert!(
(slack[0] - 1000.0).abs() < 1e-9,
"rank 0 should have staged epoch_fn slack of 1000ms; got {slack:?}",
);
assert_eq!(slack[1], 0.0, "non-firing rank slack must stay zero");
}
#[test]
fn callback_slack_skips_when_not_last_cycle() {
let mut coord = build_coord_for_slack(100);
coord.handle_epoch_fn_elapsed(0, 1000.0);
coord.maybe_apply_callback_slack_for_test();
let slack = coord.el_che_for_test().pending_callback_slack_ms();
assert_eq!(
slack,
&[0.0, 0.0],
"slack must not stage when next cycle is not the last",
);
}
#[test]
fn callback_slack_guard_filters_sub_threshold() {
let mut coord = build_coord_for_slack(25);
coord.handle_epoch_fn_elapsed(0, 50.0);
coord.maybe_apply_callback_slack_for_test();
let slack = coord.el_che_for_test().pending_callback_slack_ms();
assert_eq!(
slack,
&[0.0, 0.0],
"sub-threshold slack must be filtered out (50ms < max(50, 100))",
);
}
#[test]
fn callback_slack_skips_when_pool_empty() {
let mut coord = build_coord_for_slack(0);
coord.handle_epoch_fn_elapsed(0, 1000.0);
coord.maybe_apply_callback_slack_for_test();
let slack = coord.el_che_for_test().pending_callback_slack_ms();
assert_eq!(slack, &[0.0, 0.0]);
}
#[test]
fn callback_slack_skips_when_elche_uncalibrated() {
let world_size = 2;
let cfg = ClusterCoordinatorConfig::new(
ApplyPolicy::Sync,
AverageBackend::Cpu,
world_size,
ElChe::new(world_size, 10),
)
.no_divergence_guard();
let mut coord = ClusterCoordinator::for_test(cfg);
coord.set_callback_roles_for_test(0, 0, 0);
coord.install_chunk_pool_for_test(0, 5);
coord.handle_epoch_fn_elapsed(0, 1000.0);
coord.maybe_apply_callback_slack_for_test();
let slack = coord.el_che_for_test().pending_callback_slack_ms();
assert_eq!(
slack,
&[0.0, 0.0],
"uncalibrated ElChe → no slack staging (partition is uniform anyway)",
);
}
#[test]
fn epoch_fn_per_rank_isolation() {
let world_size = 3usize;
let cfg = cfg_sync_cpu(world_size);
let mut coord = ClusterCoordinator::for_test(cfg);
coord.set_wall_ms_accum_for_test(0, 100.0);
coord.set_wall_ms_accum_for_test(1, 200.0);
coord.set_wall_ms_accum_for_test(2, 300.0);
coord.handle_epoch_fn_elapsed(1, 25.0);
assert!((coord.wall_ms_accum_for_test(0) - 100.0).abs() < 1e-9);
assert!((coord.wall_ms_accum_for_test(1) - 175.0).abs() < 1e-9);
assert!((coord.wall_ms_accum_for_test(2) - 300.0).abs() < 1e-9);
}
#[test]
fn checkpoint_dispatched_to_role_only() {
let world_size = 2;
let r0_got = Arc::new(AtomicBool::new(false));
let r1_got = Arc::new(AtomicBool::new(false));
let r0_flag = Arc::clone(&r0_got);
let r1_flag = Arc::clone(&r1_got);
let (port, coord_handle) = spawn_coord(
world_size,
move || cfg_sync_cpu(world_size)
.total_samples(8)
.batch_size(4)
.num_epochs(2)
.checkpoint_every(1),
move |coord| {
coord.dispatch_epoch(0)?;
coord.dispatch_epoch(1)?;
let start = Instant::now();
while start.elapsed() < Duration::from_secs(2) {
coord.tick()?;
thread::sleep(Duration::from_millis(5));
}
Ok(())
},
);
fn drain_until_shutdown(
saw_checkpoint: Arc<AtomicBool>,
send_ack_for: u64,
) -> impl Fn(&mut TcpStream, &SessionSalt) -> Result<()> {
move |s, salt| {
loop {
let msg = recv_control(s, salt)?;
match msg {
ControlMsgWire::Checkpoint { version, target_rank } => {
saw_checkpoint.store(true, Ordering::Relaxed);
if target_rank == send_ack_for {
send_metrics(s, salt, MetricsMsgWire::default()).ok();
send_timing(s, salt, TimingMsgWire::CheckpointResult {
rank: send_ack_for,
version,
elapsed_ms: 1.0,
error: None,
})?;
}
}
ControlMsgWire::Shutdown
| ControlMsgWire::ShutdownWithSave { .. } => return Ok(()),
_ => {}
}
}
}
}
let r0 = fake_rank(port, 0, world_size as u32, TEST_SALT,
drain_until_shutdown(r0_flag, 0));
let r1 = fake_rank(port, 1, world_size as u32, TEST_SALT,
drain_until_shutdown(r1_flag, u64::MAX ));
r0.join().unwrap().expect("rank 0 drained cleanly");
r1.join().unwrap().expect("rank 1 drained cleanly");
coord_handle.join().unwrap().expect("coord finishes");
assert!(r0_got.load(Ordering::Relaxed),
"rank 0 (role) must receive Checkpoint frame");
assert!(!r1_got.load(Ordering::Relaxed),
"rank 1 (non-role) must NOT receive Checkpoint frame");
}
#[test]
fn checkpoint_role_failover_on_rank_death() {
let world_size = 3usize;
let dead_ranks =
crate::distributed::controller::DeadRanks::new(world_size);
let cfg = cfg_sync_cpu(world_size).dead_ranks(Arc::clone(&dead_ranks));
let mut coord = ClusterCoordinator::for_test(cfg);
assert_eq!(coord.checkpoint_role(), 0);
let stale = Instant::now()
- Duration::from_secs(coord.heartbeat_timeout_secs() * 2 + 5);
coord.set_last_heartbeat_for_test(0, stale);
coord.check_dead_ranks_for_test();
assert!(dead_ranks.is_dead(0), "rank 0 must be declared dead");
assert_eq!(
coord.checkpoint_role(),
1,
"checkpoint_role must fail over to next live rank (1)"
);
}
#[test]
fn cooperative_intent_sets_and_folds() {
use crate::distributed::wire::IntentKind;
let world_size = 2;
let cfg = cfg_sync_cpu(world_size)
.total_samples(8)
.batch_size(4)
.num_epochs(3);
let mut coord = ClusterCoordinator::for_test(cfg);
assert!(!coord.pending_eval_intent_for_test());
assert!(!coord.pending_checkpoint_intent_for_test());
coord.process_timing_msg(TimingMsgWire::Intent {
rank: 1,
kind: IntentKind::EvalNow,
});
coord.process_timing_msg(TimingMsgWire::Intent {
rank: 0,
kind: IntentKind::CheckpointNow,
});
assert!(
coord.pending_eval_intent_for_test(),
"EvalNow intent must set the pending flag"
);
assert!(
coord.pending_checkpoint_intent_for_test(),
"CheckpointNow intent must set the pending flag"
);
let _ = coord.dispatch_epoch(1);
assert!(
!coord.pending_eval_intent_for_test(),
"dispatch_epoch must fold + clear the eval intent"
);
assert!(
!coord.pending_checkpoint_intent_for_test(),
"dispatch_epoch must fold + clear the checkpoint intent"
);
}