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, 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,
}
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,
}
}
}
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 {
if 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,
)));
}
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,
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 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 MembershipLedger {
config: JoinConfig,
expected_dataset_sig: Option<[u8; 32]>,
members: Vec<JoinedMember>,
next_rank: usize,
}
impl MembershipLedger {
pub fn new(
config: JoinConfig,
expected_dataset_sig: Option<[u8; 32]>,
) -> Result<Self> {
config.validate()?;
Ok(MembershipLedger {
config,
expected_dataset_sig,
members: Vec::new(),
next_rank: 0,
})
}
pub fn joined_ranks(&self) -> usize {
self.next_rank
}
pub fn admit(
&mut self,
host: &str,
local_devices: Vec<u8>,
gpus: Vec<String>,
libtorch: String,
dataset_sig: [u8; 32],
elapsed: Duration,
) -> std::result::Result<Vec<usize>, String> {
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(_) => {}
}
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,
) -> WindowVerdict {
let joined = self.next_rank;
if let Some(target) = self.config.target_ranks {
if joined >= target {
return WindowVerdict::Formed("target ranks reached");
}
}
if elapsed < window {
return WindowVerdict::Open;
}
if joined >= self.config.min_rank_start {
return WindowVerdict::Formed("join window closed with quorum");
}
if elapsed < cap {
return WindowVerdict::Open;
}
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 snapshot(&self, phase: ClusterPhase, elapsed: Duration) -> 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),
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 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,
}
pub(crate) fn run_join_window(
source: &StreamSource,
config: &JoinConfig,
salt: &SessionSalt,
pre_shared_salt: bool,
expected_dataset_sig: Option<[u8; 32]>,
abort: &AtomicBool,
status: &crate::distributed::status::StatusBoard,
) -> Result<FormedWorld> {
let mut ledger = MembershipLedger::new(config.clone(), expected_dataset_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(ClusterPhase::Waiting, Duration::ZERO));
loop {
let elapsed = started.elapsed();
match ledger.verdict(elapsed, window, cap) {
WindowVerdict::Open => {}
WindowVerdict::Formed(reason) => {
let world_size = ledger.joined_ranks();
let snapshot = ledger.snapshot(ClusterPhase::Forming, elapsed);
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));
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()),
);
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(ClusterPhase::Waiting, started.elapsed()),
);
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 } = msg
else {
let why = "first join-channel message must be Hello".to_string();
reject(stream, join_key, &why);
return Err(why);
};
let ranks = match ledger.admit(
&host,
local_devices,
gpus,
libtorch,
dataset_sig,
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;