use std::net::TcpStream;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{Duration, Instant};
use serde::Serialize;
use crate::distributed::port_mux::StreamSource;
use crate::distributed::wire::{
CHANNEL_MAGIC_JOIN, ControlFrame, JoinMsgWire, MsgKind, SESSION_SALT_BYTES, SessionSalt,
expect_channel_magic, salt_to_hex, scaled_deadline_secs,
};
use crate::tensor::{Result, TensorError};
const ACCEPT_POLL: Duration = Duration::from_millis(20);
const JOIN_HANDSHAKE_TIMEOUT_SECS: u64 = 10;
const MAX_REJECTED_JOINS: usize = 1024;
const MAX_RANKS_PER_JOIN: usize = 256;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum StartMode {
#[default]
Auto,
Manual,
Hybrid,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct JoinConfig {
pub min_rank_start: usize,
pub join_timeout_secs: u64,
pub target_ranks: Option<usize>,
pub max_join_timeout_secs: u64,
pub open_admission: bool,
pub start_mode: StartMode,
pub nccl_backend: bool,
}
impl Default for JoinConfig {
fn default() -> Self {
JoinConfig {
min_rank_start: 1,
join_timeout_secs: 300,
target_ranks: None,
max_join_timeout_secs: 600,
open_admission: false,
start_mode: StartMode::Auto,
nccl_backend: true,
}
}
}
impl JoinConfig {
pub fn validate(&self) -> Result<()> {
if self.min_rank_start == 0 {
return Err(TensorError::new(
"cluster join: min_rank_start must be >= 1 (a world of zero \
ranks cannot train)",
));
}
if let Some(target) = self.target_ranks
&& target < self.min_rank_start
{
return Err(TensorError::new(&format!(
"cluster join: target_ranks ({target}) must be >= \
min_rank_start ({}) — the early-close target cannot sit \
below the quorum",
self.min_rank_start,
)));
}
if self.max_join_timeout_secs < self.join_timeout_secs {
return Err(TensorError::new(&format!(
"cluster join: max_join_timeout ({}s) must be >= join_timeout \
({}s) — the hard cap cannot expire before the window",
self.max_join_timeout_secs, self.join_timeout_secs,
)));
}
if self.start_mode == StartMode::Manual && self.target_ranks.is_some() {
return Err(TensorError::new(
"cluster join: `start: manual` and `target_ranks` contradict \
— the target is a clock-side auto-close and manual mode \
hands the close to the operator; drop one of them (hybrid \
keeps both)",
));
}
Ok(())
}
}
pub(crate) fn join_frame_key(open_admission: bool, salt: &SessionSalt) -> SessionSalt {
if open_admission {
[0u8; SESSION_SALT_BYTES]
} else {
*salt
}
}
pub(crate) fn resolve_open_admission(config: &JoinConfig, bind_is_loopback: bool) -> bool {
if bind_is_loopback {
return true;
}
if config.open_admission {
eprintln!(
"flodl: WARNING: open_admission is enabled on a NON-loopback \
controller bind — any peer that can reach the join port can \
join (and therefore influence) this training run. Sound only \
on a fully trusted network segment; prefer tunneled workers \
(`tunnel: true`), which flip the bind to loopback and make \
reachability itself the authentication."
);
return true;
}
false
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[non_exhaustive]
#[serde(rename_all = "snake_case")]
pub enum ClusterPhase {
Waiting,
Staging,
Forming,
Training,
Done,
Failed,
}
#[derive(Debug, Clone)]
pub struct JoinedMember {
pub host: String,
pub ranks: Vec<usize>,
pub local_devices: Vec<u8>,
pub gpus: Vec<String>,
pub libtorch: String,
pub joined_at_secs: u64,
}
#[derive(Debug, Clone, Serialize)]
pub struct MembershipSnapshot {
pub phase: ClusterPhase,
pub joined_ranks: usize,
pub joined_hosts: usize,
pub min_rank_start: usize,
pub target_ranks: Option<usize>,
pub window_remaining_secs: Option<u64>,
pub cap_remaining_secs: Option<u64>,
pub start_mode: StartMode,
pub start_armed: bool,
pub members: Vec<MemberSnapshot>,
}
#[derive(Debug, Clone, Serialize)]
pub struct MemberSnapshot {
pub host: String,
pub ranks: Vec<usize>,
pub local_devices: Vec<u8>,
pub gpus: Vec<String>,
pub libtorch: String,
pub joined_at_secs: u64,
}
#[derive(Debug, PartialEq, Eq)]
pub(crate) enum WindowVerdict {
Open,
Formed(&'static str),
Failed(String),
}
#[derive(Debug)]
pub(crate) struct JoinOffer {
pub host: String,
pub local_devices: Vec<u8>,
pub gpus: Vec<String>,
pub libtorch: String,
pub dataset_sig: [u8; 32],
pub run_id: Option<String>,
pub nccl_version: Option<(u32, u32, u32)>,
pub model_sig: Option<[u8; 32]>,
}
#[derive(Debug)]
pub(crate) struct MembershipLedger {
config: JoinConfig,
expected_dataset_sig: Option<[u8; 32]>,
expected_vendor: Option<flodl_hw::GpuVendor>,
expected_run_id: Option<String>,
expected_nccl: Option<(u32, u32)>,
expected_model_sig: Option<[u8; 32]>,
members: Vec<JoinedMember>,
next_rank: usize,
}
impl MembershipLedger {
pub fn new(
config: JoinConfig,
expected_dataset_sig: Option<[u8; 32]>,
expected_model_sig: Option<[u8; 32]>,
) -> Result<Self> {
config.validate()?;
Ok(MembershipLedger {
config,
expected_dataset_sig,
expected_vendor: None,
expected_run_id: None,
expected_nccl: None,
expected_model_sig,
members: Vec::new(),
next_rank: 0,
})
}
pub fn joined_ranks(&self) -> usize {
self.next_rank
}
pub fn admit(
&mut self,
offer: JoinOffer,
elapsed: Duration,
) -> std::result::Result<Vec<usize>, String> {
let JoinOffer {
host,
local_devices,
gpus,
libtorch,
dataset_sig,
run_id,
nccl_version,
model_sig,
} = offer;
let host = host.as_str();
if host.trim().is_empty() {
return Err("host name must be non-empty".to_string());
}
let rank_count = local_devices.len();
if rank_count == 0 {
return Err("local_devices must be non-empty (a worker with no \
ranks cannot train)"
.to_string());
}
if rank_count > MAX_RANKS_PER_JOIN {
return Err(format!(
"{rank_count} local devices exceeds the per-worker cap \
{MAX_RANKS_PER_JOIN}"
));
}
{
let mut seen = [false; 256];
for d in &local_devices {
if std::mem::replace(&mut seen[*d as usize], true) {
return Err(format!(
"duplicate local device {d} — each rank must pin a \
distinct physical GPU"
));
}
}
}
if self.members.iter().any(|m| m.host == host) {
return Err(format!(
"host {host:?} already joined this run (duplicate join — a \
stale worker from a previous launch, or two workers sharing \
one host name)"
));
}
match self.expected_dataset_sig {
None => self.expected_dataset_sig = Some(dataset_sig),
Some(ref expected) if expected != &dataset_sig => {
return Err(format!(
"dataset signature mismatch (worker {}…, run {}…) — the \
worker was built against a different dataset shard \
layout",
hex_prefix(&dataset_sig),
hex_prefix(expected),
));
}
Some(_) => {}
}
if self.config.nccl_backend
&& let flodl_hw::VariantClass::Vendor(vendor) =
flodl_hw::classify_variant_label(&libtorch)
{
match self.expected_vendor {
None => self.expected_vendor = Some(vendor),
Some(expected) if expected != vendor => {
return Err(format!(
"GPU vendor mismatch: this run's cohort is {expected} \
(libtorch {:?}) and {host:?} offers {vendor} \
({libtorch:?}) — NCCL and RCCL cannot form one \
communicator, so a mixed cohort hangs at formation. \
Use a CPU ElChe mode (cpu_sync / cpu_cadence / \
cpu_async), or a one-vendor fleet",
self.members
.iter()
.map(|m| m.libtorch.as_str())
.find(|l| {
matches!(
flodl_hw::classify_variant_label(l),
flodl_hw::VariantClass::Vendor(v) if v == expected
)
})
.unwrap_or(""),
));
}
Some(_) => {}
}
}
if let Some(run) = &run_id {
match &self.expected_run_id {
None => self.expected_run_id = Some(run.clone()),
Some(expected) if expected != run => {
return Err(format!(
"run identity mismatch: this cohort prepared run \
{}… and {host:?} prepared {}… — a publish landed \
between their fetches, so they hold two different \
runs. The stale side picks the new run up on its \
next dial",
id_prefix(expected),
id_prefix(run),
));
}
Some(_) => {}
}
}
if self.config.nccl_backend
&& let Some((maj, min, _)) = nccl_version
{
match self.expected_nccl {
None => self.expected_nccl = Some((maj, min)),
Some((emaj, emin)) if (emaj, emin) != (maj, min) => {
return Err(format!(
"NCCL version skew: this cohort loads \
{emaj}.{emin}.x and {host:?} loads {maj}.{min}.x \
— NCCL refuses its handshake across major.minor \
skew, at formation. Align the libtorch variants, \
or bridge with `fdl nccl build`",
));
}
Some(_) => {}
}
}
if let Some(sig) = &model_sig {
match &self.expected_model_sig {
None => self.expected_model_sig = Some(*sig),
Some(expected) if expected != sig => {
return Err(format!(
"model mismatch: the model this box builds ({}…) \
differs from the cohort's ({}…) — parameter names, \
shapes or dtypes disagree, so the boxes would \
corrupt each other at the first averaging step. A \
stale source tree or a wrong `bin:` is the usual \
cause",
hex_prefix(sig),
hex_prefix(expected),
));
}
Some(_) => {}
}
}
let ranks: Vec<usize> = (self.next_rank..self.next_rank + rank_count).collect();
self.next_rank += rank_count;
self.members.push(JoinedMember {
host: host.to_string(),
ranks: ranks.clone(),
local_devices,
gpus,
libtorch,
joined_at_secs: elapsed.as_secs(),
});
Ok(ranks)
}
pub fn retract_last(&mut self, host: &str) -> Result<()> {
match self.members.last() {
Some(m) if m.host == host => {
let m = self.members.pop().expect("last() was Some");
self.next_rank -= m.ranks.len();
Ok(())
}
_ => Err(TensorError::new(&format!(
"cluster join: retract_last({host:?}) is not the most recent \
admission — rank ids are assigned contiguously and only the \
tail can be returned"
))),
}
}
pub fn verdict(
&self,
elapsed: Duration,
window: Duration,
cap: Duration,
start_requested: bool,
) -> WindowVerdict {
let joined = self.next_rank;
let manual = self.config.start_mode == StartMode::Manual;
if start_requested
&& self.config.start_mode != StartMode::Auto
&& joined >= self.config.min_rank_start
{
return WindowVerdict::Formed("operator start");
}
if let Some(target) = self.config.target_ranks
&& !manual
&& joined >= target
{
return WindowVerdict::Formed("target ranks reached");
}
if elapsed < window {
return WindowVerdict::Open;
}
if !manual && joined >= self.config.min_rank_start {
return WindowVerdict::Formed("join window closed with quorum");
}
if elapsed < cap {
return WindowVerdict::Open;
}
if manual && joined >= self.config.min_rank_start {
return WindowVerdict::Failed(format!(
"operator start not received: {joined} rank(s) staged with \
quorum met, but `start: manual` got no `fdl start` within \
the max_join_timeout hard cap ({}s scaled to {}s) — the \
window bounds the exposure; raise max_join_timeout for a \
longer hold",
self.config.max_join_timeout_secs,
cap.as_secs(),
));
}
WindowVerdict::Failed(format!(
"quorum not met: {joined}/{} ranks joined within the \
max_join_timeout hard cap ({}s scaled to {}s)",
self.config.min_rank_start,
self.config.max_join_timeout_secs,
cap.as_secs(),
))
}
pub fn open_phase(&self) -> ClusterPhase {
if self.config.start_mode != StartMode::Auto && self.next_rank >= self.config.min_rank_start
{
ClusterPhase::Staging
} else {
ClusterPhase::Waiting
}
}
pub fn snapshot(
&self,
phase: ClusterPhase,
elapsed: Duration,
start_armed: bool,
) -> MembershipSnapshot {
let window = Duration::from_secs(scaled_deadline_secs(self.config.join_timeout_secs));
let cap = Duration::from_secs(scaled_deadline_secs(self.config.max_join_timeout_secs));
let remaining =
|limit: Duration| -> Option<u64> { limit.checked_sub(elapsed).map(|d| d.as_secs()) };
MembershipSnapshot {
phase,
joined_ranks: self.next_rank,
joined_hosts: self.members.len(),
min_rank_start: self.config.min_rank_start,
target_ranks: self.config.target_ranks,
window_remaining_secs: remaining(window),
cap_remaining_secs: remaining(cap),
start_mode: self.config.start_mode,
start_armed,
members: self
.members
.iter()
.map(|m| MemberSnapshot {
host: m.host.clone(),
ranks: m.ranks.clone(),
local_devices: m.local_devices.clone(),
gpus: m.gpus.clone(),
libtorch: m.libtorch.clone(),
joined_at_secs: m.joined_at_secs,
})
.collect(),
}
}
fn into_members(self) -> Vec<JoinedMember> {
self.members
}
}
fn id_prefix(id: &str) -> &str {
&id[..id.len().min(8)]
}
pub(crate) fn hex_prefix(sig: &[u8; 32]) -> String {
use std::fmt::Write as _;
let mut s = String::with_capacity(8);
for b in &sig[..4] {
let _ = write!(s, "{b:02x}");
}
s
}
#[derive(Debug)]
pub(crate) struct AdmittedWorker {
pub member: JoinedMember,
pub stream: TcpStream,
}
#[derive(Debug)]
pub(crate) struct FormedWorld {
pub workers: Vec<AdmittedWorker>,
pub world_size: usize,
pub snapshot: MembershipSnapshot,
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn run_join_window(
source: &StreamSource,
config: &JoinConfig,
salt: &SessionSalt,
pre_shared_salt: bool,
expected_dataset_sig: Option<[u8; 32]>,
expected_model_sig: Option<[u8; 32]>,
abort: &AtomicBool,
status: &crate::distributed::status::StatusBoard,
) -> Result<FormedWorld> {
let mut ledger =
MembershipLedger::new(config.clone(), expected_dataset_sig, expected_model_sig)?;
let join_key = join_frame_key(!pre_shared_salt, salt);
let window = Duration::from_secs(scaled_deadline_secs(config.join_timeout_secs));
let cap = Duration::from_secs(scaled_deadline_secs(config.max_join_timeout_secs));
let started = Instant::now();
let mut admitted: Vec<AdmittedWorker> = Vec::new();
let mut rejected = 0usize;
eprintln!(
"cluster join: window open (quorum {} ranks, target {}, window {}s, \
cap {}s, admission: {})",
config.min_rank_start,
config
.target_ranks
.map(|t| t.to_string())
.unwrap_or_else(|| "none".to_string()),
window.as_secs(),
cap.as_secs(),
if pre_shared_salt {
"pre-shared salt"
} else {
"open"
},
);
status.publish(&ledger.snapshot(ledger.open_phase(), Duration::ZERO, false));
let mut armed_seen = false;
loop {
let elapsed = started.elapsed();
let armed = status.start_requested();
if armed && !armed_seen {
armed_seen = true;
eprintln!(
"cluster join: operator start armed ({} rank(s) in) — the \
world forms at the next quorum-met poll",
ledger.joined_ranks(),
);
status.publish(&ledger.snapshot(ledger.open_phase(), elapsed, true));
}
match ledger.verdict(elapsed, window, cap, armed) {
WindowVerdict::Open => {}
WindowVerdict::Formed(reason) => {
let world_size = ledger.joined_ranks();
let snapshot = ledger.snapshot(ClusterPhase::Forming, elapsed, armed);
let members = ledger.into_members();
eprintln!(
"cluster join: world formed — {world_size} ranks across \
{} host(s) after {}s ({reason})",
members.len(),
elapsed.as_secs(),
);
status.publish(&snapshot);
debug_assert_eq!(admitted.len(), members.len());
return Ok(FormedWorld {
workers: admitted,
world_size,
snapshot,
});
}
WindowVerdict::Failed(why) => {
let msg = format!("cluster join: FAILED — {why}");
eprintln!("{msg}");
status.publish(&ledger.snapshot(ClusterPhase::Failed, elapsed, armed));
abort_admitted(&mut admitted, salt, &why);
return Err(TensorError::new(&msg));
}
}
if abort.load(Ordering::SeqCst) {
let why = "launcher aborted before the world formed".to_string();
status.publish(&ledger.snapshot(ClusterPhase::Failed, started.elapsed(), armed));
abort_admitted(&mut admitted, salt, &why);
return Err(TensorError::new(&format!("cluster join: {why}")));
}
let mut stream = match source.try_accept("cluster join")? {
Some(s) => s,
None => {
std::thread::sleep(ACCEPT_POLL);
continue;
}
};
let _ = stream.set_nodelay(true);
let handshake = Duration::from_secs(scaled_deadline_secs(JOIN_HANDSHAKE_TIMEOUT_SECS));
if stream.set_read_timeout(Some(handshake)).is_err()
|| stream
.set_write_timeout(Some(crate::distributed::wire::write_stall_timeout()))
.is_err()
{
rejected += 1;
check_reject_cap(rejected, &ledger)?;
continue;
}
match handle_join_dial(
&mut stream,
&mut ledger,
&join_key,
salt,
pre_shared_salt,
started,
cap,
) {
Ok(member) => {
let snap_ranks = ledger.joined_ranks();
eprintln!(
"cluster join: host {:?} joined with ranks {:?} ({} GPU(s), \
libtorch {:?}) — {snap_ranks} rank(s) in{}",
member.host,
member.ranks,
member.gpus.len(),
member.libtorch,
match config.target_ranks {
Some(t) => format!(", target {t}"),
None => format!(", quorum {}", config.min_rank_start),
},
);
status.publish(&ledger.snapshot(
ledger.open_phase(),
started.elapsed(),
armed_seen,
));
admitted.push(AdmittedWorker { member, stream });
}
Err(why) => {
eprintln!("cluster join: rejected a join attempt: {why}");
rejected += 1;
check_reject_cap(rejected, &ledger)?;
}
}
}
}
fn check_reject_cap(rejected: usize, ledger: &MembershipLedger) -> Result<()> {
if rejected > MAX_REJECTED_JOINS {
return Err(TensorError::new(&format!(
"cluster join: aborting after {rejected} rejected join attempts \
({} ranks admitted) — the join port is being hammered by \
something that is not a flodl worker",
ledger.joined_ranks(),
)));
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn handle_join_dial(
stream: &mut TcpStream,
ledger: &mut MembershipLedger,
join_key: &SessionSalt,
salt: &SessionSalt,
pre_shared_salt: bool,
started: Instant,
cap: Duration,
) -> std::result::Result<JoinedMember, String> {
expect_channel_magic(stream, CHANNEL_MAGIC_JOIN, "cluster join").map_err(|e| e.to_string())?;
let frame = match ControlFrame::read_from(stream, join_key) {
Ok(Some(f)) => f,
Ok(None) => return Err("connection closed before hello".to_string()),
Err(e) => return Err(e.to_string()),
};
if frame.kind != MsgKind::Join {
let why = format!("expected a Join frame, got {:?}", frame.kind);
reject(stream, join_key, &why);
return Err(why);
}
let msg: JoinMsgWire = match frame.decode() {
Ok(m) => m,
Err(e) => {
let why = format!("hello decode failed: {e}");
reject(stream, join_key, &why);
return Err(why);
}
};
let JoinMsgWire::Hello {
host,
local_devices,
gpus,
libtorch,
dataset_sig,
run_id,
nccl_version,
model_sig,
} = msg
else {
let why = "first join-channel message must be Hello".to_string();
reject(stream, join_key, &why);
return Err(why);
};
let offer = JoinOffer {
host: host.clone(),
local_devices,
gpus,
libtorch,
dataset_sig,
run_id,
nccl_version,
model_sig,
};
let ranks = match ledger.admit(offer, started.elapsed()) {
Ok(r) => r,
Err(why) => {
reject(stream, join_key, &why);
return Err(format!("host {host:?}: {why}"));
}
};
let accept = JoinMsgWire::Accept {
ranks: ranks.iter().map(|r| *r as u32).collect(),
salt_hex: (!pre_shared_salt).then(|| salt_to_hex(salt)),
formation_wait_secs: cap.saturating_sub(started.elapsed()).as_secs(),
};
let write =
ControlFrame::encode(join_key, MsgKind::Join, &accept).and_then(|f| f.write_to(stream));
if let Err(e) = write {
let _ = ledger.retract_last(&host);
return Err(format!(
"host {host:?} admitted but the accept reply failed ({e}); \
admission rolled back"
));
}
let member = ledger
.members
.last()
.cloned()
.expect("admit() just pushed this member");
Ok(member)
}
fn reject(stream: &mut TcpStream, join_key: &SessionSalt, reason: &str) {
let msg = JoinMsgWire::Reject {
reason: reason.to_string(),
};
let _ = ControlFrame::encode(join_key, MsgKind::Join, &msg).and_then(|f| f.write_to(stream));
}
fn abort_admitted(admitted: &mut [AdmittedWorker], salt: &SessionSalt, reason: &str) {
let msg = JoinMsgWire::Abort {
reason: reason.to_string(),
};
for w in admitted.iter_mut() {
let _ =
ControlFrame::encode(salt, MsgKind::Join, &msg).and_then(|f| f.write_to(&mut w.stream));
}
}
#[cfg(test)]
#[path = "membership_tests.rs"]
mod tests;