use std::path::PathBuf;
use super::config::{Devices, DistConfig, Rendezvous};
use super::context::DistContext;
use super::error::DistError;
pub const ENV_RANK: &str = "MAMBA_RS_RANK";
pub const ENV_WORLD: &str = "MAMBA_RS_WORLD";
pub const ENV_DEVICE: &str = "MAMBA_RS_DEVICE";
pub const ENV_RENDEZVOUS_DIR: &str = "MAMBA_RS_RENDEZVOUS_DIR";
pub const ENV_JOB_ID: &str = "MAMBA_RS_JOB_ID";
pub const ENV_SEED: &str = "MAMBA_RS_SEED";
pub enum Bootstrap {
Supervisor(SupervisorStatus),
Rank(DistContext),
}
pub struct SupervisorStatus {
pub exit_codes: Vec<i32>,
}
impl SupervisorStatus {
pub fn all_ok(&self) -> bool {
self.exit_codes.iter().all(|&c| c == 0)
}
}
struct EnvRank {
rank: usize,
world: usize,
device: usize,
}
fn read_env(k: &str) -> Option<String> {
std::env::var(k).ok().filter(|v| !v.is_empty())
}
fn parse_usize(k: &str, v: &str) -> Result<usize, DistError> {
v.parse()
.map_err(|_| DistError::EnvContract(format!("{k}={v:?} is not a number")))
}
fn env_rank() -> Result<Option<EnvRank>, DistError> {
env_rank_from(&read_env)
}
fn env_rank_from(read_env: &dyn Fn(&str) -> Option<String>) -> Result<Option<EnvRank>, DistError> {
if let Some(r) = read_env(ENV_RANK) {
let rank = parse_usize(ENV_RANK, &r)?;
let world = match read_env(ENV_WORLD) {
Some(w) => parse_usize(ENV_WORLD, &w)?,
None => {
return Err(DistError::EnvContract(format!(
"{ENV_RANK} is set but {ENV_WORLD} is missing"
)));
}
};
let device = match read_env(ENV_DEVICE) {
Some(d) => parse_usize(ENV_DEVICE, &d)?,
None => rank,
};
return Ok(Some(EnvRank {
rank,
world,
device,
}));
}
if let (Some(r), Some(w)) = (read_env("RANK"), read_env("WORLD_SIZE")) {
let rank = parse_usize("RANK", &r)?;
let world = parse_usize("WORLD_SIZE", &w)?;
let device = match read_env("LOCAL_RANK") {
Some(l) => parse_usize("LOCAL_RANK", &l)?,
None => rank,
};
return Ok(Some(EnvRank {
rank,
world,
device,
}));
}
if let (Some(r), Some(w)) = (read_env("SLURM_PROCID"), read_env("SLURM_NTASKS")) {
let rank = parse_usize("SLURM_PROCID", &r)?;
let world = parse_usize("SLURM_NTASKS", &w)?;
let device = match read_env("SLURM_LOCALID") {
Some(l) => parse_usize("SLURM_LOCALID", &l)?,
None => rank,
};
return Ok(Some(EnvRank {
rank,
world,
device,
}));
}
if let (Some(r), Some(w)) = (
read_env("OMPI_COMM_WORLD_RANK"),
read_env("OMPI_COMM_WORLD_SIZE"),
) {
let rank = parse_usize("OMPI_COMM_WORLD_RANK", &r)?;
let world = parse_usize("OMPI_COMM_WORLD_SIZE", &w)?;
let device = match read_env("OMPI_COMM_WORLD_LOCAL_RANK") {
Some(l) => parse_usize("OMPI_COMM_WORLD_LOCAL_RANK", &l)?,
None => rank,
};
return Ok(Some(EnvRank {
rank,
world,
device,
}));
}
Ok(None)
}
fn validate_job_id(job: &str) -> Result<(), DistError> {
if job.contains('/') || job.contains('\\') || job.contains("..") {
return Err(DistError::Config(format!(
"job_id {job:?} must be a plain directory segment (no separators, no ..)"
)));
}
Ok(())
}
fn rendezvous_paths(r: &Rendezvous) -> (PathBuf, String) {
match r {
Rendezvous::File { dir, job_id } => (dir.clone(), job_id.clone()),
#[allow(
unreachable_patterns,
reason = "Rendezvous is non_exhaustive for future Tcp/Preset variants; \
today File is the only one"
)]
_ => unreachable!("unhandled rendezvous variant"),
}
}
fn resolve_devices(d: &Devices) -> Result<Vec<usize>, DistError> {
match d {
Devices::Single(ord) => Ok(vec![*ord]),
Devices::Count(n) => Ok((0..*n).collect()),
Devices::List(l) => Ok(l.clone()),
Devices::All => {
#[cfg(feature = "cuda")]
{
let n = cudarc::driver::CudaContext::device_count()
.map_err(|e| DistError::Config(format!("device count query: {e:?}")))?;
if n <= 0 {
return Err(DistError::Config("no CUDA devices visible".into()));
}
Ok((0..n as usize).collect())
}
#[cfg(not(feature = "cuda"))]
{
Err(DistError::Config(
"Devices::All needs the cuda feature to enumerate devices".into(),
))
}
}
}
}
fn child_env(cfg: &DistConfig, rank: usize, world: usize, device: usize) -> Vec<(String, String)> {
let (dir, job) = rendezvous_paths(&cfg.rendezvous);
vec![
(ENV_RANK.into(), rank.to_string()),
(ENV_WORLD.into(), world.to_string()),
(ENV_DEVICE.into(), device.to_string()),
(ENV_RENDEZVOUS_DIR.into(), dir.display().to_string()),
(ENV_JOB_ID.into(), job),
(ENV_SEED.into(), cfg.seed.to_string()),
]
}
fn rank_context(cfg: &DistConfig, er: EnvRank) -> Result<DistContext, DistError> {
if er.world == 0 || er.rank >= er.world {
return Err(DistError::EnvContract(format!(
"rank {} outside world {}",
er.rank, er.world
)));
}
if er.world == 1 {
let seed = match read_env(ENV_SEED) {
Some(s) => s
.parse()
.map_err(|_| DistError::EnvContract(format!("{ENV_SEED}={s:?} is not a number")))?,
None => cfg.seed,
};
return Ok(DistContext::single(er.device, seed));
}
if let Some(w) = cfg.logical_world
&& w != er.world
{
return Err(DistError::EnvContract(format!(
"launcher world {} != configured logical_world {w} — a wrapper \
(srun/torchrun) probably split or collapsed the world",
er.world
)));
}
let (dir, job) = match (read_env(ENV_RENDEZVOUS_DIR), read_env(ENV_JOB_ID)) {
(Some(d), Some(j)) => (PathBuf::from(d), j),
(None, None) => rendezvous_paths(&cfg.rendezvous),
(Some(_), None) => {
return Err(DistError::EnvContract(format!(
"{ENV_RENDEZVOUS_DIR} is set but {ENV_JOB_ID} is missing"
)));
}
(None, Some(_)) => {
return Err(DistError::EnvContract(format!(
"{ENV_JOB_ID} is set but {ENV_RENDEZVOUS_DIR} is missing"
)));
}
};
if job.is_empty() {
return Err(DistError::Rendezvous(
"no job id: multi-process ranks must share ONE rendezvous. Under an \
external launcher set MAMBA_RS_RENDEZVOUS_DIR + MAMBA_RS_JOB_ID (or \
pass Rendezvous::File with an explicit, per-launch-unique job_id) — \
a rank-local default would put every rank in its own directory"
.into(),
));
}
let seed = match read_env(ENV_SEED) {
Some(s) => s
.parse()
.map_err(|_| DistError::EnvContract(format!("{ENV_SEED}={s:?} is not a number")))?,
None => cfg.seed,
};
validate_job_id(&job)?;
let barrier_dir = dir.join(&job);
std::fs::create_dir_all(&barrier_dir)
.map_err(|e| DistError::Rendezvous(format!("create {}: {e}", barrier_dir.display())))?;
#[cfg_attr(not(feature = "nccl"), allow(unused_mut))]
let mut ctx = DistContext::process(
er.rank,
er.world,
er.device,
seed,
cfg.reduce,
barrier_dir.clone(),
cfg.collective_timeout,
);
#[cfg(feature = "nccl")]
{
use super::comm::MambaComm;
MambaComm::preflight_version()?;
let cuda_ctx = cudarc::driver::CudaContext::new(er.device)
.map_err(|e| DistError::Transport(format!("bind device {}: {e:?}", er.device)))?;
let id_path = barrier_dir.join("nccl-id");
let join_start = std::time::Instant::now();
let id = MambaComm::exchange_unique_id(&id_path, er.rank, cfg.init_timeout)?;
let remaining = cfg
.init_timeout
.saturating_sub(join_start.elapsed())
.max(std::time::Duration::from_secs(1));
let comm = MambaComm::init_with_deadline(id, er.rank, er.world, cuda_ctx, remaining)?;
ctx.set_comm(comm);
}
Ok(ctx)
}
pub fn attach(cfg: DistConfig) -> Result<DistContext, DistError> {
cfg.validate()?;
match env_rank()? {
Some(er) => rank_context(&cfg, er),
None => Err(DistError::EnvContract(
"no rank in the environment — use bootstrap() for self-spawned runs".into(),
)),
}
}
pub fn bootstrap(cfg: DistConfig) -> Result<Bootstrap, DistError> {
cfg.validate()?;
if let Some(er) = env_rank()? {
return Ok(Bootstrap::Rank(rank_context(&cfg, er)?));
}
let devices = resolve_devices(&cfg.devices)?;
let world = cfg.logical_world.unwrap_or(devices.len());
if world < devices.len() {
return Err(DistError::Config(format!(
"logical world {world} smaller than the device list ({})",
devices.len()
)));
}
if world == 1 {
return Ok(Bootstrap::Rank(DistContext::single(devices[0], cfg.seed)));
}
if world > devices.len() {
return Err(DistError::Config(format!(
"replaying logical world {world} on {} devices is planned (the \
logical-W replay mode) but not wired yet",
devices.len()
)));
}
if cfg!(not(feature = "nccl")) {
return Err(DistError::Config(format!(
"a multi-process world (W={world}) needs the nccl feature — this \
build has no transport to reduce gradients over"
)));
}
let mut cfg = cfg;
if let Rendezvous::File { job_id, .. } = &mut cfg.rendezvous
&& job_id.is_empty()
{
*job_id = format!("job-{}", std::process::id());
}
{
let (dir, job) = rendezvous_paths(&cfg.rendezvous);
validate_job_id(&job)?;
let _ = std::fs::remove_dir_all(dir.join(job));
}
let exe =
std::env::current_exe().map_err(|e| DistError::Config(format!("current_exe: {e}")))?;
let args: Vec<std::ffi::OsString> = std::env::args_os().skip(1).collect();
let mut children: Vec<std::process::Child> = Vec::with_capacity(world);
for (rank, device) in devices.iter().enumerate() {
let mut cmd = std::process::Command::new(&exe);
cmd.args(&args);
for (k, v) in child_env(&cfg, rank, world, *device) {
cmd.env(k, v);
}
#[cfg(target_os = "linux")]
unsafe {
use std::os::unix::process::CommandExt;
cmd.pre_exec(|| {
libc::prctl(libc::PR_SET_PDEATHSIG, libc::SIGKILL);
Ok(())
});
}
match cmd.spawn() {
Ok(child) => children.push(child),
Err(e) => {
for c in &mut children {
let _ = c.kill();
let _ = c.wait();
}
return Err(DistError::Config(format!("spawn rank {rank}: {e}")));
}
}
}
let mut exit_codes: Vec<Option<i32>> = vec![None; children.len()];
loop {
let mut all_done = true;
let mut fail_fast = false;
for (rank, child) in children.iter_mut().enumerate() {
if exit_codes[rank].is_some() {
continue;
}
match child.try_wait() {
Ok(Some(status)) => {
let code = status.code().unwrap_or(-1);
exit_codes[rank] = Some(code);
if code != 0 {
fail_fast = true;
}
}
Ok(None) => all_done = false,
Err(_) => {
let _ = child.kill();
let _ = child.wait();
exit_codes[rank] = Some(-1);
fail_fast = true;
}
}
}
if fail_fast {
for (rank, child) in children.iter_mut().enumerate() {
if exit_codes[rank].is_none() {
let _ = child.kill();
}
}
}
if all_done {
break;
}
std::thread::sleep(std::time::Duration::from_millis(20));
}
let exit_codes: Vec<i32> = exit_codes.into_iter().map(|c| c.unwrap_or(-1)).collect();
Ok(Bootstrap::Supervisor(SupervisorStatus { exit_codes }))
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
const TEST_TIMEOUT: Duration = Duration::from_secs(10);
#[test]
fn child_env_carries_the_full_contract() {
let cfg = DistConfig::default()
.with_devices(Devices::Count(2))
.with_seed(7)
.with_rendezvous(Rendezvous::File {
dir: PathBuf::from("/tmp/rdzv"),
job_id: "j1".into(),
});
let env = child_env(&cfg, 1, 2, 3);
let get = |k: &str| {
env.iter()
.find(|(key, _)| key == k)
.map(|(_, v)| v.clone())
.unwrap()
};
assert_eq!(get(ENV_RANK), "1");
assert_eq!(get(ENV_WORLD), "2");
assert_eq!(get(ENV_DEVICE), "3");
assert_eq!(get(ENV_RENDEZVOUS_DIR), "/tmp/rdzv");
assert_eq!(get(ENV_JOB_ID), "j1");
assert_eq!(get(ENV_SEED), "7");
}
#[test]
fn env_rank_parses_all_launcher_conventions() {
let of = |pairs: &[(&str, &str)]| {
let owned: Vec<(String, String)> = pairs
.iter()
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect();
move |k: &str| -> Option<String> {
owned.iter().find(|(kk, _)| kk == k).map(|(_, v)| v.clone())
}
};
let er = env_rank_from(&of(&[("MAMBA_RS_RANK", "1"), ("MAMBA_RS_WORLD", "4")]))
.unwrap()
.unwrap();
assert_eq!((er.rank, er.world, er.device), (1, 4, 1));
let er = env_rank_from(&of(&[
("RANK", "3"),
("WORLD_SIZE", "4"),
("LOCAL_RANK", "1"),
]))
.unwrap()
.unwrap();
assert_eq!((er.rank, er.world, er.device), (3, 4, 1));
let er = env_rank_from(&of(&[("SLURM_PROCID", "2"), ("SLURM_NTASKS", "8")]))
.unwrap()
.unwrap();
assert_eq!((er.rank, er.world, er.device), (2, 8, 2));
let er = env_rank_from(&of(&[
("OMPI_COMM_WORLD_RANK", "0"),
("OMPI_COMM_WORLD_SIZE", "2"),
("OMPI_COMM_WORLD_LOCAL_RANK", "0"),
]))
.unwrap()
.unwrap();
assert_eq!((er.rank, er.world, er.device), (0, 2, 0));
assert!(env_rank_from(&of(&[])).unwrap().is_none());
assert!(env_rank_from(&of(&[("MAMBA_RS_RANK", "1")])).is_err());
assert!(env_rank_from(&of(&[("RANK", "x"), ("WORLD_SIZE", "2")])).is_err());
}
#[test]
fn job_id_validation_rejects_escapes() {
assert!(validate_job_id("run-42-fixedorder").is_ok());
assert!(validate_job_id("a/b").is_err());
assert!(validate_job_id("..").is_err());
assert!(validate_job_id("x..y").is_err());
assert!(validate_job_id("a\\b").is_err());
}
#[test]
fn file_barrier_two_threads_meet() {
let dir =
std::env::temp_dir().join(format!("mamba-rs-barrier-test-{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
let mk = |rank: usize| {
DistContext::process(
rank,
2,
rank,
0,
crate::dist::ReduceContract::default(),
dir.clone(),
TEST_TIMEOUT,
)
};
let a = std::thread::spawn({
let ctx = mk(0);
move || {
ctx.barrier().unwrap();
ctx.barrier().unwrap();
ctx.barrier().unwrap();
}
});
let b = std::thread::spawn({
let ctx = mk(1);
move || {
ctx.barrier().unwrap();
ctx.barrier().unwrap();
ctx.barrier().unwrap();
}
});
a.join().unwrap();
b.join().unwrap();
assert!(
!dir.join("gen-0").exists(),
"stale barrier generation must be cleaned up"
);
assert!(dir.join("gen-1").exists());
assert!(dir.join("gen-2").exists());
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn barrier_times_out_without_peers() {
let dir =
std::env::temp_dir().join(format!("mamba-rs-barrier-solo-{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
let ctx = DistContext::process(
0,
2,
0,
0,
crate::dist::ReduceContract::default(),
dir.clone(),
Duration::from_millis(50),
);
let err = ctx.barrier().unwrap_err();
assert!(matches!(err, DistError::Rendezvous(_)), "{err}");
let _ = std::fs::remove_dir_all(&dir);
}
}