use super::*;
use super::spawn::{
build_remote_agent_bash_command, build_slim_envelope_for,
shell_quote, supervise_children, PerHostPrebuild,
};
use serde_json::json;
use std::process::Command;
fn sample_relay_spec() -> RelaySpec {
RelaySpec {
host: "pascal".into(),
controller_host: "192.168.122.1".into(),
controller_port: 1337,
ranks: vec![1, 2],
salt_hex: "0123456789abcdef0123456789abcdef".into(),
world_size: 3,
data_channel: true,
frame_ceiling_bytes: 0,
}
}
#[test]
fn relay_spec_hex_json_round_trips() {
let spec = sample_relay_spec();
let hex = crate::distributed::cluster::hex_encode(
serde_json::to_string(&spec).unwrap().as_bytes(),
);
let bytes = crate::distributed::cluster::hex_decode(&hex).unwrap();
let back: RelaySpec = serde_json::from_slice(&bytes).unwrap();
assert_eq!(back, spec);
}
#[test]
fn agent_bash_command_exports_agent_env_only() {
let cmd = build_remote_agent_bash_command(
"/opt/flodl",
"pascal",
None,
"ddp-bench",
&["--model".into(), "resnet-graph".into()],
&std::collections::BTreeMap::new(),
&std::collections::BTreeMap::new(),
None,
);
assert!(
cmd.contains("IFS= read -r __FLODL_ENVELOPE"),
"agent spec must be read from stdin: {cmd}"
);
assert!(
cmd.contains("FLODL_INTERNAL_AGENT_JSON=\"$__FLODL_ENVELOPE\""),
"agent env must expand the stdin var: {cmd}"
);
assert!(cmd.contains("FLODL_HOST_NAME='pascal'"), "missing host override: {cmd}");
assert!(cmd.contains("cd '/opt/flodl'"), "missing cd: {cmd}");
assert!(cmd.contains("fdl 'ddp-bench'"), "missing fdl cmd: {cmd}");
assert!(cmd.contains("--model") && cmd.contains("resnet-graph"));
assert!(!cmd.contains("FLODL_INTERNAL_CLUSTER_JSON="), "leaked rank envelope: {cmd}");
assert!(!cmd.contains("FLODL_INTERNAL_LOCAL_RANK="), "leaked rank slot: {cmd}");
assert!(!cmd.contains("FLODL_INTERNAL_RELAY_JSON="), "leaked relay spec: {cmd}");
assert!(!cmd.contains("CUDA_VISIBLE_DEVICES="), "agent must not scope CUDA: {cmd}");
assert!(cmd.contains("trap ") && cmd.contains("__flodl_pid"), "missing trap: {cmd}");
}
#[test]
fn join_knobs_parse_round_trip_and_reject_typos() {
let mut val = canonical_full_json();
val["controller"]["join"] = json!({
"min_rank_start": 2,
"join_timeout": 120,
"open_admission": true,
});
let full = FullCluster::from_value(&val).unwrap();
let knobs = full.controller.join.as_ref().expect("join block parsed");
assert_eq!(knobs.min_rank_start, Some(2));
assert_eq!(knobs.join_timeout_secs, Some(120));
assert_eq!(knobs.target_ranks, None);
assert_eq!(knobs.max_join_timeout_secs, None);
assert_eq!(knobs.open_admission, Some(true));
let back = FullCluster::from_value(&full.to_json()).unwrap();
assert_eq!(back.controller.join, full.controller.join);
let mut bad = canonical_full_json();
bad["controller"]["join"] = json!({ "join_timeout_secs": 120 });
let msg = FullCluster::from_value(&bad).unwrap_err().to_string();
assert!(msg.contains("unknown field"), "got: {msg}");
let mut bad = canonical_full_json();
bad["controller"]["join"] = json!({ "min_rank_start": "two" });
let msg = FullCluster::from_value(&bad).unwrap_err().to_string();
assert!(msg.contains("non-negative integer"), "got: {msg}");
}
#[test]
fn join_config_derivation_defaults_to_capacity_all_or_nothing() {
let cfg = super::derive_join_config(None, 3);
assert_eq!(cfg.min_rank_start, 3);
assert_eq!(cfg.target_ranks, Some(3));
assert_eq!(cfg.join_timeout_secs, 300);
assert_eq!(cfg.max_join_timeout_secs, 600);
assert!(!cfg.open_admission);
let knobs = JoinKnobs { min_rank_start: Some(2), ..Default::default() };
let cfg = super::derive_join_config(Some(&knobs), 3);
assert_eq!(cfg.min_rank_start, 2);
assert_eq!(cfg.target_ranks, Some(3));
let knobs = JoinKnobs { join_timeout_secs: Some(900), ..Default::default() };
let cfg = super::derive_join_config(Some(&knobs), 3);
assert_eq!(cfg.join_timeout_secs, 900);
assert_eq!(cfg.max_join_timeout_secs, 900);
cfg.validate().unwrap();
}
#[test]
fn synthesized_world_reranks_config_hosts_and_admits_walk_ins() {
let full = FullCluster::from_value(&canonical_full_json()).unwrap();
let salt = [5u8; crate::distributed::wire::SESSION_SALT_BYTES];
let members = [
crate::distributed::membership::JoinedMember {
host: "host-b".into(),
ranks: vec![0, 1],
local_devices: vec![0, 1],
gpus: vec!["B".into(); 2],
libtorch: String::new(),
joined_at_secs: 0,
},
crate::distributed::membership::JoinedMember {
host: "cloud-worker".into(),
ranks: vec![2],
local_devices: vec![0],
gpus: vec!["C".into()],
libtorch: String::new(),
joined_at_secs: 1,
},
crate::distributed::membership::JoinedMember {
host: "host-a".into(),
ranks: vec![3],
local_devices: vec![1],
gpus: vec!["A".into()],
libtorch: String::new(),
joined_at_secs: 2,
},
];
let world = super::synthesize_world(&full, members.iter(), salt);
assert_eq!(world.world_size(), 4);
assert_eq!(world.salt, salt);
assert_eq!(world.workers[0].host, "host-b");
assert_eq!(world.workers[0].ranks, vec![0, 1]);
assert_eq!(world.workers[0].local_devices, Some(vec![0, 1]));
let cfg_b = full.workers.iter().find(|w| w.host == "host-b").unwrap();
assert_eq!(world.workers[0].path, cfg_b.path);
assert_eq!(world.workers[0].nccl_socket_ifname, cfg_b.nccl_socket_ifname);
assert!(world.workers[0].ssh.is_some());
assert_eq!(world.workers[1].host, "cloud-worker");
assert_eq!(world.workers[1].ranks, vec![2]);
assert!(world.workers[1].ssh.is_none());
assert!(!world.workers[1].tunnel);
assert_eq!(world.workers[2].host, "host-a");
assert_eq!(world.workers[2].ranks, vec![3]);
}
fn canonical_full_json() -> serde_json::Value {
json!({
"controller": {
"host": "192.168.122.1",
"port": 29500,
"path": "/opt/flodl"
},
"workers": [
{
"host": "host-a",
"ranks": [0],
"local_devices": [0],
"nccl_socket_ifname": "virbr0",
"path": "/opt/flodl",
"arch": "precompiled/cu128"
},
{
"host": "host-b",
"ssh": { "target": "host-b" },
"ranks": [1, 2],
"local_devices": "all",
"nccl_socket_ifname": "enp1s0",
"path": "/srv/flodl"
}
]
})
}
#[test]
fn parses_full_topology() {
let c = FullCluster::from_value(&canonical_full_json()).unwrap();
assert_eq!(c.controller.host, "192.168.122.1");
assert_eq!(c.controller.port, 29500);
assert_eq!(c.world_size(), 3);
assert!(c.spans_multiple_workers());
assert_eq!(c.workers.len(), 2);
assert_eq!(c.workers[0].host, "host-a");
assert_eq!(c.workers[0].ranks, vec![0]);
assert_eq!(c.workers[0].local_devices, Some(vec![0]));
assert_eq!(c.workers[0].ssh, None);
assert_eq!(c.workers[1].host, "host-b");
assert_eq!(c.workers[1].ranks, vec![1, 2]);
assert_eq!(c.workers[1].local_devices, None);
assert_eq!(c.workers[1].ssh_target(), "host-b");
}
#[test]
fn rejects_empty_workers() {
let mut v = canonical_full_json();
v["workers"] = json!([]);
let err = FullCluster::from_value(&v).unwrap_err();
assert!(err.to_string().contains("workers must be non-empty"), "got: {err}");
}
#[test]
fn rejects_rank_gap_across_hosts() {
let mut v = canonical_full_json();
v["workers"][1]["ranks"] = json!([2, 3]); let err = FullCluster::from_value(&v).unwrap_err();
assert!(
err.to_string().contains("duplicates or gaps"),
"got: {err}"
);
}
#[test]
fn rejects_duplicate_ranks() {
let mut v = canonical_full_json();
v["workers"][1]["ranks"] = json!([0, 1]); let err = FullCluster::from_value(&v).unwrap_err();
assert!(
err.to_string().contains("duplicates or gaps"),
"got: {err}"
);
}
#[test]
fn rejects_local_devices_length_mismatch_for_explicit() {
let mut v = canonical_full_json();
v["workers"][1]["local_devices"] = json!([0]); let err = FullCluster::from_value(&v).unwrap_err();
assert!(err.to_string().contains("length mismatch"), "got: {err}");
}
#[test]
fn accepts_local_devices_all_at_launcher_parse_time() {
let mut v = canonical_full_json();
v["workers"][0]["local_devices"] = json!("all");
let c = FullCluster::from_value(&v).unwrap();
assert_eq!(c.workers[0].local_devices, None);
}
#[test]
fn rejects_unknown_local_devices_string() {
let mut v = canonical_full_json();
v["workers"][0]["local_devices"] = json!("every");
let err = FullCluster::from_value(&v).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("local_devices") && msg.contains("every"),
"got: {msg}"
);
}
#[test]
fn rejects_controller_port_overflow() {
let mut v = canonical_full_json();
v["controller"]["port"] = json!(100_000);
let err = FullCluster::from_value(&v).unwrap_err();
assert!(err.to_string().contains("u16"), "got: {err}");
}
#[test]
fn slim_envelope_strips_ssh_carries_metadata() {
let full = FullCluster::from_value(&canonical_full_json()).unwrap();
let worker = full.workers.iter().find(|h| h.host == "host-b").unwrap();
let env = build_slim_envelope_for(&full, worker, &full.controller.host, false);
assert_eq!(env["controller"]["host"], "192.168.122.1");
assert_eq!(env["controller"]["port"], 29500);
assert_eq!(env["world_size"], 3);
assert_eq!(env["num_workers"], 2);
assert_eq!(env["worker"]["host"], "host-b");
assert_eq!(env["worker"]["ranks"], serde_json::json!([1, 2]));
assert_eq!(env["worker"]["local_devices"], serde_json::json!("all"));
assert_eq!(env["worker"]["nccl_socket_ifname"], "enp1s0");
assert!(env["worker"].get("ssh").is_none(), "ssh must be stripped");
}
#[test]
fn slim_envelope_emits_explicit_local_devices_when_present() {
let full = FullCluster::from_value(&canonical_full_json()).unwrap();
let host_a = full.workers.iter().find(|h| h.host == "host-a").unwrap();
let env = build_slim_envelope_for(&full, host_a, &full.controller.host, false);
assert_eq!(env["worker"]["local_devices"], serde_json::json!([0]));
}
#[test]
fn slim_envelope_rank_resources_round_trips() {
let full = FullCluster::from_value(&canonical_full_json()).unwrap();
let host_a = full.workers.iter().find(|h| h.host == "host-a").unwrap();
let off = build_slim_envelope_for(&full, host_a, &full.controller.host, false);
assert!(off.get("rank_resources").is_none(), "off run must not emit the key");
let parsed = crate::distributed::LocalCluster::from_value(&off).unwrap();
assert!(!parsed.rank_resources);
let on = build_slim_envelope_for(&full, host_a, &full.controller.host, true);
assert_eq!(on["rank_resources"], serde_json::json!(true));
let parsed = crate::distributed::LocalCluster::from_value(&on).unwrap();
assert!(parsed.rank_resources);
}
#[test]
fn tunnel_field_parses_defaults_false_and_round_trips() {
let c = FullCluster::from_value(&canonical_full_json()).unwrap();
assert!(!c.workers[0].tunnel);
assert!(!c.workers[1].tunnel);
let mut v = canonical_full_json();
v["workers"][1]["tunnel"] = json!(true);
let c = FullCluster::from_value(&v).unwrap();
assert!(!c.workers[0].tunnel);
assert!(c.workers[1].tunnel);
let back = FullCluster::from_value(&c.to_json()).unwrap();
assert!(back.workers[1].tunnel);
let mut v = canonical_full_json();
v["workers"][1]["tunnel"] = json!("yes");
let err = FullCluster::from_value(&v).unwrap_err();
assert!(err.to_string().contains("tunnel must be a boolean"), "got: {err}");
}
#[test]
fn tunnel_topology_rejects_launcher_local_host() {
let mut v = canonical_full_json();
v["workers"][0]["tunnel"] = json!(true);
let full = FullCluster::from_value(&v).unwrap();
let err = validate_tunnel_topology(&full, "host-a", false).unwrap_err();
assert!(err.to_string().contains("launcher host"), "got: {err}");
}
#[test]
fn tunnel_topology_rejects_nccl_backend() {
let mut v = canonical_full_json();
v["workers"][1]["tunnel"] = json!(true);
let full = FullCluster::from_value(&v).unwrap();
let err = validate_tunnel_topology(&full, "host-a", true).unwrap_err();
assert!(err.to_string().contains("NCCL"), "got: {err}");
}
#[test]
fn tunnel_topology_bind_scope() {
let full = FullCluster::from_value(&canonical_full_json()).unwrap();
assert!(!validate_tunnel_topology(&full, "host-a", false).unwrap());
let mut v = canonical_full_json();
v["workers"][1]["tunnel"] = json!(true);
let full = FullCluster::from_value(&v).unwrap();
assert!(validate_tunnel_topology(&full, "host-a", false).unwrap());
let full = FullCluster::from_value(&v).unwrap();
assert!(!validate_tunnel_topology(&full, "some-other-launcher", false).unwrap());
}
#[test]
fn slim_envelope_honors_controller_dial_host() {
let full = FullCluster::from_value(&canonical_full_json()).unwrap();
let worker = full.workers.iter().find(|h| h.host == "host-b").unwrap();
let env = build_slim_envelope_for(&full, worker, "127.0.0.1", false);
assert_eq!(env["controller"]["host"], "127.0.0.1");
assert_eq!(env["controller"]["port"], 29500);
}
#[test]
fn shell_quote_simple() {
assert_eq!(shell_quote("foo"), "'foo'");
}
#[test]
fn shell_quote_with_spaces() {
assert_eq!(shell_quote("foo bar"), "'foo bar'");
}
#[test]
fn shell_quote_escapes_internal_quotes() {
assert_eq!(shell_quote("don't"), "'don'\\''t'");
}
fn empty_env() -> std::collections::BTreeMap<String, String> {
std::collections::BTreeMap::new()
}
#[test]
fn agent_bash_command_shape() {
let cluster_env = empty_env();
let host_env = empty_env();
let s = build_remote_agent_bash_command(
"/srv/flodl",
"host-b",
Some("cluster"),
"train",
&["--epochs".to_string(), "10".to_string()],
&cluster_env,
&host_env,
None,
);
assert!(s.starts_with("IFS= read -r __FLODL_ENVELOPE\n"));
assert!(s.contains("cd '/srv/flodl' && "));
assert!(s.contains("FLODL_INTERNAL_AGENT_JSON=\"$__FLODL_ENVELOPE\""));
assert!(s.contains("FLODL_HOST_NAME='host-b'"));
assert!(s.contains("FDL_ENV='cluster'"));
assert!(s.contains("fdl 'train' '--epochs' '10' &\n"));
assert!(s.contains("trap 'kill -TERM \"$__flodl_pid\" 2>/dev/null;"));
assert!(s.contains("kill -KILL \"$__flodl_pid\""));
assert!(s.contains("wait \"$__flodl_pid\""));
assert!(s.ends_with("exit $?\n"));
}
#[test]
fn agent_bash_command_omits_fdl_env_when_none() {
let cluster_env = empty_env();
let host_env = empty_env();
let s = build_remote_agent_bash_command(
"/srv/flodl",
"worker",
None,
"train",
&[],
&cluster_env,
&host_env,
None,
);
assert!(
!s.contains("FDL_ENV"),
"FDL_ENV must be absent when overlay_env is None; got: {s}"
);
}
#[test]
fn agent_bash_command_uses_trap_wrapper() {
let cluster_env = empty_env();
let host_env = empty_env();
let s = build_remote_agent_bash_command(
"/srv", "w", None, "train", &[],
&cluster_env, &host_env, None,
);
assert!(s.contains(" fdl "), "missing `fdl` invocation: {s}");
assert!(s.contains(" &\n"), "missing background `&`: {s}");
assert!(
s.contains("__flodl_pid=$!"),
"missing `__flodl_pid=$!`: {s}"
);
assert!(
s.contains("trap 'kill -TERM \"$__flodl_pid\""),
"missing trap line: {s}"
);
assert!(s.contains("wait \"$__flodl_pid\""), "missing wait: {s}");
assert!(s.ends_with("exit $?\n"), "missing exit prop: {s}");
}
#[test]
fn agent_bash_command_quotes_dangerous_path() {
let cluster_env = empty_env();
let host_env = empty_env();
let s = build_remote_agent_bash_command(
"/srv/it's", "w", None, "train", &[],
&cluster_env, &host_env, None,
);
assert!(
s.contains("cd '/srv/it'\\''s'"),
"path with single quote not properly escaped: {s}"
);
}
#[test]
fn agent_bash_command_uses_prebuild_binary_and_ld_path() {
let cluster_env = empty_env();
let host_env = empty_env();
let pb = PerHostPrebuild {
bin: "target/cluster/worker/release/ddp-bench".into(),
ld_library_path: "/opt/libtorch/lib".into(),
cwd_subpath: "ddp-bench".into(),
};
let s = build_remote_agent_bash_command(
"/srv/flodl",
"worker",
None,
"ddp-bench",
&["--mode".into(), "nccl-sync".into()],
&cluster_env,
&host_env,
Some(&pb),
);
assert!(
s.contains("LD_LIBRARY_PATH='/opt/libtorch/lib'"),
"missing prebuild LD_LIBRARY_PATH: {s}",
);
assert!(
s.contains("cd '/srv/flodl/ddp-bench'"),
"remote cwd must cd into <host.path>/<cwd_subpath>: {s}",
);
assert!(
s.contains(" '/srv/flodl/target/cluster/worker/release/ddp-bench'"),
"binary path must be absolute (independent of cwd offset): {s}",
);
assert!(
!s.contains("fdl 'ddp-bench'"),
"prebuild path must NOT re-enter fdl on remote: {s}",
);
assert!(
s.contains("'--mode' 'nccl-sync' &\n"),
"user args must be appended ahead of the trap wrapper: {s}",
);
assert!(s.ends_with("exit $?\n"), "trap wrapper must end the cmd: {s}");
}
#[test]
fn agent_bash_command_prebuild_yields_to_host_env_ld_path() {
let cluster_env = empty_env();
let mut host_env = empty_env();
host_env.insert(
"LD_LIBRARY_PATH".into(),
"/opt/libtorch/lib:/usr/local/lib".into(),
);
let pb = PerHostPrebuild {
bin: "target/cluster/worker/release/ddp-bench".into(),
ld_library_path: "/opt/libtorch/lib".into(),
cwd_subpath: String::new(),
};
let s = build_remote_agent_bash_command(
"/srv", "w", None, "ddp-bench", &[],
&cluster_env, &host_env, Some(&pb),
);
let host_pos = s.find("LD_LIBRARY_PATH='/opt/libtorch/lib:/usr/local/lib'").unwrap();
assert!(
!s.contains(" LD_LIBRARY_PATH='/opt/libtorch/lib' "),
"auto-derived LD_LIBRARY_PATH should yield to host_env: {s}",
);
let _ = host_pos;
}
#[test]
fn agent_bash_command_exports_cluster_and_host_env() {
let mut cluster_env = empty_env();
cluster_env.insert("NCCL_P2P_DISABLE".into(), "1".into());
cluster_env.insert("SHARED_FLAG".into(), "cluster-wins".into());
let mut host_env = empty_env();
host_env.insert("HOST_FLAG".into(), "host-val".into());
host_env.insert("SHARED_FLAG".into(), "host-wins".into());
let s = build_remote_agent_bash_command(
"/srv", "w", None, "train", &[],
&cluster_env, &host_env, None,
);
assert!(s.contains("NCCL_P2P_DISABLE='1'"));
assert!(s.contains("HOST_FLAG='host-val'"));
let cluster_pos = s.find("SHARED_FLAG='cluster-wins'").unwrap();
let host_pos = s.find("SHARED_FLAG='host-wins'").unwrap();
assert!(cluster_pos < host_pos, "host env must export after cluster env");
}
#[test]
fn slim_envelope_round_trips_through_local_cluster_parser() {
let full = FullCluster::from_value(&canonical_full_json()).unwrap();
let host_a = full.workers.iter().find(|h| h.host == "host-a").unwrap();
let env = build_slim_envelope_for(&full, host_a, &full.controller.host, false);
let parsed = crate::distributed::cluster::LocalCluster::from_value(&env)
.expect("slim envelope must parse via LocalCluster::from_value");
assert_eq!(parsed.world_size(), 3);
assert_eq!(parsed.controller.host, "192.168.122.1");
assert_eq!(parsed.worker.host, "host-a");
assert_eq!(parsed.worker.ranks, vec![0]);
assert_eq!(parsed.worker.local_devices, vec![0]);
assert_eq!(parsed.salt, [0u8; crate::distributed::wire::SESSION_SALT_BYTES]);
}
#[test]
fn slim_envelope_propagates_session_salt() {
let mut full = FullCluster::from_value(&canonical_full_json()).unwrap();
let salt: crate::distributed::wire::SessionSalt = [
0xde, 0xad, 0xbe, 0xef, 0x01, 0x02, 0x03, 0x04,
0xfe, 0xed, 0xfa, 0xce, 0x05, 0x06, 0x07, 0x08,
];
full = full.with_session_salt(salt);
let host_a = full.workers.iter().find(|h| h.host == "host-a").unwrap();
let env = build_slim_envelope_for(&full, host_a, &full.controller.host, false);
let hex = env
.get("salt")
.and_then(|v| v.as_str())
.expect("envelope.salt is a string");
assert_eq!(hex.len(), 32);
let parsed = crate::distributed::cluster::LocalCluster::from_value(&env).unwrap();
assert_eq!(parsed.salt, salt);
}
#[test]
fn supervise_children_clean_exit_returns_none() {
let mut children: Vec<super::spawn::SupervisedChild> = Vec::new();
for lr in 0..2 {
let child = Command::new("true")
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.expect("spawn `true`");
children.push(("host".to_string(), lr, vec![lr], child, Vec::new()));
}
assert!(supervise_children(children, None, None).is_none());
}
#[test]
fn supervise_children_failure_terminates_peers() {
let fail_child = Command::new("sh")
.args(["-c", "exit 1"])
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.expect("spawn `sh -c 'exit 1'`");
let sleep_child = Command::new("sleep")
.arg("60")
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.expect("spawn `sleep 60`");
let children = vec![
("host-fail".to_string(), 0, vec![0], fail_child, Vec::new()),
("host-sleep".to_string(), 1, vec![1], sleep_child, Vec::new()),
];
let start = std::time::Instant::now();
let err = supervise_children(children, None, None).expect("expected failure attribution");
let elapsed = start.elapsed();
assert!(
err.to_string().contains("host-fail"),
"attribution should name the first failed rank: {err}"
);
assert!(
elapsed < std::time::Duration::from_secs(15),
"SIGTERM-on-failure must reap the sleeper well before its 60s budget; took {elapsed:?}"
);
}
fn clear_role_env() {
unsafe {
std::env::remove_var(ENV_FULL_CLUSTER_JSON);
std::env::remove_var(ENV_RELAY_JSON);
std::env::remove_var(crate::distributed::cluster::ENV_CLUSTER_JSON);
std::env::remove_var(crate::distributed::cluster::ENV_LOCAL_RANK);
}
}
#[test]
fn role_env_pristine_matches_dispatch_truth_table() {
let _guard = crate::distributed::cluster::ENV_MUTEX.lock().unwrap();
clear_role_env();
assert!(role_env_pristine());
for k in [ENV_FULL_CLUSTER_JSON, ENV_RELAY_JSON] {
unsafe { std::env::set_var(k, "deadbeef") };
assert!(!role_env_pristine(), "{k} set must not be pristine");
clear_role_env();
}
unsafe {
std::env::set_var(crate::distributed::cluster::ENV_CLUSTER_JSON, "deadbeef");
std::env::set_var(crate::distributed::cluster::ENV_LOCAL_RANK, "0");
}
assert!(!role_env_pristine());
unsafe { std::env::set_var(ENV_FULL_CLUSTER_JSON, "deadbeef") };
assert!(!role_env_pristine());
clear_role_env();
}
#[test]
fn programmatic_promotion_is_role_gated() {
let _guard = crate::distributed::cluster::ENV_MUTEX.lock().unwrap();
clear_role_env();
let full = FullCluster::from_value(&canonical_full_json()).unwrap();
unsafe {
std::env::set_var(crate::distributed::cluster::ENV_CLUSTER_JSON, "deadbeef");
std::env::set_var(crate::distributed::cluster::ENV_LOCAL_RANK, "0");
}
assert!(!promote_programmatic_cluster(&full));
assert!(std::env::var_os(ENV_FULL_CLUSTER_JSON).is_none());
clear_role_env();
unsafe { std::env::set_var(ENV_RELAY_JSON, "deadbeef") };
assert!(!promote_programmatic_cluster(&full));
assert!(std::env::var_os(ENV_FULL_CLUSTER_JSON).is_none());
clear_role_env();
assert!(promote_programmatic_cluster(&full));
assert!(matches!(dispatch(), Ok(Role::Launcher)));
let back = FullCluster::from_env().unwrap();
assert_eq!(back.controller.host, full.controller.host);
assert_eq!(back.controller.port, full.controller.port);
assert_eq!(back.world_size(), full.world_size());
assert!(!promote_programmatic_cluster(&full));
clear_role_env();
}
#[test]
fn env_block_rejects_reserved_and_malformed_keys() {
let with_env = |k: &str, v: &str| {
let mut val = canonical_full_json();
val["env"] = serde_json::json!({ k: v });
FullCluster::from_value(&val)
};
for reserved in [
"CUDA_VISIBLE_DEVICES",
"CUDA_DEVICE_ORDER",
"FLODL_INTERNAL_LOCAL_RANK",
"FLODL_HOST_NAME",
"FDL_ENV",
] {
let err = with_env(reserved, "x").unwrap_err();
assert!(
err.to_string().contains("reserved"),
"{reserved}: expected reserved-key rejection, got: {err}"
);
}
for bad in ["HAS SPACE", "1LEADING_DIGIT", "DASH-ED", ""] {
let err = with_env(bad, "x").unwrap_err();
assert!(
err.to_string().contains("valid env var name"),
"{bad:?}: expected charset rejection, got: {err}"
);
}
for ok in [
"NCCL_DEBUG",
"LD_PRELOAD",
"LD_LIBRARY_PATH",
"_UNDER",
"FLODL_VERBOSITY",
"FLODL_STAGER",
] {
assert!(
with_env(ok, "x").is_ok(),
"{ok}: legitimate tuning key must be accepted"
);
}
}
#[test]
fn worker_env_block_rejects_reserved_keys() {
let mut val = canonical_full_json();
val["workers"][0]["env"] =
serde_json::json!({ "FLODL_INTERNAL_CLUSTER_JSON": "deadbeef" });
let err = FullCluster::from_value(&val).unwrap_err();
assert!(err.to_string().contains("reserved"), "got: {err}");
}
#[test]
fn supervise_elastic_tolerates_death_and_reports_ranks() {
use crate::distributed::cluster_coordinator::ReportedDeaths;
let fail_child = Command::new("sh")
.args(["-c", "exit 7"])
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.expect("spawn failing child");
let survivor = Command::new("sh")
.args(["-c", "sleep 1; exit 0"])
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.expect("spawn survivor");
let children = vec![
("host-a".to_string(), 0, vec![0], fail_child, Vec::new()),
("host-b".to_string(), 1, vec![1], survivor, Vec::new()),
];
let queue: ReportedDeaths =
std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let dead = crate::distributed::controller::DeadRanks::new(2);
let elastic = ElasticSupervision {
reported_deaths: std::sync::Arc::clone(&queue),
dead_ranks: std::sync::Arc::clone(&dead),
max_failure: Some(
crate::distributed::max_failure::MaxFailureThreshold::Absolute(2),
),
world_size: 2,
cohort_formed: std::sync::Arc::new(
std::sync::atomic::AtomicBool::new(true),
),
};
let verdict = supervise_children(children, Some(elastic), None);
assert!(
verdict.is_none(),
"within-tolerance death must not fail the run: {verdict:?}"
);
assert_eq!(
queue.lock().unwrap().as_slice(),
&[0],
"failed child's global ranks must be reported to the coordinator"
);
}
#[test]
fn supervise_elastic_pre_formation_kills_all() {
let fail_child = Command::new("sh")
.args(["-c", "exit 1"])
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.expect("spawn failing child");
let sleeper = Command::new("sleep")
.arg("60")
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.expect("spawn sleeper");
let children = vec![
("host-a".to_string(), 0, vec![0], fail_child, Vec::new()),
("host-b".to_string(), 1, vec![1], sleeper, Vec::new()),
];
let elastic = ElasticSupervision {
reported_deaths: std::sync::Arc::new(std::sync::Mutex::new(Vec::new())),
dead_ranks: crate::distributed::controller::DeadRanks::new(2),
max_failure: None,
world_size: 2,
cohort_formed: std::sync::Arc::new(
std::sync::atomic::AtomicBool::new(false),
),
};
let start = std::time::Instant::now();
let err = supervise_children(children, Some(elastic), None)
.expect("pre-formation failure must fail the run");
assert!(err.to_string().contains("host-a"), "got: {err}");
assert!(
start.elapsed() < std::time::Duration::from_secs(15),
"sleeper must be SIGTERMed promptly pre-formation"
);
}
#[test]
fn supervise_elastic_verdict_fails_past_threshold() {
let c0 = Command::new("sh")
.args(["-c", "exit 3"])
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.expect("spawn");
let children = vec![("host-a".to_string(), 0, vec![0], c0, Vec::new())];
let dead = crate::distributed::controller::DeadRanks::new(3);
dead.declare_dead(0);
dead.declare_dead(1);
let elastic = ElasticSupervision {
reported_deaths: std::sync::Arc::new(std::sync::Mutex::new(Vec::new())),
dead_ranks: std::sync::Arc::clone(&dead),
max_failure: Some(
crate::distributed::max_failure::MaxFailureThreshold::Absolute(2),
),
world_size: 3,
cohort_formed: std::sync::Arc::new(
std::sync::atomic::AtomicBool::new(true),
),
};
let err = supervise_children(children, Some(elastic), None)
.expect("threshold breach must fail the run");
assert!(err.to_string().contains("max_failure exceeded"), "got: {err}");
}