use std::path::Path;
use std::process::Command;
use crate::config::{self, ClusterConfig, ProjectConfig};
pub const ENV_FULL_CLUSTER_JSON: &str = "FLODL_INTERNAL_FULL_CLUSTER_JSON";
pub const ENV_FDL_CMD: &str = "FLODL_INTERNAL_FDL_CMD";
pub const ENV_FDL_ENV: &str = "FDL_ENV";
pub const ENV_CLUSTER_JSON: &str = "FLODL_INTERNAL_CLUSTER_JSON";
pub const ENV_CLUSTER_EXTRA_HOSTS: &str = "FLODL_INTERNAL_CLUSTER_EXTRA_HOSTS";
pub const ENV_HOST_USER: &str = "FLODL_INTERNAL_HOST_USER";
pub const ENV_HOST_OVERRIDE: &str = "FLODL_HOST_NAME";
pub const ENV_LOCAL_RANK: &str = "FLODL_INTERNAL_LOCAL_RANK";
pub fn is_reserved_cluster_env_key(key: &str) -> bool {
key.starts_with("FLODL_INTERNAL_")
|| key == "CUDA_VISIBLE_DEVICES"
|| key == "CUDA_DEVICE_ORDER"
|| key == ENV_HOST_OVERRIDE
|| key == ENV_FDL_ENV
}
pub const ENV_NET_TIMEOUT_SCALE: &str = "FLODL_NET_TIMEOUT_SCALE";
pub fn validate_net_timeout_scale() -> Result<(), String> {
validate_net_timeout_scale_value(std::env::var(ENV_NET_TIMEOUT_SCALE).ok().as_deref())
}
fn validate_net_timeout_scale_value(raw: Option<&str>) -> Result<(), String> {
let Some(raw) = raw else { return Ok(()) };
let trimmed = raw.trim();
match trimmed.parse::<f64>() {
Ok(v) if v.is_finite() && v >= 0.1 => Ok(()),
Ok(_) => Err(format!(
"{ENV_NET_TIMEOUT_SCALE}={trimmed} is out of range; expected a \
finite scale factor >= 0.1 (0.1 keeps every deadline above the \
1s heartbeat cadence)"
)),
Err(_) => Err(format!(
"{ENV_NET_TIMEOUT_SCALE}={trimmed:?} is not a number; expected a \
scale factor >= 0.1 (e.g. 3 for a slow WAN link, 0.5 for a \
fast-failure test rig)"
)),
}
}
pub fn should_dispatch(project: &ProjectConfig, chain: &[Option<bool>]) -> bool {
if is_recursive_invocation() {
return false;
}
config::cluster_dispatch_enabled(project, chain)
}
pub fn is_recursive_invocation() -> bool {
std::env::var_os(ENV_CLUSTER_JSON).is_some()
}
pub fn prepare_cluster_env(
cluster: &ClusterConfig,
overlay_env: Option<&str>,
cmd: &str,
) -> Result<Vec<String>, String> {
cluster.validate()?;
let mut warnings: Vec<String> = Vec::new();
let mut shippable = cluster.clone();
let (controller_ip, controller_warning) =
resolve_host_to_ip(&shippable.controller.host);
if let Some(ip) = controller_ip {
shippable.controller.host = ip;
}
if let Some(w) = controller_warning {
warnings.push(w);
}
let counts = probe_worker_device_counts(&shippable)?;
shippable.populate_ranks(&counts)?;
shippable.validate()?;
let json = shippable.canonical_json()?;
let hex = hex_encode(json.as_bytes());
let (extra_hosts, host_warnings) = resolve_cluster_extra_hosts(cluster);
warnings.extend(host_warnings);
let host_user = resolve_local_user();
if host_user.is_none() {
warnings.push(
"could not determine the controller's user (USER unset, whoami \
unavailable); ssh will use its own defaults — set `ssh.user:` \
per host in fdl.cluster.yml if remote accounts differ"
.to_string(),
);
}
unsafe {
std::env::set_var(ENV_FULL_CLUSTER_JSON, &hex);
std::env::set_var(ENV_FDL_CMD, cmd);
if let Some(u) = &host_user {
std::env::set_var(ENV_HOST_USER, u);
}
if !extra_hosts.is_empty() {
std::env::set_var(ENV_CLUSTER_EXTRA_HOSTS, extra_hosts.join(" "));
}
if let Some(e) = overlay_env {
if !e.trim().is_empty() {
std::env::set_var(ENV_FDL_ENV, e);
}
}
}
Ok(warnings)
}
fn probe_local_device_counts(cluster: &ClusterConfig) -> Result<Vec<usize>, String> {
let mut counts = Vec::with_capacity(cluster.workers.len());
let mut cached_local: Option<usize> = None;
for (i, w) in cluster.workers.iter().enumerate() {
let count = match &w.local_devices {
config::LocalDevices::Explicit(v) => v.len(),
config::LocalDevices::All => {
if cached_local.is_none() {
cached_local = Some(
crate::gpus::count_visible_gpus_via_nvidia_smi().map_err(|e| {
format!(
"cluster.workers[{i}] ({:?}): local nvidia-smi \
probe failed: {e}",
w.host,
)
})?,
);
}
cached_local.unwrap()
}
};
if count == 0 {
return Err(format!(
"cluster.workers[{i}] ({:?}): 0 CUDA devices visible \
(local_devices: all). Run `fdl @cluster-test <cmd>` on a \
host with visible GPUs, or use an explicit \
`local_devices: [...]` list.",
w.host,
));
}
counts.push(count);
}
Ok(counts)
}
pub fn prepare_test_cluster_env(cluster: &ClusterConfig) -> Result<String, String> {
cluster.validate()?;
let mut shippable = cluster.clone();
let counts = probe_local_device_counts(&shippable)?;
shippable.populate_ranks(&counts)?;
shippable.validate()?;
let json = shippable.canonical_json()?;
Ok(hex_encode(json.as_bytes()))
}
fn probe_worker_device_counts(cluster: &ClusterConfig) -> Result<Vec<usize>, String> {
let mut counts = Vec::with_capacity(cluster.workers.len());
for (i, w) in cluster.workers.iter().enumerate() {
let count = match &w.local_devices {
config::LocalDevices::Explicit(v) => v.len(),
config::LocalDevices::All => ssh_query_gpu_count(w).map_err(|e| {
format!(
"cluster.workers[{i}] ({:?}): probe failed: {e}",
w.host,
)
})?,
};
if count == 0 {
return Err(format!(
"cluster.workers[{i}] ({:?}): probed 0 CUDA devices \
(local_devices: all). Either the host has no GPUs visible \
(check nvidia-smi + CUDA_VISIBLE_DEVICES) or it's a \
misconfiguration — provide an explicit `local_devices: [...]` \
list instead.",
w.host,
));
}
counts.push(count);
}
Ok(counts)
}
pub(crate) fn apply_worker_ssh_opts(cmd: &mut Command, worker: &config::ClusterWorker) {
if let Some(ssh) = worker.ssh.as_ref() {
if let Some(port) = ssh.port {
cmd.arg("-p").arg(port.to_string());
}
if let Some(user) = ssh.user.as_deref() {
cmd.arg("-l").arg(user);
}
if let Some(id) = ssh.identity_file.as_deref() {
cmd.arg("-i").arg(id);
}
if let Some(warning) = batchmode_override_warning(&ssh.options, &worker.host) {
eprintln!("{warning}");
}
for opt in &ssh.options {
cmd.arg("-o").arg(opt);
}
}
}
fn batchmode_override_warning(opts: &[String], host: &str) -> Option<String> {
opts.iter().find_map(|opt| {
let (k, v) = opt.split_once('=')?;
(k.trim().eq_ignore_ascii_case("BatchMode")
&& !v.trim().eq_ignore_ascii_case("yes"))
.then(|| {
format!(
"fdl: host {host:?} ssh.options set `{}` — flodl's ssh is \
non-interactive and will hang on any prompt (passphrase, \
host-key). Proceeding as requested.",
opt.trim()
)
})
})
}
fn ssh_query_gpu_count(worker: &config::ClusterWorker) -> Result<usize, String> {
let target = worker
.ssh
.as_ref()
.and_then(|s| s.target.as_deref())
.unwrap_or(&worker.host);
let mut cmd = Command::new("ssh");
apply_worker_ssh_opts(&mut cmd, worker);
cmd.args([
"-T",
"-o",
"BatchMode=yes",
"-o",
"ConnectTimeout=5",
]);
cmd.arg(target);
cmd.arg("nvidia-smi --query-gpu=index --format=csv,noheader 2>/dev/null | wc -l");
let output = cmd
.output()
.map_err(|e| format!("ssh spawn failed: {e}"))?;
if !output.status.success() {
let stderr = String::from_utf8_lossy(&output.stderr).trim().to_string();
return Err(format!(
"ssh to {target:?} exited {} (stderr: {stderr})",
output.status,
));
}
let stdout = String::from_utf8_lossy(&output.stdout);
let count_str = stdout.trim();
count_str.parse::<usize>().map_err(|e| {
format!(
"could not parse nvidia-smi output as device count: {e:?} \
(got {count_str:?})"
)
})
}
fn resolve_cluster_extra_hosts(cluster: &ClusterConfig) -> (Vec<String>, Vec<String>) {
let mut hosts = Vec::new();
let mut warnings = Vec::new();
for w in &cluster.workers {
let (ip, warning) = resolve_host_to_ip(&w.host);
if let Some(ip) = ip {
hosts.push(format!("{}:{ip}", w.host));
}
let has_explicit_ssh_target = w
.ssh
.as_ref()
.and_then(|s| s.target.as_deref())
.is_some();
if let Some(msg) = warning
&& !has_explicit_ssh_target
{
warnings.push(msg);
}
}
(hosts, warnings)
}
fn resolve_host_to_ip(host: &str) -> (Option<String>, Option<String>) {
use std::net::ToSocketAddrs;
if host.parse::<std::net::IpAddr>().is_ok() {
return (Some(host.to_string()), None);
}
match (host, 0u16).to_socket_addrs() {
Ok(iter) => {
let (ip, only_loopback) = select_preferred_ip(iter.map(|sa| sa.ip()));
let warning = match (&ip, only_loopback) {
(Some(ip), true) => Some(format!(
"host {host:?} only resolves to loopback {ip} on the controller \
— remote ranks will fail to connect. Set the controller's `host:` \
in fdl.cluster.yml to a non-loopback IP reachable from peer nodes \
(e.g. the libvirt bridge IP 192.168.122.1 for virbr0). An explicit \
IP is used as-is (it is the address ranks dial), bypassing this \
hostname resolution; on a multi-NIC controller this is also how you \
pin which address is advertised."
)),
_ => None,
};
(ip, warning)
}
Err(e) => (
None,
Some(format!(
"host {host:?} did not resolve on controller: {e} (remote ranks \
will retry via their own NSS — fix host-side resolution if they \
also fail)"
)),
),
}
}
fn select_preferred_ip<I: IntoIterator<Item = std::net::IpAddr>>(
iter: I,
) -> (Option<String>, bool) {
let mut loopback_fallback: Option<String> = None;
for ip in iter {
if !ip.is_loopback() {
return (Some(ip.to_string()), false);
}
loopback_fallback.get_or_insert_with(|| ip.to_string());
}
let only_loopback = loopback_fallback.is_some();
(loopback_fallback, only_loopback)
}
pub fn cluster_compose_overlay_arg(project_root: &Path) -> String {
let raw = match std::env::var(ENV_CLUSTER_EXTRA_HOSTS) {
Ok(s) => s,
Err(_) => return String::new(),
};
let pairs: Vec<&str> = raw.split_whitespace().filter(|p| !p.is_empty()).collect();
if pairs.is_empty() {
return String::new();
}
let mut entries = String::new();
for pair in &pairs {
entries.push_str(" - \"");
for ch in pair.chars() {
match ch {
'\\' => entries.push_str("\\\\"),
'"' => entries.push_str("\\\""),
_ => entries.push(ch),
}
}
entries.push_str("\"\n");
}
let overlay = format!(
"# Generated by fdl-cli (cluster mode) — DO NOT EDIT BY HAND.\n\
# Regenerated on every `fdl @cluster ...` invocation.\n\
services:\n\
\x20\x20cuda:\n\
\x20\x20\x20\x20extra_hosts:\n{entries}\
\x20\x20dev:\n\
\x20\x20\x20\x20extra_hosts:\n{entries}\
\x20\x20bench:\n\
\x20\x20\x20\x20extra_hosts:\n{entries}",
);
let overlay_path = project_root.join(".fdl-cluster-overlay.yml");
if let Err(e) = std::fs::write(&overlay_path, overlay) {
eprintln!(
"fdl: warning: failed to write cluster compose overlay at {:?}: {e} \
(continuing without --add-host injection — remote hostnames \
may not resolve inside the container)",
overlay_path
);
return String::new();
}
format!(
" -f docker-compose.yml -f {}",
crate::util::shell::posix_quote(&overlay_path.display().to_string())
)
}
pub fn hex_encode(bytes: &[u8]) -> String {
const TABLE: &[u8; 16] = b"0123456789abcdef";
let mut s = String::with_capacity(bytes.len() * 2);
for &b in bytes {
s.push(TABLE[(b >> 4) as usize] as char);
s.push(TABLE[(b & 0x0F) as usize] as char);
}
s
}
pub fn resolve_local_user() -> Option<String> {
resolve_user_from(
std::env::var("USER").ok().as_deref(),
Command::new("whoami").output().ok().and_then(|out| {
if out.status.success() {
String::from_utf8(out.stdout).ok()
} else {
None
}
}),
)
}
fn resolve_user_from(user_env: Option<&str>, whoami_out: Option<String>) -> Option<String> {
if let Some(s) = user_env {
let s = s.trim();
if !s.is_empty() {
return Some(s.to_string());
}
}
whoami_out
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
}
pub fn resolve_local_hostname() -> String {
if let Ok(s) = std::env::var(ENV_HOST_OVERRIDE) {
let s = s.trim().to_string();
if !s.is_empty() {
return s;
}
}
Command::new("hostname")
.output()
.ok()
.and_then(|out| {
if out.status.success() {
String::from_utf8(out.stdout)
.ok()
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
} else {
None
}
})
.unwrap_or_else(|| "unknown-host".to_string())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::util::test_env::env_lock;
#[test]
fn net_timeout_scale_validation_mirrors_library_rule() {
assert!(validate_net_timeout_scale_value(None).is_ok());
assert!(validate_net_timeout_scale_value(Some("3")).is_ok());
assert!(validate_net_timeout_scale_value(Some("0.1")).is_ok());
assert!(validate_net_timeout_scale_value(Some(" 2.0 ")).is_ok());
assert!(validate_net_timeout_scale_value(Some("0.05")).is_err());
assert!(validate_net_timeout_scale_value(Some("-1")).is_err());
assert!(validate_net_timeout_scale_value(Some("inf")).is_err());
assert!(validate_net_timeout_scale_value(Some("abc")).is_err());
}
#[test]
fn resolve_user_never_fabricates() {
assert_eq!(resolve_user_from(Some("fab"), None).as_deref(), Some("fab"));
assert_eq!(resolve_user_from(Some(" fab \n"), None).as_deref(), Some("fab"));
assert_eq!(
resolve_user_from(Some(""), Some("who\n".into())).as_deref(),
Some("who"),
);
assert_eq!(resolve_user_from(None, Some("who".into())).as_deref(), Some("who"));
assert_eq!(resolve_user_from(None, None), None);
assert_eq!(resolve_user_from(Some(" "), Some(" ".into())), None);
}
#[test]
fn should_dispatch_returns_false_when_cluster_json_set() {
let _guard = env_lock();
unsafe {
std::env::set_var(ENV_CLUSTER_JSON, "deadbeef");
}
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /opt/flodl
workers:
- host: solo
local_devices: [0]
nccl_socket_ifname: lo
path: /opt/flodl
commands:
x: { cluster: true, run: \"echo hi\" }
";
let project: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
assert!(
!should_dispatch(&project, &[Some(true)]),
"recursion guard: must return false when FLODL_INTERNAL_CLUSTER_JSON is set"
);
unsafe {
std::env::remove_var(ENV_CLUSTER_JSON);
}
}
#[test]
fn should_dispatch_delegates_when_env_unset() {
let _guard = env_lock();
unsafe {
std::env::remove_var(ENV_CLUSTER_JSON);
}
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /opt/flodl
workers:
- host: solo
local_devices: [0]
nccl_socket_ifname: lo
path: /opt/flodl
commands:
x: { run: \"echo hi\" }
";
let project: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
assert!(!should_dispatch(&project, &[None]));
assert!(should_dispatch(&project, &[Some(true)]));
}
#[test]
fn hex_encode_matches_library() {
assert_eq!(hex_encode(b""), "");
assert_eq!(hex_encode(&[0x00]), "00");
assert_eq!(hex_encode(&[0xff]), "ff");
assert_eq!(hex_encode(&[0x0f, 0xa0]), "0fa0");
assert_eq!(hex_encode(b"hi"), "6869");
}
#[test]
fn prepare_cluster_env_sets_required_vars() {
let _guard = env_lock();
unsafe {
std::env::remove_var(ENV_FULL_CLUSTER_JSON);
std::env::remove_var(ENV_FDL_CMD);
std::env::remove_var(ENV_FDL_ENV);
}
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /opt/flodl
workers:
- host: solo
local_devices: [0]
nccl_socket_ifname: lo
path: /opt/flodl
commands:
train: { cluster: true, run: \"true\" }
";
let project: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
let cluster = project.cluster.as_ref().unwrap();
prepare_cluster_env(cluster, Some("cluster"), "train").expect("prepare OK");
assert!(!std::env::var(ENV_FULL_CLUSTER_JSON).unwrap().is_empty());
assert_eq!(std::env::var(ENV_FDL_CMD).unwrap(), "train");
assert_eq!(std::env::var(ENV_FDL_ENV).unwrap(), "cluster");
let hex = std::env::var(ENV_FULL_CLUSTER_JSON).unwrap();
assert!(hex.chars().all(|c| c.is_ascii_hexdigit()));
unsafe {
std::env::remove_var(ENV_FULL_CLUSTER_JSON);
std::env::remove_var(ENV_FDL_CMD);
std::env::remove_var(ENV_FDL_ENV);
}
}
#[test]
fn prepare_cluster_env_skips_fdl_env_when_blank() {
let _guard = env_lock();
unsafe {
std::env::remove_var(ENV_FDL_ENV);
}
let yaml = "\
cluster:
controller:
host: 127.0.0.1
port: 29500
path: /opt/flodl
workers:
- host: solo
local_devices: [0]
nccl_socket_ifname: lo
path: /opt/flodl
commands:
train: { cluster: true, run: \"true\" }
";
let project: ProjectConfig = serde_yaml_ng::from_str(yaml).unwrap();
let cluster = project.cluster.as_ref().unwrap();
prepare_cluster_env(cluster, None, "train").unwrap();
assert!(std::env::var_os(ENV_FDL_ENV).is_none());
prepare_cluster_env(cluster, Some(" "), "train").unwrap();
assert!(std::env::var_os(ENV_FDL_ENV).is_none());
unsafe {
std::env::remove_var(ENV_FULL_CLUSTER_JSON);
std::env::remove_var(ENV_FDL_CMD);
}
}
#[test]
fn select_preferred_ip_prefers_non_loopback() {
use std::net::IpAddr;
let ips: Vec<IpAddr> = vec![
"127.0.1.1".parse().unwrap(),
"192.168.122.1".parse().unwrap(),
];
let (ip, only_loopback) = select_preferred_ip(ips);
assert_eq!(ip.as_deref(), Some("192.168.122.1"));
assert!(!only_loopback);
}
#[test]
fn select_preferred_ip_falls_back_to_loopback_with_flag() {
use std::net::IpAddr;
let ips: Vec<IpAddr> = vec!["127.0.1.1".parse().unwrap(), "::1".parse().unwrap()];
let (ip, only_loopback) = select_preferred_ip(ips);
assert_eq!(ip.as_deref(), Some("127.0.1.1"));
assert!(only_loopback);
}
#[test]
fn select_preferred_ip_empty_iterator() {
let (ip, only_loopback) = select_preferred_ip(std::iter::empty());
assert!(ip.is_none());
assert!(!only_loopback);
}
#[test]
fn select_preferred_ip_skips_ipv6_loopback() {
use std::net::IpAddr;
let ips: Vec<IpAddr> = vec!["::1".parse().unwrap(), "10.0.0.5".parse().unwrap()];
let (ip, only_loopback) = select_preferred_ip(ips);
assert_eq!(ip.as_deref(), Some("10.0.0.5"));
assert!(!only_loopback);
}
#[test]
fn prepare_cluster_env_validates_cluster() {
let _guard = env_lock();
let cluster = ClusterConfig {
controller: crate::config::ClusterController {
host: String::new(),
port: 1337,
path: String::new(),
docker: None,
arch: None,
data_path: None,
join: None,
},
workers: Vec::new(),
env: std::collections::BTreeMap::new(),
};
let err = prepare_cluster_env(&cluster, None, "train").unwrap_err();
assert!(err.contains("controller.host"), "got: {err}");
}
#[test]
fn resolve_cluster_extra_hosts_suppresses_warning_when_ssh_target_explicit() {
use crate::config::{
ClusterController, ClusterWorker, LocalDevices, SshConfig,
};
let cluster = ClusterConfig {
controller: ClusterController {
host: "127.0.0.1".into(),
port: 1337,
path: "/tmp".into(),
docker: None,
arch: None,
data_path: None,
join: None,
},
workers: vec![ClusterWorker {
host: "nonexistent.invalid.".into(),
ranks: vec![0],
local_devices: LocalDevices::Explicit(vec![0]),
nccl_socket_ifname: "lo".into(),
path: "/tmp".into(),
ssh: Some(SshConfig {
target: Some("127.0.0.1".into()),
..SshConfig::default()
}),
tunnel: false,
arch: None,
data_path: None,
docker: None,
env: std::collections::BTreeMap::new(),
}],
env: std::collections::BTreeMap::new(),
};
let (_hosts, warnings) = resolve_cluster_extra_hosts(&cluster);
assert!(
warnings.is_empty(),
"explicit ssh.target should suppress the resolution warning, \
got warnings: {warnings:?}"
);
}
#[test]
fn resolve_cluster_extra_hosts_warns_when_ssh_target_absent() {
use crate::config::{ClusterController, ClusterWorker, LocalDevices};
let cluster = ClusterConfig {
controller: ClusterController {
host: "127.0.0.1".into(),
port: 1337,
path: "/tmp".into(),
docker: None,
arch: None,
data_path: None,
join: None,
},
workers: vec![ClusterWorker {
host: "nonexistent.invalid.".into(),
ranks: vec![0],
local_devices: LocalDevices::Explicit(vec![0]),
nccl_socket_ifname: "lo".into(),
path: "/tmp".into(),
ssh: None,
tunnel: false,
arch: None,
data_path: None,
docker: None,
env: std::collections::BTreeMap::new(),
}],
env: std::collections::BTreeMap::new(),
};
let (_hosts, warnings) = resolve_cluster_extra_hosts(&cluster);
assert!(
warnings.iter().any(|w| w.contains("did not resolve")),
"missing ssh.target should keep the warning, got warnings: {warnings:?}"
);
}
}