use super::*;
#[test]
fn cluster_overlay_parses() {
let mut cfg: ProjectConfig =
serde_yaml_ng::from_str(canonical_cluster_yaml()).expect("parse cluster overlay");
populate_canonical_ranks(cfg.cluster.as_mut().unwrap());
let cluster = cfg.cluster.as_ref().expect("cluster: block present");
assert_eq!(cluster.controller.host, "192.168.122.1");
assert_eq!(cluster.controller.port, 29500);
assert_eq!(cluster.world_size(), 3);
assert_eq!(cluster.workers.len(), 2);
assert!(cluster.spans_multiple_hosts());
let worker = &cluster.workers[1];
assert_eq!(worker.host, "host-b");
assert_eq!(worker.ranks, vec![1, 2]);
assert_eq!(worker.local_devices, LocalDevices::Explicit(vec![0, 1]));
assert_eq!(worker.nccl_socket_ifname, "enp1s0");
assert_eq!(worker.path, "/srv/flodl");
assert_eq!(
worker.ssh.as_ref().and_then(|s| s.target.as_deref()),
Some("host-b"),
);
let test_cmd = cfg.commands.get("cuda-test").expect("cuda-test command");
assert_eq!(test_cmd.cluster, Some(true));
let train_cmd = cfg.commands.get("train").expect("train command");
assert_eq!(train_cmd.cluster, Some(true));
assert_eq!(
train_cmd.run.as_deref(),
Some("cargo run --release --bin my-training-app")
);
}
#[test]
fn cluster_block_optional() {
let yaml = "commands: { foo: { run: \"echo hi\" } }\n";
let cfg: ProjectConfig = serde_yaml_ng::from_str(yaml).expect("parse without cluster");
assert!(cfg.cluster.is_none());
assert_eq!(cfg.commands.get("foo").and_then(|c| c.cluster), None);
}
#[test]
fn cluster_overlay_merges_via_deep_merge() {
let base_yaml = "commands:\n train:\n run: cargo run --release\n";
let overlay_yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /tmp/test-solo
workers:
- host: solo
local_devices: [0]
nccl_socket_ifname: lo
path: /tmp/test-solo
commands:
train:
cluster: true
";
let base: serde_yaml_ng::Value = serde_yaml_ng::from_str(base_yaml).unwrap();
let overlay: serde_yaml_ng::Value = serde_yaml_ng::from_str(overlay_yaml).unwrap();
let merged = crate::overlay::deep_merge(base, overlay);
let merged_yaml = serde_yaml_ng::to_string(&merged).unwrap();
let mut cfg: ProjectConfig =
serde_yaml_ng::from_str(&merged_yaml).expect("merged config parses");
cfg.cluster.as_mut().unwrap().populate_ranks(&[1]).unwrap();
let cluster = cfg.cluster.as_ref().expect("cluster: from overlay");
assert_eq!(cluster.world_size(), 1);
assert_eq!(cluster.workers[0].host, "solo");
let train = cfg.commands.get("train").expect("train command");
assert_eq!(train.cluster, Some(true));
assert_eq!(train.run.as_deref(), Some("cargo run --release"));
}
#[test]
fn validate_rejects_duplicate_ranks() {
let mut cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
populate_canonical_ranks(cfg.cluster.as_mut().unwrap());
cfg.cluster.as_mut().unwrap().workers[1].ranks = vec![1, 1];
let err = cfg.cluster.as_ref().unwrap().validate().unwrap_err();
assert!(err.contains("duplicates or gaps"), "got: {err}");
}
#[test]
fn validate_rejects_rank_gap() {
let mut cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
populate_canonical_ranks(cfg.cluster.as_mut().unwrap());
cfg.cluster.as_mut().unwrap().workers[1].ranks = vec![2, 3];
let err = cfg.cluster.as_ref().unwrap().validate().unwrap_err();
assert!(err.contains("duplicates or gaps"), "got: {err}");
}
#[test]
fn validate_rejects_len_mismatch() {
let mut cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
populate_canonical_ranks(cfg.cluster.as_mut().unwrap());
cfg.cluster.as_mut().unwrap().workers[1].local_devices = LocalDevices::Explicit(vec![0]);
let err = cfg.cluster.as_ref().unwrap().validate().unwrap_err();
assert!(err.contains("length mismatch"), "got: {err}");
}
#[test]
fn validate_rejects_reserved_cluster_env_key() {
let mut cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
cfg.cluster
.as_mut()
.unwrap()
.env
.insert("CUDA_VISIBLE_DEVICES".into(), "3".into());
let err = cfg.cluster.as_ref().unwrap().validate().unwrap_err();
assert!(err.contains("reserved"), "got: {err}");
assert!(err.contains("CUDA_VISIBLE_DEVICES"), "got: {err}");
}
#[test]
fn validate_rejects_reserved_worker_env_key() {
let mut cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
cfg.cluster.as_mut().unwrap().workers[0]
.env
.insert("FLODL_INTERNAL_LOCAL_RANK".into(), "0".into());
let err = cfg.cluster.as_ref().unwrap().validate().unwrap_err();
assert!(err.contains("reserved"), "got: {err}");
}
#[test]
fn validate_allows_non_reserved_env_key() {
let mut cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
let c = cfg.cluster.as_mut().unwrap();
c.env.insert("NCCL_P2P_DISABLE".into(), "1".into());
c.env
.insert("FLODL_DASHBOARD_BIND".into(), "0.0.0.0".into());
cfg.cluster
.as_ref()
.unwrap()
.validate()
.expect("non-reserved env must pass");
}
#[test]
fn validate_rejects_empty_hosts() {
let mut cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
cfg.cluster.as_mut().unwrap().workers.clear();
let err = cfg.cluster.as_ref().unwrap().validate().unwrap_err();
assert!(err.contains("non-empty"), "got: {err}");
}
#[test]
fn discovery_allows_empty_workers_and_parses_its_knobs() {
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 1337
path: /tmp/test-discovery
join:
discovery: true
min_rank_start: 2
token: 0123456789abcdef0123456789abcdef
tunnel_only: true
workers: []
";
let cfg: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
let cluster = cfg.cluster.as_ref().unwrap();
cluster
.validate()
.expect("empty workers must pass under discovery");
let join = cluster.controller.join.as_ref().expect("join block");
assert_eq!(join.discovery, Some(true));
assert_eq!(join.min_rank_start, Some(2));
assert_eq!(
join.token.as_deref(),
Some("0123456789abcdef0123456789abcdef")
);
assert_eq!(join.tunnel_only, Some(true));
let json = serde_json::to_value(cluster).unwrap();
assert_eq!(json["controller"]["join"]["discovery"], true);
assert_eq!(json["controller"]["join"]["tunnel_only"], true);
assert_eq!(
json["controller"]["join"]["token"],
"0123456789abcdef0123456789abcdef"
);
}
#[test]
fn join_start_value_is_validated() {
let yaml_for = |start: &str| {
format!(
"\
cluster:
controller:
host: 127.0.0.1
port: 1337
path: /tmp/test-start
join:
start: {start}
workers:
- host: solo
local_devices: [0]
nccl_socket_ifname: \"\"
path: /tmp/test-start
"
)
};
for ok in ["auto", "manual", "hybrid"] {
let cfg: ProjectConfig = serde_yaml_ng::from_str(&yaml_for(ok)).unwrap();
cfg.cluster.as_ref().unwrap().validate().expect(ok);
}
let cfg: ProjectConfig = serde_yaml_ng::from_str(&yaml_for("operator")).unwrap();
let err = cfg.cluster.as_ref().unwrap().validate().unwrap_err();
assert!(err.contains("auto | manual | hybrid"), "got: {err}");
let cfg: ProjectConfig = serde_yaml_ng::from_str(&yaml_for("manual")).unwrap();
let json = serde_json::to_value(cfg.cluster.as_ref().unwrap()).unwrap();
assert_eq!(json["controller"]["join"]["start"], "manual");
}
#[test]
fn worker_join_block_parses_and_defaults() {
let cfg: ProjectConfig = serde_yaml_ng::from_str(
"\
join:
controller: 127.0.0.1:1337
ssh:
target: ctrl.example.com
user: flodl-join
identity_file: /etc/flodl/join_key
token: 0123456789abcdef0123456789abcdef
bin: target/release/train
host: worker-7
devices: [0, 1]
persist: true
args: [\"--model\", \"lenet\"]
",
)
.unwrap();
let join = cfg.join.expect("join block parsed");
assert_eq!(join.controller.as_deref(), Some("127.0.0.1:1337"));
let ssh = join.ssh.as_ref().unwrap();
assert_eq!(ssh.target.as_deref(), Some("ctrl.example.com"));
assert_eq!(ssh.user.as_deref(), Some("flodl-join"));
assert_eq!(join.bin.as_deref(), Some("target/release/train"));
assert_eq!(join.devices, Some(vec![0, 1]));
assert!(join.persist);
assert_eq!(join.args, vec!["--model".to_string(), "lenet".into()]);
let cfg: ProjectConfig = serde_yaml_ng::from_str("join:\n bin: t/bin\n").unwrap();
let join = cfg.join.expect("minimal join block");
assert_eq!(join.bin.as_deref(), Some("t/bin"));
assert!(join.controller.is_none() && join.ssh.is_none());
assert!(!join.persist && join.args.is_empty());
let err = serde_yaml_ng::from_str::<ProjectConfig>("join:\n binn: t/bin\n")
.unwrap_err()
.to_string();
assert!(err.contains("binn"), "got: {err}");
}
#[test]
fn validate_rejects_missing_socket_ifname_when_multi_host() {
let mut cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
cfg.cluster.as_mut().unwrap().workers[0].nccl_socket_ifname = String::new();
let err = cfg.cluster.as_ref().unwrap().validate().unwrap_err();
assert!(err.contains("nccl_socket_ifname"), "got: {err}");
assert!(err.contains("multiple workers"), "got: {err}");
}
#[test]
fn validate_rejects_empty_path() {
let mut cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
cfg.cluster.as_mut().unwrap().workers[0].path = String::new();
let err = cfg.cluster.as_ref().unwrap().validate().unwrap_err();
assert!(err.contains("path"), "got: {err}");
assert!(err.contains("host-a"), "got: {err}");
}
#[test]
fn validate_allows_empty_socket_ifname_when_single_host() {
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /tmp/test-solo
workers:
- host: solo
local_devices: [0]
nccl_socket_ifname: \"\"
path: /tmp/test-solo
";
let cfg: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
cfg.cluster
.as_ref()
.unwrap()
.validate()
.expect("single-host with empty ifname must pass");
}
#[test]
fn validate_passes_canonical() {
let cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
cfg.cluster
.as_ref()
.unwrap()
.validate()
.expect("canonical topology must validate");
}
#[test]
fn canonical_json_stable_and_round_trips() {
let cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
let cluster = cfg.cluster.as_ref().unwrap();
let json = cluster.canonical_json().expect("serialize");
let json2 = cluster.canonical_json().unwrap();
assert_eq!(json, json2);
let parsed: ClusterConfig = serde_json::from_str(&json).expect("round-trip");
assert_eq!(parsed.controller.host, cluster.controller.host);
assert_eq!(parsed.controller.port, cluster.controller.port);
assert_eq!(parsed.workers.len(), cluster.workers.len());
assert_eq!(parsed.world_size(), cluster.world_size());
assert_eq!(
parsed.workers[1]
.ssh
.as_ref()
.and_then(|s| s.target.as_deref()),
Some("host-b"),
);
}
#[test]
fn local_devices_all_yaml_parses_as_marker() {
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /tmp/solo
workers:
- host: solo
local_devices: all
nccl_socket_ifname: lo
path: /tmp/solo
";
let cfg: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
let host = &cfg.cluster.unwrap().workers[0];
assert_eq!(host.local_devices, LocalDevices::All);
assert!(host.local_devices.is_all());
assert!(host.local_devices.as_explicit().is_none());
}
#[test]
fn local_devices_explicit_parses_as_explicit() {
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /tmp/solo
workers:
- host: solo
local_devices: [3]
nccl_socket_ifname: lo
path: /tmp/solo
";
let cfg: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
let host = &cfg.cluster.unwrap().workers[0];
assert_eq!(host.local_devices, LocalDevices::Explicit(vec![3]));
assert!(!host.local_devices.is_all());
assert_eq!(host.local_devices.as_explicit(), Some(&[3u8][..]));
}
#[test]
fn local_devices_all_skips_length_check_in_validate() {
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /tmp/solo
workers:
- host: solo
local_devices: all
nccl_socket_ifname: lo
path: /tmp/solo
";
let cfg: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
cfg.cluster
.as_ref()
.unwrap()
.validate()
.expect("validate must pass for local_devices: all");
}
#[test]
fn local_envelope_for_all_emits_all_string() {
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /tmp/solo
workers:
- host: solo
local_devices: all
nccl_socket_ifname: lo
path: /tmp/solo
";
let cfg: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
let cluster = cfg.cluster.unwrap();
let env = cluster.local_envelope_for(&cluster.workers[0]);
assert_eq!(env["worker"]["local_devices"], serde_json::json!("all"));
}
#[test]
fn local_devices_all_round_trips_through_canonical_json() {
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /tmp/solo
workers:
- host: solo
local_devices: all
nccl_socket_ifname: lo
path: /tmp/solo
";
let cfg: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
let cluster = cfg.cluster.unwrap();
let json = cluster.canonical_json().unwrap();
let parsed: ClusterConfig = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.workers[0].local_devices, LocalDevices::All);
}
#[test]
fn local_envelope_carries_data_path_only_when_declared() {
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /tmp/solo
workers:
- host: declares
local_devices: all
nccl_socket_ifname: lo
path: /tmp/solo
data_path: /mnt/corpus
- host: silent
local_devices: all
nccl_socket_ifname: lo
path: /tmp/solo
";
let cfg: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
let cluster = cfg.cluster.unwrap();
let declares = cluster.local_envelope_for(&cluster.workers[0]);
let silent = cluster.local_envelope_for(&cluster.workers[1]);
assert_eq!(declares["worker"]["data_path"], "/mnt/corpus");
assert!(
silent["worker"].get("data_path").is_none(),
"an undeclared host must not receive the convention default: {}",
silent["worker"]
);
}
#[test]
fn local_devices_rejects_unknown_string() {
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /tmp/solo
workers:
- host: solo
local_devices: every
nccl_socket_ifname: lo
path: /tmp/solo
";
let err = serde_yaml_ng::from_str::<ProjectConfig>(yaml).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("local_devices") && msg.contains("every"),
"expected loud error mentioning the bad value, got: {msg}"
);
}
#[test]
fn local_envelope_strips_ssh_adds_world_metadata() {
let mut cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
populate_canonical_ranks(cfg.cluster.as_mut().unwrap());
let cluster = cfg.cluster.as_ref().unwrap();
let worker = &cluster.workers[1];
let env = cluster.local_envelope_for(worker);
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);
let w = &env["worker"];
assert_eq!(w["host"], "host-b");
assert_eq!(w["ranks"], serde_json::json!([1, 2]));
assert_eq!(w["local_devices"], serde_json::json!([0, 1]));
assert_eq!(w["nccl_socket_ifname"], "enp1s0");
assert_eq!(w["path"], "/srv/flodl");
assert_eq!(w["arch"], "builds/sm61-sm120");
assert!(w.get("ssh").is_none(), "ssh must not appear in envelope");
}
#[test]
fn local_envelope_omits_optional_arch() {
let cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
let mut cluster = cfg.cluster.unwrap();
cluster.workers[0].arch = None;
let env = cluster.local_envelope_for(&cluster.workers[0]);
assert!(
env["worker"].get("arch").is_none(),
"arch should be omitted when None"
);
}
#[test]
fn local_envelope_first_worker_carries_rank_zero() {
let mut cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
populate_canonical_ranks(cfg.cluster.as_mut().unwrap());
let cluster = cfg.cluster.as_ref().unwrap();
let host_a = &cluster.workers[0];
let env = cluster.local_envelope_for(host_a);
let ranks = env["worker"]["ranks"].as_array().unwrap();
assert!(
ranks.iter().any(|r| r.as_u64() == Some(0)),
"first worker's envelope must include rank 0"
);
}
#[test]
fn ssh_target_defaults_to_name() {
let cfg: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
let cluster = cfg.cluster.as_ref().unwrap();
assert_eq!(cluster.ssh_target(&cluster.workers[0]), "host-a"); assert_eq!(cluster.ssh_target(&cluster.workers[1]), "host-b"); }
#[test]
fn cluster_dispatch_empty_chain_is_false() {
assert!(!resolve_cluster_dispatch(&[]));
}
#[test]
fn cluster_dispatch_all_unset_is_false() {
assert!(!resolve_cluster_dispatch(&[None, None, None]));
}
#[test]
fn cluster_dispatch_leaf_true_wins_over_unset_ancestors() {
assert!(resolve_cluster_dispatch(&[None, Some(true)]));
}
#[test]
fn cluster_dispatch_leaf_false_wins_over_unset_ancestors() {
assert!(!resolve_cluster_dispatch(&[None, Some(false)]));
}
#[test]
fn cluster_dispatch_inherits_from_path_command_at_root() {
assert!(resolve_cluster_dispatch(&[Some(true), None]));
}
#[test]
fn cluster_dispatch_leaf_overrides_ancestor_true_to_false() {
assert!(!resolve_cluster_dispatch(&[Some(true), Some(false)]));
}
#[test]
fn cluster_dispatch_leaf_overrides_ancestor_false_to_true() {
assert!(resolve_cluster_dispatch(&[Some(false), Some(true)]));
}
#[test]
fn cluster_dispatch_intermediate_override_wins_at_sub_sub_depth() {
assert!(!resolve_cluster_dispatch(&[
Some(true),
Some(false),
None,
None
]));
}
#[test]
fn cluster_dispatch_deepest_override_wins_through_chain() {
assert!(resolve_cluster_dispatch(&[
Some(true),
Some(false),
None,
Some(true)
]));
}
#[test]
fn cluster_dispatch_enabled_requires_cluster_block() {
let no_cluster: ProjectConfig =
serde_yaml_ng::from_str("commands: { foo: { run: \"echo hi\" } }\n").unwrap();
assert!(!cluster_dispatch_enabled(&no_cluster, &[Some(true)]));
assert!(!cluster_dispatch_enabled(
&no_cluster,
&[Some(true), Some(true)]
));
}
#[test]
fn cluster_dispatch_enabled_when_block_and_chain_agree() {
let with_cluster: ProjectConfig = serde_yaml_ng::from_str(canonical_cluster_yaml()).unwrap();
assert!(cluster_dispatch_enabled(&with_cluster, &[Some(true)]));
assert!(cluster_dispatch_enabled(&with_cluster, &[Some(true), None]));
assert!(!cluster_dispatch_enabled(&with_cluster, &[Some(false)]));
assert!(!cluster_dispatch_enabled(
&with_cluster,
&[None, Some(false)]
));
assert!(!cluster_dispatch_enabled(&with_cluster, &[]));
assert!(!cluster_dispatch_enabled(&with_cluster, &[None, None]));
}