use std::env;
use std::io::Write;
use std::net::TcpStream;
#[cfg(test)]
use std::net::TcpListener;
use std::time::{Duration, Instant};
use crate::{Device, Result, TensorError};
use super::launcher::FullCluster;
use super::wire::{
ControlFrame, MsgKind, RendezvousMsgWire, RendezvousRole, SessionSalt,
};
use super::{LocalCluster, NCCL_UNIQUE_ID_BYTES, NcclUniqueId, WorkerBlock};
const HOSTNAME_MAX_LEN: usize = 255;
const IO_TIMEOUT: Duration = Duration::from_secs(30);
const ENV_NCCL_SOCKET_IFNAME: &str = "NCCL_SOCKET_IFNAME";
const RENDEZVOUS_IDLE_TIMEOUT: Duration = Duration::from_secs(120);
const RENDEZVOUS_POLL_INTERVAL: Duration = Duration::from_millis(200);
const MAX_REJECTED_CONNECTIONS: usize = 1024;
#[derive(Debug)]
pub struct TcpRendezvous {
world_size: usize,
local_ranks: Vec<usize>,
local_devices: Vec<Device>,
unique_id: NcclUniqueId,
}
impl TcpRendezvous {
pub fn world_size(&self) -> usize {
self.world_size
}
pub fn local_ranks(&self) -> &[usize] {
&self.local_ranks
}
pub fn local_devices(&self) -> &[Device] {
&self.local_devices
}
pub fn unique_id(&self) -> &NcclUniqueId {
&self.unique_id
}
pub(crate) fn establish<F>(
cluster: &LocalCluster,
dataset_signature: [u8; 32],
gen_uid: F,
) -> Result<Self>
where
F: FnOnce() -> Result<NcclUniqueId>,
{
let this_host = cluster.this_worker()?;
validate_socket_ifname(cluster, this_host)?;
let local_ranks = this_host.ranks.clone();
let local_devices: Vec<Device> = this_host
.local_devices
.iter()
.map(|&d| Device::CUDA(d))
.collect();
let (my_global_rank, _) = cluster.my_rank()?;
let host_name = this_host.host.clone();
let uid_bytes = run_rank_rendezvous(
cluster,
&cluster.salt,
dataset_signature,
u32::try_from(my_global_rank).map_err(|_| {
TensorError::new(&format!(
"rendezvous: global_rank {my_global_rank} does not fit in u32"
))
})?,
&host_name,
gen_uid,
)?;
crate::msg!("cluster: {}", cluster_mapping(cluster));
Ok(TcpRendezvous {
world_size: cluster.world_size(),
local_ranks,
local_devices,
unique_id: NcclUniqueId::from_bytes(uid_bytes),
})
}
}
fn run_rank_rendezvous<F>(
cluster: &LocalCluster,
salt: &SessionSalt,
dataset_sig: [u8; 32],
global_rank: u32,
host_name: &str,
gen_uid: F,
) -> Result<[u8; NCCL_UNIQUE_ID_BYTES]>
where
F: FnOnce() -> Result<NcclUniqueId>,
{
if host_name.len() > HOSTNAME_MAX_LEN {
return Err(TensorError::new(&format!(
"rendezvous: host name {host_name:?} exceeds {HOSTNAME_MAX_LEN} bytes"
)));
}
let addr = crate::distributed::wire::join_host_port(
&cluster.controller.host,
cluster.controller.port,
);
let mut stream =
crate::distributed::wire::connect_with_retry(addr.as_str(), "rendezvous")?;
if let Ok(peer) = stream.peer_addr() {
crate::distributed::wire::warn_cleartext_public_peer("rendezvous", peer);
}
stream
.set_read_timeout(Some(IO_TIMEOUT))
.and_then(|()| stream.set_write_timeout(Some(IO_TIMEOUT)))
.map_err(|e| {
TensorError::new(&format!("rendezvous: setting timeouts failed: {e}"))
})?;
crate::distributed::wire::write_channel_magic(
&mut stream,
crate::distributed::wire::CHANNEL_MAGIC_RENDEZVOUS,
)?;
let hello = RendezvousMsgWire::Hello {
dataset_sig,
global_rank,
host_name: host_name.to_string(),
};
write_rendezvous_frame(&mut stream, salt, &hello)?;
let role = match read_rendezvous_frame(&mut stream, salt)? {
RendezvousMsgWire::Role(r) => r,
other => {
return Err(TensorError::new(&format!(
"rendezvous: expected Role frame from controller, got {other:?}"
)));
}
};
let uid_vec = match role {
RendezvousRole::Generate => {
let uid = gen_uid()?;
let bytes = *uid.as_bytes();
write_rendezvous_frame(
&mut stream,
salt,
&RendezvousMsgWire::Uid { uid_bytes: bytes.to_vec() },
)?;
bytes.to_vec()
}
RendezvousRole::Wait => match read_rendezvous_frame(&mut stream, salt)? {
RendezvousMsgWire::Uid { uid_bytes } => uid_bytes,
other => {
return Err(TensorError::new(&format!(
"rendezvous: expected Uid frame from controller, got {other:?}"
)));
}
},
};
let mut uid = [0u8; NCCL_UNIQUE_ID_BYTES];
if uid_vec.len() != NCCL_UNIQUE_ID_BYTES {
return Err(TensorError::new(&format!(
"rendezvous: UID payload length {} != {NCCL_UNIQUE_ID_BYTES}",
uid_vec.len()
)));
}
uid.copy_from_slice(&uid_vec);
Ok(uid)
}
#[cfg(test)]
pub fn run_controller_rendezvous(
full: &FullCluster,
local_host_name: &str,
) -> Result<()> {
run_controller_rendezvous_with(
full,
local_host_name,
bind_rendezvous_listener(full)?,
RENDEZVOUS_IDLE_TIMEOUT,
&std::sync::atomic::AtomicBool::new(false),
)
}
#[cfg(test)]
fn bind_rendezvous_listener(
full: &FullCluster,
) -> Result<crate::distributed::port_mux::StreamSource> {
let bind_addr = format!("0.0.0.0:{}", full.controller.port);
let listener = TcpListener::bind(&bind_addr).map_err(|e| {
TensorError::new(&format!(
"rendezvous: controller failed to bind {bind_addr}: {e}"
))
})?;
crate::distributed::port_mux::StreamSource::from_listener(listener, "rendezvous")
}
pub(crate) fn run_controller_rendezvous_aborting(
full: &FullCluster,
local_host_name: &str,
source: crate::distributed::port_mux::StreamSource,
abort: &std::sync::atomic::AtomicBool,
) -> Result<()> {
run_controller_rendezvous_with(
full,
local_host_name,
source,
RENDEZVOUS_IDLE_TIMEOUT,
abort,
)
}
fn run_controller_rendezvous_with(
full: &FullCluster,
local_host_name: &str,
source: crate::distributed::port_mux::StreamSource,
idle_timeout: Duration,
abort: &std::sync::atomic::AtomicBool,
) -> Result<()> {
let world_size = full.world_size();
if world_size == 0 {
return Err(TensorError::new(
"rendezvous: empty cluster (world_size = 0)",
));
}
let controller_at =
format!("{}:{}", full.controller.host, full.controller.port);
let designated_rank = pick_designated_rank(full, local_host_name);
eprintln!(
"cluster launcher: rendezvous server up on port {} \
(world_size={world_size}, generator=rank {designated_rank})",
full.controller.port,
);
let mut streams: Vec<Option<TcpStream>> = (0..world_size).map(|_| None).collect();
let mut reference_sig: Option<[u8; 32]> = None;
let mut accepted = 0usize;
let mut rejected = 0usize;
let mut last_progress = Instant::now();
while accepted < world_size {
let mut stream = match source.try_accept("rendezvous")? {
Some(s) => s,
None => {
if abort.load(std::sync::atomic::Ordering::SeqCst) {
return Err(TensorError::new(&format!(
"rendezvous: aborted by launcher \
({accepted}/{world_size} ranks in)"
)));
}
if last_progress.elapsed() > idle_timeout {
return Err(TensorError::new(&format!(
"rendezvous: timed out after {}s with no new rank \
connecting ({accepted}/{world_size} ranks in). Check \
that every rank process launched and can reach the \
controller at {controller_at}.",
idle_timeout.as_secs(),
)));
}
std::thread::sleep(RENDEZVOUS_POLL_INTERVAL);
continue;
}
};
let peer = stream
.peer_addr()
.map(|a| a.to_string())
.unwrap_or_else(|_| "<unknown>".to_string());
if stream
.set_read_timeout(Some(IO_TIMEOUT))
.and_then(|()| stream.set_write_timeout(Some(IO_TIMEOUT)))
.is_err()
{
rejected += 1;
if rejected > MAX_REJECTED_CONNECTIONS {
return Err(rejected_cap_error(world_size, accepted));
}
continue;
}
if let Err(e) = crate::distributed::wire::expect_channel_magic(
&mut stream,
crate::distributed::wire::CHANNEL_MAGIC_RENDEZVOUS,
"rendezvous",
) {
eprintln!(
"cluster launcher: rendezvous rejected connection from \
{peer} ({e}); continuing to accept"
);
rejected += 1;
if rejected > MAX_REJECTED_CONNECTIONS {
return Err(rejected_cap_error(world_size, accepted));
}
continue;
}
let msg = match read_rendezvous_frame(&mut stream, &full.salt) {
Ok(m) => m,
Err(e) => {
eprintln!(
"cluster launcher: rendezvous rejected connection from \
{peer} (bad frame: {e}); continuing to accept"
);
rejected += 1;
if rejected > MAX_REJECTED_CONNECTIONS {
return Err(rejected_cap_error(world_size, accepted));
}
continue;
}
};
let (dataset_sig, global_rank, host_name) = match msg {
RendezvousMsgWire::Hello { dataset_sig, global_rank, host_name } => {
(dataset_sig, global_rank, host_name)
}
other => {
eprintln!(
"cluster launcher: rendezvous rejected connection from \
{peer} (expected Hello, got {other:?}); continuing to accept"
);
rejected += 1;
if rejected > MAX_REJECTED_CONNECTIONS {
return Err(rejected_cap_error(world_size, accepted));
}
continue;
}
};
match reference_sig {
None => reference_sig = Some(dataset_sig),
Some(ref expected) if &dataset_sig != expected => {
return Err(TensorError::new(&format!(
"rendezvous: dataset_signature mismatch from host {host_name:?} \
rank {global_rank} (peer {peer}). Each rank must read from the \
same dataset; silent fan-out across diverging shards is the \
worst class of bug."
)));
}
_ => {}
}
let rank_idx = usize::try_from(global_rank).map_err(|_| {
TensorError::new(&format!(
"rendezvous: rank {global_rank} from {host_name:?} (peer {peer}) \
does not fit in usize"
))
})?;
if rank_idx >= world_size {
return Err(TensorError::new(&format!(
"rendezvous: rank {global_rank} from {host_name:?} (peer {peer}) \
out of bounds for world_size {world_size}"
)));
}
if streams[rank_idx].is_some() {
eprintln!(
"cluster launcher: rendezvous rank {global_rank} from \
{host_name:?} (peer {peer}) reconnected before role dispatch; \
replacing stale stream (transient TCP blip, not cohort-fatal)"
);
streams[rank_idx] = Some(stream);
last_progress = Instant::now();
continue;
}
streams[rank_idx] = Some(stream);
accepted += 1;
last_progress = Instant::now();
}
for (rank_idx, slot) in streams.iter_mut().enumerate() {
let stream = slot.as_mut().expect("every slot filled by accept loop");
let role = if rank_idx as u32 == designated_rank {
RendezvousRole::Generate
} else {
RendezvousRole::Wait
};
write_rendezvous_frame(stream, &full.salt, &RendezvousMsgWire::Role(role))?;
}
let designated_idx = designated_rank as usize;
let uid_bytes = {
let stream = streams[designated_idx]
.as_mut()
.expect("designated stream filled");
match read_rendezvous_frame(stream, &full.salt)? {
RendezvousMsgWire::Uid { uid_bytes } => uid_bytes,
other => {
return Err(TensorError::new(&format!(
"rendezvous: expected Uid from generator rank {designated_rank}, got {other:?}"
)));
}
}
};
if uid_bytes.len() != NCCL_UNIQUE_ID_BYTES {
return Err(TensorError::new(&format!(
"rendezvous: generator rank {designated_rank} sent UID of length {} \
(expected {NCCL_UNIQUE_ID_BYTES})",
uid_bytes.len()
)));
}
for (rank_idx, slot) in streams.iter_mut().enumerate() {
if rank_idx == designated_idx {
continue;
}
let stream = slot.as_mut().expect("every slot filled");
write_rendezvous_frame(
stream,
&full.salt,
&RendezvousMsgWire::Uid { uid_bytes: uid_bytes.clone() },
)?;
}
Ok(())
}
fn rejected_cap_error(world_size: usize, accepted: usize) -> TensorError {
TensorError::new(&format!(
"rendezvous: aborting after {MAX_REJECTED_CONNECTIONS} rejected \
pre-auth connections with only {accepted}/{world_size} ranks in. \
Something is hammering the rendezvous port (scanner, health \
checker, or a peer from another session/cluster)."
))
}
pub fn pick_designated_rank(full: &FullCluster, local_host_name: &str) -> u32 {
for worker in &full.workers {
if worker.host == local_host_name {
if let Some(&r) = worker.ranks.first() {
return r as u32;
}
}
}
full.workers
.first()
.and_then(|w| w.ranks.first())
.copied()
.map(|r| r as u32)
.unwrap_or(0)
}
fn write_rendezvous_frame<W: Write>(
w: &mut W,
salt: &SessionSalt,
msg: &RendezvousMsgWire,
) -> Result<()> {
let frame = ControlFrame::encode(salt, MsgKind::Rendezvous, msg)?;
frame.write_to(w)
}
fn read_rendezvous_frame(
stream: &mut TcpStream,
salt: &SessionSalt,
) -> Result<RendezvousMsgWire> {
let frame = ControlFrame::read_from(stream, salt)?.ok_or_else(|| {
TensorError::new("rendezvous: peer closed stream before sending frame")
})?;
if frame.kind != MsgKind::Rendezvous {
return Err(TensorError::new(&format!(
"rendezvous: expected MsgKind::Rendezvous, got {:?}",
frame.kind
)));
}
frame.decode()
}
fn validate_socket_ifname(cluster: &LocalCluster, this_host: &WorkerBlock) -> Result<()> {
if !cluster.spans_multiple_workers() {
return Ok(());
}
if env::var(ENV_NCCL_SOCKET_IFNAME).is_ok() {
return Ok(());
}
if !this_host.nccl_socket_ifname.trim().is_empty() {
unsafe {
env::set_var(ENV_NCCL_SOCKET_IFNAME, &this_host.nccl_socket_ifname);
}
return Ok(());
}
Err(TensorError::new(&format!(
"rendezvous: {ENV_NCCL_SOCKET_IFNAME} must be set when the cluster spans \
multiple hosts (auto-detection rejected -- interface naming is \
config-specific and silent fallthrough costs hours)"
)))
}
fn cluster_mapping(cluster: &LocalCluster) -> String {
let h: &WorkerBlock = &cluster.worker;
let parts: Vec<String> = h
.ranks
.iter()
.zip(h.local_devices.iter())
.map(|(r, d)| format!("{}:{} -> r{}", h.host, d, r))
.collect();
parts.join(", ")
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU16, Ordering};
use std::thread;
static ENV_MUTEX: Mutex<()> = Mutex::new(());
static NEXT_PORT: AtomicU16 = AtomicU16::new(29500);
fn next_port() -> u16 {
NEXT_PORT.fetch_add(1, Ordering::Relaxed)
}
fn slim_envelope_for(host_name: &str, port: u16) -> LocalCluster {
let (ranks, devices) = match host_name {
"host-a" => (vec![0], vec![0]),
"host-b" => (vec![1], vec![0]),
other => panic!("unknown test host {other:?}"),
};
let v = json!({
"controller": { "host": "127.0.0.1", "port": port },
"world_size": 2,
"num_workers": 2,
"worker": {
"host": host_name,
"ranks": ranks,
"local_devices": devices,
"nccl_socket_ifname": "lo",
"path": format!("/tmp/test-{host_name}"),
}
});
LocalCluster::from_value(&v).expect("test slim envelope")
}
fn full_cluster_for_test(port: u16) -> FullCluster {
let v = json!({
"controller": {
"host": "127.0.0.1",
"port": port,
"path": "/tmp/test-controller"
},
"workers": [
{
"host": "host-a",
"ranks": [0],
"local_devices": [0],
"nccl_socket_ifname": "lo",
"path": "/tmp/test-host-a",
"arch": "precompiled/cu128"
},
{
"host": "host-b",
"ranks": [1],
"local_devices": [0],
"nccl_socket_ifname": "lo",
"path": "/tmp/test-host-b",
"arch": "precompiled/cu128"
}
]
});
FullCluster::from_value(&v).expect("test full envelope")
}
#[test]
fn pick_designated_rank_prefers_local_worker() {
let full = full_cluster_for_test(next_port());
assert_eq!(pick_designated_rank(&full, "host-b"), 1);
assert_eq!(pick_designated_rank(&full, "host-a"), 0);
assert_eq!(pick_designated_rank(&full, "192.168.122.1"), 0);
}
#[test]
fn controller_rendezvous_times_out_when_a_rank_never_connects() {
let full = full_cluster_for_test(next_port());
let start = Instant::now();
let result = run_controller_rendezvous_with(
&full,
"test-controller-host",
bind_rendezvous_listener(&full).unwrap(),
Duration::from_secs(1),
&std::sync::atomic::AtomicBool::new(false),
);
let err = result.expect_err("must time out, not hang");
assert!(
err.to_string().contains("timed out"),
"unexpected error: {err}"
);
assert!(
start.elapsed() < Duration::from_secs(15),
"took too long: {:?}",
start.elapsed()
);
}
#[test]
fn controller_rendezvous_aborts_within_a_poll_interval() {
let full = full_cluster_for_test(next_port());
let abort = std::sync::atomic::AtomicBool::new(true); let start = Instant::now();
let err = run_controller_rendezvous_with(
&full,
"test-controller-host",
bind_rendezvous_listener(&full).unwrap(),
Duration::from_secs(120), &abort,
)
.expect_err("must abort, not wait for the idle ceiling");
assert!(err.to_string().contains("aborted"), "got: {err}");
assert!(
start.elapsed() < Duration::from_secs(5),
"abort took too long: {:?}",
start.elapsed()
);
}
#[test]
fn full_rendezvous_via_controller_and_two_ranks() {
let port = next_port();
let mut full = full_cluster_for_test(port);
full.salt = [0x77u8; 16];
let salt = full.salt;
let env_a = slim_envelope_for("host-a", port);
let env_b = slim_envelope_for("host-b", port);
let env_a = with_salt(env_a, salt);
let env_b = with_salt(env_b, salt);
let sig = [0x42u8; 32];
let stub_uid_bytes = [0xabu8; NCCL_UNIQUE_ID_BYTES];
let _guard = ENV_MUTEX.lock().unwrap();
let prev_ifname = env::var(ENV_NCCL_SOCKET_IFNAME).ok();
unsafe {
env::set_var(ENV_NCCL_SOCKET_IFNAME, "lo");
}
let ctrl_handle = thread::spawn(move || {
run_controller_rendezvous(&full, "test-controller-host")
});
let rank_a_handle = thread::spawn(move || {
crate::distributed::cluster::set_thread_hostname_override(Some("host-a"));
crate::distributed::cluster::set_thread_local_rank_override(Some(0));
TcpRendezvous::establish(&env_a, sig, || {
Ok(NcclUniqueId::from_bytes(stub_uid_bytes))
})
});
let rank_b_handle = thread::spawn(move || {
crate::distributed::cluster::set_thread_hostname_override(Some("host-b"));
crate::distributed::cluster::set_thread_local_rank_override(Some(0));
TcpRendezvous::establish(&env_b, sig, || {
panic!("host-b rank must not be the generator (controller picked host-a)")
})
});
let ctrl_res = ctrl_handle.join().expect("controller thread");
let rdv_a = rank_a_handle.join().expect("host-a thread").expect("host-a ok");
let rdv_b = rank_b_handle.join().expect("host-b thread").expect("host-b ok");
if let Some(v) = prev_ifname {
unsafe { env::set_var(ENV_NCCL_SOCKET_IFNAME, v); }
} else {
unsafe { env::remove_var(ENV_NCCL_SOCKET_IFNAME); }
}
ctrl_res.expect("controller rendezvous ok");
assert_eq!(rdv_a.world_size(), 2);
assert_eq!(rdv_b.world_size(), 2);
assert_eq!(rdv_a.local_ranks(), &[0usize]);
assert_eq!(rdv_b.local_ranks(), &[1usize]);
assert_eq!(rdv_a.unique_id().as_bytes(), &stub_uid_bytes);
assert_eq!(rdv_b.unique_id().as_bytes(), &stub_uid_bytes);
}
#[test]
fn duplicate_hello_reconnect_replaces_stream_not_cohort_fatal() {
let port = next_port();
let mut full = full_cluster_for_test(port);
full.salt = [0x77u8; 16];
let salt = full.salt;
let sig = [0x42u8; 32];
let stub_uid = [0xabu8; NCCL_UNIQUE_ID_BYTES];
let ctrl_handle =
thread::spawn(move || run_controller_rendezvous(&full, "test-controller-host"));
let connect = move || -> TcpStream {
let addr = format!("127.0.0.1:{port}");
for _ in 0..100 {
if let Ok(mut s) = TcpStream::connect(&addr) {
s.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
s.set_write_timeout(Some(Duration::from_secs(5))).unwrap();
crate::distributed::wire::write_channel_magic(
&mut s,
crate::distributed::wire::CHANNEL_MAGIC_RENDEZVOUS,
)
.unwrap();
return s;
}
thread::sleep(Duration::from_millis(20));
}
panic!("could not connect to controller at {addr}");
};
let hello = |rank: u32, host: &str| RendezvousMsgWire::Hello {
dataset_sig: sig,
global_rank: rank,
host_name: host.to_string(),
};
let settle = Duration::from_millis(400);
let mut conn0a = connect();
write_rendezvous_frame(&mut conn0a, &salt, &hello(0, "host-a")).unwrap();
thread::sleep(settle);
let mut conn0b = connect();
write_rendezvous_frame(&mut conn0b, &salt, &hello(0, "host-a")).unwrap();
thread::sleep(settle);
let mut conn1 = connect();
write_rendezvous_frame(&mut conn1, &salt, &hello(1, "host-b")).unwrap();
match read_rendezvous_frame(&mut conn0b, &salt).unwrap() {
RendezvousMsgWire::Role(RendezvousRole::Generate) => {}
other => panic!("rank 0 (live stream) expected Role::Generate, got {other:?}"),
}
write_rendezvous_frame(
&mut conn0b,
&salt,
&RendezvousMsgWire::Uid { uid_bytes: stub_uid.to_vec() },
)
.unwrap();
match read_rendezvous_frame(&mut conn1, &salt).unwrap() {
RendezvousMsgWire::Role(RendezvousRole::Wait) => {}
other => panic!("rank 1 expected Role::Wait, got {other:?}"),
}
match read_rendezvous_frame(&mut conn1, &salt).unwrap() {
RendezvousMsgWire::Uid { uid_bytes } => {
assert_eq!(uid_bytes, stub_uid.to_vec(), "rank 1 must receive the UID");
}
other => panic!("rank 1 expected Uid, got {other:?}"),
}
ctrl_handle
.join()
.expect("controller thread")
.expect("rendezvous must complete despite the mid-rendezvous reconnect");
drop(conn0a);
}
#[test]
fn rendezvous_rejects_signature_mismatch() {
let port = next_port();
let mut full = full_cluster_for_test(port);
full.salt = [0x55u8; 16];
let salt = full.salt;
let env_a = with_salt(slim_envelope_for("host-a", port), salt);
let env_b = with_salt(slim_envelope_for("host-b", port), salt);
let sig_a = [0x42u8; 32];
let sig_b = [0x43u8; 32]; let stub_uid = [0xabu8; NCCL_UNIQUE_ID_BYTES];
let _guard = ENV_MUTEX.lock().unwrap();
let prev_ifname = env::var(ENV_NCCL_SOCKET_IFNAME).ok();
unsafe {
env::set_var(ENV_NCCL_SOCKET_IFNAME, "lo");
}
let ctrl_handle = thread::spawn(move || {
run_controller_rendezvous(&full, "test-controller-host")
});
let rank_a_handle = thread::spawn(move || {
crate::distributed::cluster::set_thread_hostname_override(Some("host-a"));
crate::distributed::cluster::set_thread_local_rank_override(Some(0));
TcpRendezvous::establish(&env_a, sig_a, || {
Ok(NcclUniqueId::from_bytes(stub_uid))
})
});
let rank_b_handle = thread::spawn(move || {
crate::distributed::cluster::set_thread_hostname_override(Some("host-b"));
crate::distributed::cluster::set_thread_local_rank_override(Some(0));
TcpRendezvous::establish(&env_b, sig_b, || {
Ok(NcclUniqueId::from_bytes(stub_uid))
})
});
let ctrl_err = ctrl_handle.join().expect("controller thread");
let rdv_a = rank_a_handle.join().expect("host-a thread");
let rdv_b = rank_b_handle.join().expect("host-b thread");
if let Some(v) = prev_ifname {
unsafe { env::set_var(ENV_NCCL_SOCKET_IFNAME, v); }
} else {
unsafe { env::remove_var(ENV_NCCL_SOCKET_IFNAME); }
}
let err = ctrl_err.expect_err("controller must reject sig mismatch");
let msg = err.to_string();
assert!(msg.contains("dataset_signature mismatch"), "got: {msg}");
assert!(
msg.contains("host-a") || msg.contains("host-b"),
"got: {msg}"
);
assert!(rdv_a.is_err() || rdv_b.is_err(), "at least one rank must fail");
}
fn with_salt(mut c: LocalCluster, salt: SessionSalt) -> LocalCluster {
c.salt = salt;
c
}
#[test]
fn cluster_rendezvous_single_host_no_socket_ifname_required() {
let v = json!({
"controller": { "host": "127.0.0.1", "port": next_port() },
"world_size": 1,
"num_workers": 1,
"worker": {
"host": "solo", "ranks": [0], "local_devices": [0],
"nccl_socket_ifname": "lo", "path": "/tmp/test-solo"
}
});
let c = LocalCluster::from_value(&v).expect("parse");
let _guard = ENV_MUTEX.lock().unwrap();
let prev_ifname = env::var(ENV_NCCL_SOCKET_IFNAME).ok();
unsafe {
env::remove_var(ENV_NCCL_SOCKET_IFNAME);
}
let this_host = c.worker.clone();
assert!(
validate_socket_ifname(&c, &this_host).is_ok(),
"single-host must not require ifname"
);
if let Some(v) = prev_ifname {
unsafe {
env::set_var(ENV_NCCL_SOCKET_IFNAME, v);
}
}
}
#[test]
fn multi_host_auto_exports_socket_ifname_from_cluster_config() {
let cluster = slim_envelope_for("host-a", next_port());
let this_host = cluster.worker.clone();
let _guard = ENV_MUTEX.lock().unwrap();
let prev_ifname = env::var(ENV_NCCL_SOCKET_IFNAME).ok();
unsafe {
env::remove_var(ENV_NCCL_SOCKET_IFNAME);
}
let result = validate_socket_ifname(&cluster, &this_host);
let exported = env::var(ENV_NCCL_SOCKET_IFNAME).ok();
unsafe {
env::remove_var(ENV_NCCL_SOCKET_IFNAME);
if let Some(v) = prev_ifname {
env::set_var(ENV_NCCL_SOCKET_IFNAME, v);
}
}
assert!(result.is_ok(), "auto-export must succeed: {result:?}");
assert_eq!(exported.as_deref(), Some("lo"));
}
#[test]
fn multi_host_loud_error_when_cluster_config_ifname_empty() {
let v = json!({
"controller": { "host": "127.0.0.1", "port": next_port() },
"world_size": 2,
"num_workers": 2,
"worker": {
"host": "a", "ranks": [0], "local_devices": [0],
"nccl_socket_ifname": "", "path": "/tmp/test-a"
}
});
let cluster = LocalCluster::from_value(&v).expect("parse");
let this_host = cluster.worker.clone();
let _guard = ENV_MUTEX.lock().unwrap();
let prev_ifname = env::var(ENV_NCCL_SOCKET_IFNAME).ok();
unsafe {
env::remove_var(ENV_NCCL_SOCKET_IFNAME);
}
let err = validate_socket_ifname(&cluster, &this_host).expect_err("empty ifname must error");
if let Some(v) = prev_ifname {
unsafe {
env::set_var(ENV_NCCL_SOCKET_IFNAME, v);
}
}
let msg = err.to_string();
assert!(msg.contains("NCCL_SOCKET_IFNAME"), "got: {msg}");
assert!(msg.contains("multiple hosts"), "got: {msg}");
}
#[test]
fn single_host_process_per_rank_round_trip() {
let port = next_port();
let salt: SessionSalt = [0x33u8; 16];
let full = {
let v = json!({
"controller": {
"host": "127.0.0.1",
"port": port,
"path": "/tmp/test-controller"
},
"workers": [
{
"host": "single-host",
"ranks": [0, 1],
"local_devices": [0, 1],
"nccl_socket_ifname": "lo",
"path": "/tmp/test-single-host",
"arch": "precompiled/cu128"
}
]
});
let mut f = FullCluster::from_value(&v).expect("test full envelope");
f.salt = salt;
f
};
let slim = || -> LocalCluster {
let v = json!({
"controller": { "host": "127.0.0.1", "port": port },
"world_size": 2,
"num_workers": 1,
"worker": {
"host": "single-host",
"ranks": [0, 1],
"local_devices": [0, 1],
"nccl_socket_ifname": "lo",
"path": "/tmp/test-single-host",
}
});
let mut c = LocalCluster::from_value(&v).expect("single-host envelope");
c.salt = salt;
c
};
let sig = [0x42u8; 32];
let stub_uid_bytes = [0xcdu8; NCCL_UNIQUE_ID_BYTES];
let _guard = ENV_MUTEX.lock().unwrap();
let ctrl_handle = thread::spawn(move || {
run_controller_rendezvous(&full, "test-controller-host")
});
let env_0 = slim();
let rank_0 = thread::spawn(move || {
crate::distributed::cluster::set_thread_hostname_override(Some("single-host"));
crate::distributed::cluster::set_thread_local_rank_override(Some(0));
TcpRendezvous::establish(&env_0, sig, || {
Ok(NcclUniqueId::from_bytes(stub_uid_bytes))
})
});
let env_1 = slim();
let rank_1 = thread::spawn(move || {
crate::distributed::cluster::set_thread_hostname_override(Some("single-host"));
crate::distributed::cluster::set_thread_local_rank_override(Some(1));
TcpRendezvous::establish(&env_1, sig, || {
panic!("non-zero rank must not be the generator (controller picked rank 0)")
})
});
let ctrl_res = ctrl_handle.join().expect("controller thread");
let r0 = rank_0.join().expect("rank 0 thread").expect("rank 0 ok");
let r1 = rank_1.join().expect("rank 1 thread").expect("rank 1 ok");
ctrl_res.expect("controller rendezvous ok");
assert_eq!(r0.world_size(), 2);
assert_eq!(r1.world_size(), 2);
assert_eq!(r0.local_ranks(), &[0usize, 1]);
assert_eq!(r1.local_ranks(), &[0usize, 1]);
assert_eq!(r0.unique_id().as_bytes(), &stub_uid_bytes);
assert_eq!(r1.unique_id().as_bytes(), &stub_uid_bytes);
}
#[test]
fn cluster_mapping_format() {
let single_rank = json!({
"controller": { "host": "127.0.0.1", "port": 29500 },
"world_size": 3,
"num_workers": 2,
"worker": {
"host": "node-a", "ranks": [0], "local_devices": [0],
"nccl_socket_ifname": "virbr0", "path": "/tmp/test-a"
}
});
let multi_rank = json!({
"controller": { "host": "127.0.0.1", "port": 29500 },
"world_size": 3,
"num_workers": 2,
"worker": {
"host": "node-b", "ranks": [1, 2], "local_devices": [0, 1],
"nccl_socket_ifname": "enp1s0", "path": "/tmp/test-b"
}
});
let c1 = LocalCluster::from_value(&single_rank).unwrap();
let c2 = LocalCluster::from_value(&multi_rank).unwrap();
assert_eq!(cluster_mapping(&c1), "node-a:0 -> r0");
assert_eq!(cluster_mapping(&c2), "node-b:0 -> r1, node-b:1 -> r2");
}
}