use super::*;
#[test]
fn heartbeat_stale_declares_rank_dead_and_unblocks_should_average() {
let world_size = 3;
let dead_ranks = crate::distributed::controller::DeadRanks::new(world_size);
let dead_for_coord = Arc::clone(&dead_ranks);
let (port, coord_handle) = spawn_coord(
world_size,
move || {
ClusterCoordinatorConfig::new(
ApplyPolicy::Sync,
AverageBackend::Cpu,
world_size,
ElChe::new(world_size, 1),
)
.no_divergence_guard()
.dead_ranks(dead_for_coord)
.heartbeat_timeout_secs(1)
},
|coord| {
let start = Instant::now();
while coord.avg_count() == 0 {
if start.elapsed() > Duration::from_secs(10) {
return Err(TensorError::new(
"heartbeat_stale: avg_count never advanced",
));
}
coord.tick()?;
thread::sleep(Duration::from_millis(20));
}
assert!(
coord.avg_count() >= 1,
"cycle finalized with surviving ranks"
);
Ok(())
},
);
let dead_for_assertion = Arc::clone(&dead_ranks);
let r2 = fake_rank(port, 2, world_size as u32, TEST_SALT, move |_s, _salt| {
thread::sleep(Duration::from_millis(3500));
Ok(())
});
let body = |rank: u64| {
move |s: &mut TcpStream, salt: &SessionSalt| -> Result<()> {
send_timing(
s,
salt,
TimingMsgWire::Batch {
rank,
batch_ms: 10.0, data_ms: 0.0,
step_count: 1,
param_norm: None,
batch_loss: 0.5,
sync_divergence: None,
},
)?;
let _ = recv_control_keepalive(s, salt, rank, 1)?; send_timing(
s,
salt,
TimingMsgWire::SyncAck {
rank,
step_count: 2,
divergence: Some(0.05),
post_norm: Some(1.0),
pre_norm: Some(1.05),
},
)?;
let _ = recv_control_keepalive(s, salt, rank, 2)?; let _ = recv_control_keepalive(s, salt, rank, 2)?; Ok(())
}
};
let r0 = fake_rank(port, 0, world_size as u32, TEST_SALT, body(0));
let r1 = fake_rank(port, 1, world_size as u32, TEST_SALT, body(1));
r0.join().unwrap().expect("rank 0 completes averaging");
r1.join().unwrap().expect("rank 1 completes averaging");
let _ = r2.join();
coord_handle.join().unwrap().expect("coord drives clean");
assert!(
dead_for_assertion.is_dead(2),
"rank 2 must be flagged dead in shared ledger"
);
}
#[test]
fn dead_rank_remainder_redistributed_via_extend_partition() {
let world_size = 3;
let total_samples = 30;
let batch_size = 1;
let dead_ranks = crate::distributed::controller::DeadRanks::new(world_size);
let dead_for_coord = Arc::clone(&dead_ranks);
let (port, coord_handle) = spawn_coord(
world_size,
move || {
ClusterCoordinatorConfig::new(
ApplyPolicy::Sync,
AverageBackend::Cpu,
world_size,
ElChe::new(world_size, 1),
)
.no_divergence_guard()
.dead_ranks(dead_for_coord)
.heartbeat_timeout_secs(1)
.total_samples(total_samples)
.batch_size(batch_size)
.num_epochs(1)
},
|coord| {
coord.dispatch_epoch(0)?;
let start = Instant::now();
while !coord.dead_ranks.as_ref().unwrap().is_dead(2) {
if start.elapsed() > Duration::from_secs(5) {
return Err(TensorError::new(
"dead_rank_redistribute: rank 2 never declared dead",
));
}
coord.tick()?;
thread::sleep(Duration::from_millis(20));
}
Ok(())
},
);
use std::sync::atomic::AtomicU64;
let r0_extension = Arc::new(AtomicU64::new(0));
let r1_extension = Arc::new(AtomicU64::new(0));
let r0_acc = Arc::clone(&r0_extension);
let r1_acc = Arc::clone(&r1_extension);
let make_alive = |rank: u64, acc: Arc<AtomicU64>| {
move |s: &mut TcpStream, salt: &SessionSalt| -> Result<()> {
let mut received_start_epoch = false;
let read_deadline = Instant::now() + Duration::from_secs(4);
while Instant::now() < read_deadline {
let _ = send_timing(s, salt, TimingMsgWire::Heartbeat { rank, step_count: 1 });
s.set_read_timeout(Some(Duration::from_millis(200))).ok();
match recv_frame(s, salt) {
Ok(Some(frame)) => match frame.decode::<ControlMsgWire>() {
Ok(ControlMsgWire::StartEpoch(_)) => {
received_start_epoch = true;
send_timing(
s,
salt,
TimingMsgWire::Batch {
rank,
batch_ms: 5.0, data_ms: 0.0,
step_count: 1,
param_norm: None,
batch_loss: 0.1,
sync_divergence: None,
},
)?;
}
Ok(ControlMsgWire::ExtendPartition {
partition_size,
..
}) => {
acc.fetch_add(partition_size, Ordering::SeqCst);
}
Ok(_other) => {
}
Err(_) => break,
},
Ok(None) => break,
Err(_) => continue,
}
}
assert!(received_start_epoch, "rank {rank} got StartEpoch");
Ok(())
}
};
let r2 = fake_rank(port, 2, world_size as u32, TEST_SALT, |_s, _salt| {
thread::sleep(Duration::from_millis(3500));
Ok(())
});
let r0 = fake_rank(port, 0, world_size as u32, TEST_SALT, make_alive(0, r0_acc));
let r1 = fake_rank(port, 1, world_size as u32, TEST_SALT, make_alive(1, r1_acc));
r0.join().unwrap().expect("rank 0 path");
r1.join().unwrap().expect("rank 1 path");
let _ = r2.join();
coord_handle.join().unwrap().expect("coord drives clean");
let r0_total = r0_extension.load(Ordering::SeqCst);
let r1_total = r1_extension.load(Ordering::SeqCst);
let total_redistributed = r0_total + r1_total;
assert_eq!(
total_redistributed, 10,
"dead rank 2's un-processed remainder (10) must be reshared \
across survivors; got r0={r0_total}, r1={r1_total}"
);
}
#[test]
fn dead_ranks_optional_default_disables_elastic_membership() {
let world_size = 2;
let (port, coord_handle) = spawn_coord(
world_size,
move || {
cfg_sync_cpu(world_size).heartbeat_timeout_secs(0)
},
move |coord| {
thread::sleep(Duration::from_millis(50));
coord.tick()?;
assert_eq!(
coord.active_count(),
world_size,
"dead-rank detection must be off without ledger"
);
Ok(())
},
);
let r0 = fake_rank(port, 0, world_size as u32, TEST_SALT, |_, _| Ok(()));
let r1 = fake_rank(port, 1, world_size as u32, TEST_SALT, |_, _| Ok(()));
r0.join().unwrap().expect("rank 0 handshake");
r1.join().unwrap().expect("rank 1 handshake");
coord_handle.join().unwrap().expect("coord drives clean");
}
#[test]
fn observe_meta_runs_during_averaging_cycle_no_anchor_change_in_probe() {
let world_size = 2;
let initial_anchor = 1usize;
let (port, coord_handle) = spawn_coord(
world_size,
move || cfg_sync_nccl(world_size).meta_controller(true),
move |coord| {
let start = Instant::now();
while coord.avg_count() == 0 {
if start.elapsed() > Duration::from_secs(5) {
return Err(TensorError::new(
"observe_meta_runs: avg_count never advanced",
));
}
coord.tick()?;
thread::sleep(Duration::from_millis(10));
}
assert_eq!(
coord.el_che().anchor(),
initial_anchor,
"Probe-phase meta must NOT nudge the anchor on first cycle"
);
let lrs = coord.last_lr_per_rank_for_test();
assert!(lrs.iter().all(|lr| lr.is_some()), "LRs captured");
Ok(())
},
);
let body = |rank: u32| {
let rank = rank as u64;
move |s: &mut TcpStream, salt: &SessionSalt| -> Result<()> {
send_timing(
s,
salt,
TimingMsgWire::LrUpdate { rank, lr: 0.01 },
)?;
send_timing(
s,
salt,
TimingMsgWire::Batch {
rank,
batch_ms: 10.0, data_ms: 0.0,
step_count: 1,
param_norm: None,
batch_loss: 1.0,
sync_divergence: None,
},
)?;
let _ = recv_control(s, salt)?;
send_timing(
s,
salt,
TimingMsgWire::SyncAck {
rank,
step_count: 2,
divergence: Some(0.05),
post_norm: Some(1.0),
pre_norm: Some(1.05),
},
)?;
let _ = recv_control(s, salt)?; Ok(())
}
};
let r0 = fake_rank(port, 0, world_size as u32, TEST_SALT, body(0));
let r1 = fake_rank(port, 1, world_size as u32, TEST_SALT, body(1));
r0.join().unwrap().expect("rank 0 path");
r1.join().unwrap().expect("rank 1 path");
coord_handle.join().unwrap().expect("coord cycle 1 with meta on");
}
#[test]
fn max_failure_threshold_breach_dispatches_shutdown_with_save() {
let world_size = 3;
let dead_ranks =
crate::distributed::controller::DeadRanks::new(world_size);
let dead_for_coord = Arc::clone(&dead_ranks);
let (port, coord_handle) = spawn_coord(
world_size,
move || {
ClusterCoordinatorConfig::new(
ApplyPolicy::Sync,
AverageBackend::Cpu,
world_size,
ElChe::new(world_size, 1),
)
.no_divergence_guard()
.dead_ranks(dead_for_coord)
.heartbeat_timeout_secs(1)
.max_failure(
crate::distributed::max_failure::MaxFailureThreshold::Absolute(1),
)
},
|coord| {
let start = Instant::now();
while !coord.shutdown_with_save_dispatched() {
if start.elapsed() > Duration::from_secs(10) {
return Err(TensorError::new(
"max_failure: ShutdownWithSave never dispatched",
));
}
coord.tick()?;
thread::sleep(Duration::from_millis(20));
}
Ok(())
},
);
fn drain_shutdown_with_save(
s: &mut TcpStream,
salt: &SessionSalt,
) -> Result<()> {
s.set_read_timeout(Some(Duration::from_secs(5)))
.map_err(|e| TensorError::new(&format!("timeout: {e}")))?;
let msg = recv_control(s, salt)?;
match msg {
ControlMsgWire::ShutdownWithSave { reason } => {
let r = crate::distributed::SaveReason::from_u8(reason)
.expect("known SaveReason variant");
if r != crate::distributed::SaveReason::MaxFailureExceeded {
return Err(TensorError::new(&format!(
"expected MaxFailureExceeded, got {r:?}"
)));
}
Ok(())
}
other => Err(TensorError::new(&format!(
"expected ShutdownWithSave, got {other:?}"
))),
}
}
let r0 = fake_rank(
port,
0,
world_size as u32,
TEST_SALT,
drain_shutdown_with_save,
);
let r1 = fake_rank(
port,
1,
world_size as u32,
TEST_SALT,
drain_shutdown_with_save,
);
let r2 = fake_rank(
port,
2,
world_size as u32,
TEST_SALT,
drain_shutdown_with_save,
);
r0.join().unwrap().expect("rank 0 receives ShutdownWithSave");
r1.join().unwrap().expect("rank 1 receives ShutdownWithSave");
r2.join().unwrap().expect("rank 2 receives ShutdownWithSave");
coord_handle.join().unwrap().expect("coord dispatched broadcast");
}
#[test]
fn controller_writes_meta_json_on_shutdown_with_save() {
let world_size = 3;
let dir = std::env::temp_dir().join(format!(
"flodl_coord_meta_{}",
std::process::id()
));
std::fs::create_dir_all(&dir).unwrap();
let stem = dir.join("coord_ckpt");
let stem_str = stem.to_str().unwrap().to_string();
let dead_ranks =
crate::distributed::controller::DeadRanks::new(world_size);
let dead_for_coord = Arc::clone(&dead_ranks);
let stem_for_coord = stem_str.clone();
let (port, coord_handle) = spawn_coord(
world_size,
move || {
ClusterCoordinatorConfig::new(
ApplyPolicy::Sync,
AverageBackend::Cpu,
world_size,
ElChe::new(world_size, 3),
)
.no_divergence_guard()
.dead_ranks(dead_for_coord)
.heartbeat_timeout_secs(1)
.max_failure(
crate::distributed::max_failure::MaxFailureThreshold::Absolute(1),
)
.save_path(stem_for_coord.clone())
},
|coord| {
let start = Instant::now();
while !coord.shutdown_with_save_dispatched() {
if start.elapsed() > Duration::from_secs(10) {
return Err(TensorError::new(
"coord meta: ShutdownWithSave never dispatched",
));
}
coord.tick()?;
thread::sleep(Duration::from_millis(20));
}
Ok(())
},
);
let r0 = fake_rank(
port,
0,
world_size as u32,
TEST_SALT,
|s: &mut TcpStream, salt: &SessionSalt| -> Result<()> {
s.set_read_timeout(Some(Duration::from_secs(5))).ok();
let _ = recv_control(s, salt)?; Ok(())
},
);
let r1 = fake_rank(
port,
1,
world_size as u32,
TEST_SALT,
|s: &mut TcpStream, salt: &SessionSalt| -> Result<()> {
s.set_read_timeout(Some(Duration::from_secs(5))).ok();
let _ = recv_control(s, salt)?;
Ok(())
},
);
let r2 = fake_rank(
port,
2,
world_size as u32,
TEST_SALT,
|s: &mut TcpStream, salt: &SessionSalt| -> Result<()> {
s.set_read_timeout(Some(Duration::from_secs(5))).ok();
let _ = recv_control(s, salt)?;
Ok(())
},
);
r0.join().unwrap().expect("rank 0 path");
r1.join().unwrap().expect("rank 1 path");
r2.join().unwrap().expect("rank 2 path");
coord_handle.join().unwrap().expect("coord dispatched");
let meta_path =
crate::distributed::CheckpointBundle::meta_path(&stem_str);
assert!(
meta_path.exists(),
"controller meta.json missing at {}",
meta_path.display(),
);
let meta =
crate::distributed::CheckpointMeta::read_from_file(&meta_path)
.expect("controller-written meta parses");
assert_eq!(meta.world_size_at_save, world_size);
assert_eq!(
meta.save_reason,
crate::distributed::SaveReason::MaxFailureExceeded,
);
let state = meta
.elche_state
.expect("controller writes elche_state into meta");
assert_eq!(state.anchor, 3);
assert_eq!(state.smoothed_ms_per_batch.len(), world_size);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn rendezvous_retry_picks_next_survivor_on_generator_death() {
let world_size = 3;
let dead_ranks = crate::distributed::controller::DeadRanks::new(world_size);
let dead_for_coord = Arc::clone(&dead_ranks);
let dead_for_test = Arc::clone(&dead_ranks);
let (port, coord_handle) = spawn_coord(
world_size,
move || {
cfg_sync_nccl(world_size)
.dead_ranks(dead_for_coord)
.heartbeat_timeout_secs(60)
.rendezvous_timeout_secs(60)
},
move |coord| {
coord.test_seed_rendezvous_pending(0, vec![0, 1, 2], 0);
dead_for_test.declare_dead(0);
coord.tick()?; assert_eq!(
coord.rendezvous_pending_generator(),
Some(1),
"retry must pick rank 1 (next ascending survivor)"
);
assert_eq!(
coord.rendezvous_tried_generators(),
vec![0],
"rank 0 recorded as tried"
);
Ok(())
},
);
let r0 = fake_rank(port, 0, world_size as u32, TEST_SALT, |_s, _salt| Ok(()));
let r1 = fake_rank(port, 1, world_size as u32, TEST_SALT, move |s, salt| {
let msg = recv_control(s, salt)?;
match msg {
ControlMsgWire::RequestNewNcclId => Ok(()),
other => Err(TensorError::new(&format!(
"rank 1 expected RequestNewNcclId, got {other:?}"
))),
}
});
let r2 = fake_rank(port, 2, world_size as u32, TEST_SALT, |_s, _salt| Ok(()));
r0.join().unwrap().expect("rank 0 handshake");
r1.join().unwrap().expect("rank 1 receives RequestNewNcclId");
r2.join().unwrap().expect("rank 2 handshake");
coord_handle.join().unwrap().expect("coord drives clean");
}
#[test]
fn rendezvous_retry_fires_on_timeout_without_death() {
let world_size = 3;
let dead_ranks = crate::distributed::controller::DeadRanks::new(world_size);
let dead_for_coord = Arc::clone(&dead_ranks);
let (port, coord_handle) = spawn_coord(
world_size,
move || {
cfg_sync_nccl(world_size)
.dead_ranks(dead_for_coord)
.heartbeat_timeout_secs(60)
.rendezvous_timeout_secs(1)
},
move |coord| {
coord.test_seed_rendezvous_pending(0, vec![0, 1, 2], 10);
coord.tick()?; assert_eq!(
coord.rendezvous_pending_generator(),
Some(1),
"timeout retry must pick the next ascending survivor (rank 0 timed out)"
);
assert_eq!(
coord.rendezvous_tried_generators(),
vec![0],
"rank 0 recorded as tried on timeout"
);
Ok(())
},
);
let r0 = fake_rank(port, 0, world_size as u32, TEST_SALT, |_s, _salt| Ok(()));
let r1 = fake_rank(port, 1, world_size as u32, TEST_SALT, move |s, salt| {
let msg = recv_control(s, salt)?;
match msg {
ControlMsgWire::RequestNewNcclId => Ok(()),
other => Err(TensorError::new(&format!(
"rank 1 expected RequestNewNcclId on timeout retry, got {other:?}"
))),
}
});
let r2 = fake_rank(port, 2, world_size as u32, TEST_SALT, |_s, _salt| Ok(()));
r0.join().unwrap().expect("rank 0 handshake");
r1.join().unwrap().expect("rank 1 receives RequestNewNcclId on timeout");
r2.join().unwrap().expect("rank 2 handshake");
coord_handle.join().unwrap().expect("coord drives clean");
}
#[test]
fn rendezvous_exhaustion_dispatches_shutdown_with_save() {
let world_size = 3;
let dir = std::env::temp_dir().join(format!(
"flodl_rdv_exhaust_{}",
std::process::id()
));
std::fs::create_dir_all(&dir).unwrap();
let stem = dir.join("ckpt").to_string_lossy().into_owned();
let (port, coord_handle) = spawn_coord(
world_size,
move || {
cfg_sync_nccl_with_dataset(world_size, 12)
.heartbeat_timeout_secs(60)
.rendezvous_timeout_secs(1)
.save_path(stem.clone())
},
move |coord| {
coord.test_seed_rendezvous_pending(
0,
Vec::new(),
10,
);
coord.tick()?;
assert!(
coord.rendezvous_pending_generator().is_none(),
"exhausted pool must clear pending"
);
assert!(
coord.shutdown_with_save_dispatched(),
"exhausted pool must dispatch ShutdownWithSave"
);
Ok(())
},
);
fn drain_shutdown(s: &mut TcpStream, salt: &SessionSalt) -> Result<()> {
s.set_read_timeout(Some(Duration::from_secs(5)))
.map_err(|e| TensorError::new(&format!("timeout: {e}")))?;
match recv_control(s, salt)? {
ControlMsgWire::ShutdownWithSave { .. } => Ok(()),
other => Err(TensorError::new(&format!(
"expected ShutdownWithSave, got {other:?}"
))),
}
}
let r0 = fake_rank(port, 0, world_size as u32, TEST_SALT, drain_shutdown);
let r1 = fake_rank(port, 1, world_size as u32, TEST_SALT, drain_shutdown);
let r2 = fake_rank(port, 2, world_size as u32, TEST_SALT, drain_shutdown);
r0.join().unwrap().expect("rank 0 receives ShutdownWithSave");
r1.join().unwrap().expect("rank 1 receives ShutdownWithSave");
r2.join().unwrap().expect("rank 2 receives ShutdownWithSave");
coord_handle.join().unwrap().expect("coord drives clean");
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn epoch_transition_dispatches_next_then_shutdowns_at_horizon() {
let world_size = 2;
let num_epochs = 2;
let (port, coord_handle) = spawn_coord(
world_size,
move || cfg_sync_cpu(world_size)
.total_samples(8)
.batch_size(4)
.num_epochs(num_epochs),
move |coord| {
coord.dispatch_epoch(0)?;
let start = Instant::now();
loop {
if start.elapsed() > Duration::from_secs(10) {
return Err(TensorError::new(
"coord did not drain within 10s",
));
}
if !coord.tick()? {
break;
}
thread::sleep(Duration::from_millis(5));
}
assert_eq!(
coord.last_aggregated_epoch(),
Some(num_epochs - 1),
"both epochs must have aggregated",
);
Ok(())
},
);
fn rank_body(
rank: u64,
num_epochs: usize,
) -> impl Fn(&mut TcpStream, &SessionSalt) -> Result<()> {
move |s, salt| {
let mut completed = 0usize;
let mut saw_shutdown = false;
while !saw_shutdown {
let msg = recv_control(s, salt)?;
match msg {
ControlMsgWire::StartEpoch(plan) => {
send_metrics(s, salt, MetricsMsgWire {
rank,
epoch: plan.epoch,
avg_loss: 0.5,
batches_processed: 2,
epoch_ms: 50.0,
samples_processed: 4,
share_complete_ms: 0.0,
compute_only_ms: 50.0,
data_starve_ms: 0.0,
scalars: std::collections::HashMap::new(),
resources: None,
})?;
completed += 1;
}
ControlMsgWire::Shutdown
| ControlMsgWire::ShutdownWithSave { .. } => {
saw_shutdown = true;
}
_ => {}
}
}
if completed != num_epochs {
return Err(TensorError::new(&format!(
"rank {rank}: received Shutdown after {completed} epochs \
(expected {num_epochs})",
)));
}
Ok(())
}
}
let r0 = fake_rank(port, 0, world_size as u32, TEST_SALT,
rank_body(0, num_epochs));
let r1 = fake_rank(port, 1, world_size as u32, TEST_SALT,
rank_body(1, num_epochs));
r0.join().unwrap().expect("rank 0 completed all epochs");
r1.join().unwrap().expect("rank 1 completed all epochs");
coord_handle.join().unwrap().expect("coord finishes cleanly");
}
#[test]
fn reported_death_declared_via_drain_and_cycle_completes() {
let world_size = 3;
let dead_ranks = crate::distributed::controller::DeadRanks::new(world_size);
let dead_for_coord = Arc::clone(&dead_ranks);
let reported: crate::distributed::cluster_coordinator::ReportedDeaths =
Arc::new(std::sync::Mutex::new(Vec::new()));
let reported_for_coord = Arc::clone(&reported);
let (port, coord_handle) = spawn_coord(
world_size,
move || {
ClusterCoordinatorConfig::new(
ApplyPolicy::Sync,
AverageBackend::Cpu,
world_size,
ElChe::new(world_size, 1),
)
.no_divergence_guard()
.dead_ranks(dead_for_coord)
.reported_deaths(reported_for_coord)
.heartbeat_timeout_secs(30)
},
|coord| {
let start = Instant::now();
while coord.avg_count() == 0 {
if start.elapsed() > Duration::from_secs(10) {
return Err(TensorError::new(
"reported_death: avg_count never advanced",
));
}
coord.tick()?;
thread::sleep(Duration::from_millis(20));
}
Ok(())
},
);
let r2 = fake_rank(port, 2, world_size as u32, TEST_SALT, move |_s, _salt| {
thread::sleep(Duration::from_millis(3500));
Ok(())
});
let reporter = {
let q = Arc::clone(&reported);
thread::spawn(move || {
thread::sleep(Duration::from_millis(800));
q.lock().unwrap().push(2);
})
};
let body = |rank: u64| {
move |s: &mut TcpStream, salt: &SessionSalt| -> Result<()> {
send_timing(
s,
salt,
TimingMsgWire::Batch {
rank,
batch_ms: 10.0, data_ms: 0.0,
step_count: 1,
param_norm: None,
batch_loss: 0.5,
sync_divergence: None,
},
)?;
let _ = recv_control_keepalive(s, salt, rank, 1)?; send_timing(
s,
salt,
TimingMsgWire::SyncAck {
rank,
step_count: 2,
divergence: Some(0.05),
post_norm: Some(1.0),
pre_norm: Some(1.05),
},
)?;
let _ = recv_control_keepalive(s, salt, rank, 2)?; let _ = recv_control_keepalive(s, salt, rank, 2)?; Ok(())
}
};
let r0 = fake_rank(port, 0, world_size as u32, TEST_SALT, body(0));
let r1 = fake_rank(port, 1, world_size as u32, TEST_SALT, body(1));
r0.join().unwrap().expect("rank 0 completes averaging");
r1.join().unwrap().expect("rank 1 completes averaging");
let _ = r2.join();
let _ = reporter.join();
coord_handle.join().unwrap().expect("coord drives clean");
assert!(
dead_ranks.is_dead(2),
"reported death must land in the shared ledger via the drain"
);
assert!(
!dead_ranks.is_dead(0) && !dead_ranks.is_dead(1),
"survivors must not be reaped (staleness window is 30s)"
);
assert!(
reported.lock().unwrap().is_empty(),
"queue must be drained by the tick"
);
}
#[test]
fn exiting_latch_suppresses_late_death_report_exactly_once() {
let world_size = 3;
let dead_ranks = crate::distributed::controller::DeadRanks::new(world_size);
let dead_for_coord = Arc::clone(&dead_ranks);
let reported: crate::distributed::cluster_coordinator::ReportedDeaths =
Arc::new(std::sync::Mutex::new(Vec::new()));
let reported_for_coord = Arc::clone(&reported);
let (port, coord_handle) = spawn_coord(
world_size,
move || {
ClusterCoordinatorConfig::new(
ApplyPolicy::Sync,
AverageBackend::Cpu,
world_size,
ElChe::new(world_size, 1),
)
.no_divergence_guard()
.dead_ranks(dead_for_coord)
.reported_deaths(reported_for_coord)
.heartbeat_timeout_secs(30)
},
|coord| {
let start = Instant::now();
while start.elapsed() < Duration::from_secs(2) {
coord.tick()?;
thread::sleep(Duration::from_millis(20));
}
if coord.active_count() != 2 {
return Err(TensorError::new(&format!(
"active_count must decrement exactly once for a \
cleanly-exited rank (Exiting latch), got {} of {}",
coord.active_count(),
3,
)));
}
Ok(())
},
);
let r2 = fake_rank(port, 2, world_size as u32, TEST_SALT, move |s, salt| {
send_timing(s, salt, TimingMsgWire::Exiting { rank: 2 })?;
thread::sleep(Duration::from_millis(2500));
Ok(())
});
let reporter = {
let q = Arc::clone(&reported);
thread::spawn(move || {
thread::sleep(Duration::from_millis(600));
q.lock().unwrap().push(2);
})
};
let body = |rank: u64| {
move |s: &mut TcpStream, salt: &SessionSalt| -> Result<()> {
send_timing(
s,
salt,
TimingMsgWire::Batch {
rank,
batch_ms: 10.0, data_ms: 0.0,
step_count: 1,
param_norm: None,
batch_loss: 0.5,
sync_divergence: None,
},
)?;
thread::sleep(Duration::from_millis(2500));
Ok(())
}
};
let r0 = fake_rank(port, 0, world_size as u32, TEST_SALT, body(0));
let r1 = fake_rank(port, 1, world_size as u32, TEST_SALT, body(1));
let _ = r0.join();
let _ = r1.join();
let _ = r2.join();
let _ = reporter.join();
coord_handle.join().unwrap().expect("coord drives clean");
assert!(
!dead_ranks.is_dead(2),
"a cleanly-exited rank must never be declared dead by a late \
report (the drain skips exited ranks)"
);
assert!(
reported.lock().unwrap().is_empty(),
"the late report must still be drained (consumed, not left queued)"
);
}