use std::io::{ErrorKind, Read, Write};
use std::net::{SocketAddr, TcpListener, TcpStream};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Condvar, Mutex};
use std::thread::{self, JoinHandle};
use std::time::Duration;
use hmac_sha256::HMAC;
use crate::distributed::relay::mux::{MuxRead, MuxRecord, RelayControlMsg};
use crate::distributed::wire::SessionSalt;
use crate::tensor::{Result, TensorError};
pub(crate) const HANDSHAKE_MAGIC_RANK: u32 = 0xF10D_17C0;
pub(crate) const HANDSHAKE_MAGIC_CONTROLLER_ACK: u32 = 0xF10D_17C1;
pub(crate) const ROUND_FRAME_MAGIC: u32 = 0xF10D_17F1;
pub(crate) const PROTOCOL_VERSION: u32 = 2;
pub const DTYPE_F32: u8 = 0;
pub const DTYPE_BF16: u8 = 1;
pub(crate) const MAX_ROUND_FRAME_TENSORS: usize = 65_536;
mod dead_ranks;
mod round_frame;
pub use dead_ranks::DeadRanks;
pub use round_frame::{RoundFrame, RoundKind, TensorPayload};
pub(crate) use round_frame::{read_round_frame, write_round_frame};
#[cfg(test)]
pub(crate) use round_frame::sum_frames;
pub(crate) use round_frame::{f32_slice_to_payload_bytes, payload_to_f32};
pub(crate) use round_frame::{
accumulate_payload_into, payload_element_size, read_round_frame_streamed,
round_frame_wire_len, scale_payload_bytes, write_round_frame_streamed, PayloadPart,
};
use round_frame::reduce_realized_work;
#[cfg(test)]
use round_frame::{bytes_as_f32, f32_to_bytes};
#[derive(Debug)]
pub struct ClusterController {
bound_port: u16,
shutdown: Arc<AtomicBool>,
handle: Option<JoinHandle<Result<()>>>,
}
impl ClusterController {
#[allow(dead_code)]
pub fn start(
bind_addr: SocketAddr,
world_size: usize,
salt: SessionSalt,
) -> Result<Self> {
let dead_ranks = DeadRanks::new(world_size);
Self::start_with_dead_ranks(bind_addr, world_size, salt, dead_ranks, None, None)
}
pub fn start_with_dead_ranks(
bind_addr: SocketAddr,
world_size: usize,
salt: SessionSalt,
dead_ranks: Arc<DeadRanks>,
forge: Option<Arc<crate::distributed::CheckpointForge>>,
outer_optimizer: Option<Box<dyn crate::distributed::OuterOptimizer>>,
) -> Result<Self> {
let listener = TcpListener::bind(bind_addr).map_err(|e| {
TensorError::new(&format!(
"cluster_controller: bind {bind_addr} failed: {e}"
))
})?;
let bound_port = listener
.local_addr()
.map_err(|e| {
TensorError::new(&format!(
"cluster_controller: local_addr() failed: {e}"
))
})?
.port();
let source = crate::distributed::port_mux::StreamSource::from_listener(
listener,
"cluster_controller",
)?;
Self::start_from_source(
source, bound_port, world_size, salt, dead_ranks, forge,
outer_optimizer,
)
}
pub(crate) fn start_from_source(
source: crate::distributed::port_mux::StreamSource,
bound_port: u16,
world_size: usize,
salt: SessionSalt,
dead_ranks: Arc<DeadRanks>,
forge: Option<Arc<crate::distributed::CheckpointForge>>,
outer_optimizer: Option<Box<dyn crate::distributed::OuterOptimizer>>,
) -> Result<Self> {
if world_size == 0 {
return Err(TensorError::new(
"cluster_controller: world_size must be > 0",
));
}
if dead_ranks.world_size() != world_size {
return Err(TensorError::new(&format!(
"cluster_controller: dead_ranks world_size ({}) must match \
controller world_size ({})",
dead_ranks.world_size(),
world_size,
)));
}
let shutdown = Arc::new(AtomicBool::new(false));
let shutdown_cloned = Arc::clone(&shutdown);
let handle = thread::Builder::new()
.name(format!("flodl-cluster-controller:{bound_port}"))
.spawn(move || {
run_reduce_thread(
source, world_size, salt, shutdown_cloned, dead_ranks, forge,
outer_optimizer,
)
})
.map_err(|e| {
TensorError::new(&format!("cluster_controller: spawn worker failed: {e}"))
})?;
Ok(ClusterController {
bound_port,
shutdown,
handle: Some(handle),
})
}
pub fn port(&self) -> u16 {
self.bound_port
}
pub fn shutdown(mut self) -> Result<()> {
self.shutdown.store(true, Ordering::SeqCst);
if let Some(h) = self.handle.take() {
return h
.join()
.map_err(|_| TensorError::new("cluster_controller: worker panicked"))?;
}
Ok(())
}
}
impl Drop for ClusterController {
fn drop(&mut self) {
self.shutdown.store(true, Ordering::SeqCst);
if let Some(h) = self.handle.take() {
let _ = h.join();
}
}
}
const REDUCE_POLL: Duration = Duration::from_millis(100);
fn run_reduce_thread(
source: crate::distributed::port_mux::StreamSource,
world_size: usize,
salt: SessionSalt,
shutdown: Arc<AtomicBool>,
dead_ranks: Arc<DeadRanks>,
forge: Option<Arc<crate::distributed::CheckpointForge>>,
outer_optimizer: Option<Box<dyn crate::distributed::OuterOptimizer>>,
) -> Result<()> {
let mut outer_stepper = outer_optimizer
.map(crate::distributed::outer_optimizer::OuterStepper::new);
let slots = Arc::new(ReduceSlots::new());
let mut conn_writes: Vec<TcpStream> = Vec::new();
let mut rank_conn: Vec<Option<usize>> = (0..world_size).map(|_| None).collect();
let mut all_conn_ranks: Vec<Vec<usize>> = Vec::new();
let mut reader_threads: Vec<JoinHandle<()>> = Vec::new();
let mut covered = 0usize;
while covered < world_size {
if shutdown.load(Ordering::SeqCst) {
return Ok(());
}
match source.try_accept("cluster_controller") {
Ok(None) => {
thread::sleep(Duration::from_millis(20));
}
Ok(Some(mut stream)) => {
let _ = stream.set_nodelay(true);
stream
.set_write_timeout(Some(crate::distributed::wire::write_stall_timeout()))
.map_err(|e| {
TensorError::new(&format!(
"cluster_controller: set_write_timeout: {e}"
))
})?;
stream
.set_read_timeout(Some(Duration::from_secs(10)))
.map_err(|e| {
TensorError::new(&format!(
"cluster_controller: set_read_timeout: {e}"
))
})?;
crate::distributed::wire::expect_channel_magic(
&mut stream,
crate::distributed::wire::CHANNEL_MAGIC_DATA,
"cluster_controller",
)?;
let ranks = match MuxRecord::read_from(&mut stream, &salt)? {
Some(MuxRecord::Control(RelayControlMsg::Hello { host, ranks })) => {
crate::verbose!(
" cluster_controller: relay '{host}' carries ranks {ranks:?}"
);
ranks
}
Some(other) => {
return Err(TensorError::new(&format!(
"cluster_controller: expected relay Hello, got {other:?}"
)));
}
None => {
return Err(TensorError::new(
"cluster_controller: relay closed connection before Hello",
));
}
};
let conn_idx = conn_writes.len();
let mut conn_ranks: Vec<usize> = Vec::with_capacity(ranks.len());
for r in &ranks {
let r = *r as usize;
if r >= world_size {
return Err(TensorError::new(&format!(
"cluster_controller: relay announced rank {r} >= world_size {world_size}"
)));
}
if rank_conn[r].is_some() {
return Err(TensorError::new(&format!(
"cluster_controller: rank {r} announced by two relays"
)));
}
rank_conn[r] = Some(conn_idx);
conn_ranks.push(r);
}
MuxRecord::control(RelayControlMsg::HelloAck).write_to(&mut stream, &salt)?;
let read_half = stream.try_clone().map_err(|e| {
TensorError::new(&format!("cluster_controller: relay try_clone: {e}"))
})?;
read_half
.set_read_timeout(Some(REDUCE_POLL))
.map_err(|e| {
TensorError::new(&format!("cluster_controller: set_read_timeout: {e}"))
})?;
conn_writes.push(stream);
covered += conn_ranks.len();
let registered = slots.register_conn(conn_ranks.clone());
debug_assert_eq!(registered, conn_idx);
all_conn_ranks.push(conn_ranks.clone());
let slots_c = Arc::clone(&slots);
let dead_c = Arc::clone(&dead_ranks);
let shutdown_c = Arc::clone(&shutdown);
let t = thread::Builder::new()
.name(format!("flodl-controller-relay{conn_idx}"))
.spawn(move || {
reduce_reader(
read_half, conn_idx, conn_ranks, slots_c, dead_c, shutdown_c,
salt,
)
})
.map_err(|e| {
TensorError::new(&format!("cluster_controller: spawn reader: {e}"))
})?;
reader_threads.push(t);
}
Err(e) => return Err(e),
}
}
drop(source);
let mut forwarded_dead: Vec<bool> = vec![false; world_size];
let outcome = loop {
if shutdown.load(Ordering::SeqCst) {
break Ok(());
}
let round = {
let conn_writes = &mut conn_writes;
let forwarded_dead = &mut forwarded_dead;
slots.wait_for_round(&dead_ranks, &shutdown, REDUCE_POLL, || {
forward_dead_diffs(
&dead_ranks,
forwarded_dead,
&rank_conn,
conn_writes,
&salt,
);
})
};
match round {
RoundOutcome::Frames(frames) => {
if let Err(e) = average_and_scatter(
&frames,
&mut conn_writes,
&all_conn_ranks,
&dead_ranks,
&salt,
forge.as_deref(),
outer_stepper.as_mut(),
) {
break Err(e);
}
}
RoundOutcome::Shutdown => break Ok(()),
RoundOutcome::Error(e) => break Err(e),
}
};
shutdown.store(true, Ordering::SeqCst);
for t in reader_threads {
let _ = t.join();
}
outcome
}
struct ReduceSlots {
inner: Mutex<SlotsInner>,
cv: Condvar,
}
struct SlotsInner {
frames: Vec<Option<RoundFrame>>,
conn_ranks: Vec<Vec<usize>>,
shutdown: bool,
error: Option<TensorError>,
}
enum RoundOutcome {
Frames(Vec<Option<RoundFrame>>),
Shutdown,
Error(TensorError),
}
impl ReduceSlots {
fn new() -> Self {
ReduceSlots {
inner: Mutex::new(SlotsInner {
frames: Vec::new(),
conn_ranks: Vec::new(),
shutdown: false,
error: None,
}),
cv: Condvar::new(),
}
}
fn register_conn(&self, ranks: Vec<usize>) -> usize {
let mut inner = self.inner.lock().unwrap();
inner.frames.push(None);
inner.conn_ranks.push(ranks);
inner.frames.len() - 1
}
fn deposit(&self, conn: usize, frame: RoundFrame) {
let mut inner = self.inner.lock().unwrap();
if conn < inner.frames.len() {
inner.frames[conn] = Some(frame);
}
self.cv.notify_all();
}
fn request_shutdown(&self) {
let mut inner = self.inner.lock().unwrap();
inner.shutdown = true;
self.cv.notify_all();
}
fn set_error(&self, err: TensorError) {
let mut inner = self.inner.lock().unwrap();
if inner.error.is_none() {
inner.error = Some(err);
}
self.cv.notify_all();
}
fn wait_for_round(
&self,
dead: &DeadRanks,
external_shutdown: &AtomicBool,
poll: Duration,
mut on_poll: impl FnMut(),
) -> RoundOutcome {
let mut inner = self.inner.lock().unwrap();
loop {
if external_shutdown.load(Ordering::SeqCst) || inner.shutdown {
return RoundOutcome::Shutdown;
}
if let Some(e) = inner.error.take() {
return RoundOutcome::Error(e);
}
let n_conns = inner.frames.len();
let expected: Vec<usize> = (0..n_conns)
.filter(|c| inner.conn_ranks[*c].iter().any(|r| !dead.is_dead(*r)))
.collect();
if expected.is_empty() {
return RoundOutcome::Shutdown;
}
if expected.iter().all(|c| inner.frames[*c].is_some()) {
let mut out: Vec<Option<RoundFrame>> = Vec::with_capacity(n_conns);
for c in 0..n_conns {
if expected.contains(&c) {
out.push(inner.frames[c].take());
} else {
inner.frames[c] = None;
out.push(None);
}
}
return RoundOutcome::Frames(out);
}
drop(inner);
on_poll();
inner = self.inner.lock().unwrap();
let (guard, _timeout) = self.cv.wait_timeout(inner, poll).unwrap();
inner = guard;
}
}
}
fn reduce_reader(
mut read: TcpStream,
conn_idx: usize,
ranks: Vec<usize>,
slots: Arc<ReduceSlots>,
dead_ranks: Arc<DeadRanks>,
shutdown: Arc<AtomicBool>,
salt: SessionSalt,
) {
loop {
if shutdown.load(Ordering::SeqCst) {
return;
}
match MuxRecord::try_read_from(&mut read, &salt) {
Ok(MuxRead::Record(MuxRecord::HostFrame { payload })) => {
if ranks.iter().all(|r| dead_ranks.is_dead(*r)) {
continue; }
let mut slice = &payload[..];
match read_round_frame(&mut slice, &salt) {
Ok(Some(frame)) => slots.deposit(conn_idx, frame),
Ok(None) => {
slots.set_error(TensorError::new(&format!(
"cluster_controller: truncated HostFrame payload on \
connection {conn_idx} (ranks {ranks:?})"
)));
return;
}
Err(e) => {
slots.set_error(e);
return;
}
}
}
Ok(MuxRead::Record(MuxRecord::Data { rank, .. })) => {
slots.set_error(TensorError::new(&format!(
"cluster_controller: per-rank Data record (rank {rank}) on the \
data channel; the relay is expected to fold local frames into \
a HostFrame — mixed relay/controller builds?"
)));
return;
}
Ok(MuxRead::Record(MuxRecord::Control(RelayControlMsg::RankExit { rank }))) => {
if !dead_ranks.is_dead(rank as usize) {
slots.request_shutdown();
}
}
Ok(MuxRead::Record(MuxRecord::Control(_))) => {
}
Ok(MuxRead::Record(MuxRecord::Broadcast { .. })) => {
eprintln!(
"cluster_controller: unexpected Broadcast record from relay \
{conn_idx}; dropping"
);
}
Ok(MuxRead::WouldBlock) => {}
Ok(MuxRead::Eof) => {
if ranks.iter().any(|r| !dead_ranks.is_dead(*r)) {
slots.request_shutdown();
}
return;
}
Err(e) => {
if ranks.iter().any(|r| !dead_ranks.is_dead(*r)) {
slots.set_error(e);
}
return;
}
}
}
}
fn forward_dead_diffs(
dead_ranks: &DeadRanks,
forwarded: &mut [bool],
rank_conn: &[Option<usize>],
conn_writes: &mut [TcpStream],
salt: &SessionSalt,
) {
for (rank, fwd) in forwarded.iter_mut().enumerate() {
if *fwd || !dead_ranks.is_dead(rank) {
continue;
}
*fwd = true;
let Some(ci) = rank_conn.get(rank).copied().flatten() else {
continue;
};
if let Err(e) = MuxRecord::control(RelayControlMsg::DeclareDead {
rank: rank as u32,
})
.write_to(&mut conn_writes[ci], salt)
{
crate::verbose!(
" cluster_controller: DeclareDead({rank}) forward to relay \
{ci} failed ({e}); connection presumed dying",
);
}
}
}
fn average_and_scatter(
frames: &[Option<RoundFrame>],
conn_writes: &mut [TcpStream],
conn_ranks: &[Vec<usize>],
dead_ranks: &DeadRanks,
salt: &SessionSalt,
forge: Option<&crate::distributed::CheckpointForge>,
outer_stepper: Option<&mut crate::distributed::outer_optimizer::OuterStepper>,
) -> Result<()> {
let averaged = reduce_realized_work(frames)?;
let mut outer_momentum: Option<Vec<TensorPayload>> = None;
let averaged = match outer_stepper {
Some(stepper) => {
let stepped = stepper.process_frame(averaged)?;
if stepped.kind == RoundKind::Model
&& forge.is_some_and(|f| f.is_armed())
&& let Some(m) = stepper.checkpoint_state()
{
let refs: Vec<&crate::tensor::Tensor> = m.iter().collect();
outer_momentum = Some(
crate::distributed::cpu_reduce::tensors_to_round_frame(&refs, DTYPE_F32)?
.tensors,
);
}
stepped
}
None => averaged,
};
let mut buf: Vec<u8> = Vec::new();
write_round_frame(&mut buf, &averaged, salt)?;
let record = MuxRecord::broadcast(buf);
for (ci, ranks) in conn_ranks.iter().enumerate() {
if ranks.iter().all(|r| dead_ranks.is_dead(*r)) {
continue; }
if let Err(e) = record.write_to(&mut conn_writes[ci], salt) {
eprintln!(
"cluster_controller: consensus broadcast to relay {ci} \
(ranks {ranks:?}) failed ({e}); declaring its ranks dead and \
continuing with survivors"
);
for r in ranks {
dead_ranks.declare_dead(*r);
}
}
}
if let Some(f) = forge {
if averaged.kind == RoundKind::Model {
if let Some(m) = outer_momentum {
f.stash_outer_momentum(m);
}
f.accumulate(averaged);
}
}
Ok(())
}
pub(crate) fn read_handshake(stream: &mut TcpStream, expected_world_size: usize) -> Result<usize> {
let mut buf = [0u8; 16];
stream.read_exact(&mut buf).map_err(|e| {
TensorError::new(&format!("cluster_controller: handshake read failed: {e}"))
})?;
let magic = u32::from_le_bytes(buf[0..4].try_into().unwrap());
if magic != HANDSHAKE_MAGIC_RANK {
return Err(TensorError::new(&format!(
"cluster_controller: handshake magic 0x{magic:08x} != 0x{HANDSHAKE_MAGIC_RANK:08x}"
)));
}
let proto_ver = u32::from_le_bytes(buf[4..8].try_into().unwrap());
if proto_ver != PROTOCOL_VERSION {
return Err(TensorError::new(&format!(
"cluster_controller: handshake protocol_version {proto_ver} != {PROTOCOL_VERSION}"
)));
}
let rank_id = u32::from_le_bytes(buf[8..12].try_into().unwrap()) as usize;
let rank_world_size = u32::from_le_bytes(buf[12..16].try_into().unwrap()) as usize;
if rank_world_size != expected_world_size {
return Err(TensorError::new(&format!(
"cluster_controller: handshake world_size {rank_world_size} != expected {expected_world_size}"
)));
}
Ok(rank_id)
}
pub(crate) fn write_handshake_ack(stream: &mut TcpStream) -> Result<()> {
let mut buf = [0u8; 8];
buf[0..4].copy_from_slice(&HANDSHAKE_MAGIC_CONTROLLER_ACK.to_le_bytes());
buf[4..8].copy_from_slice(&PROTOCOL_VERSION.to_le_bytes());
stream.write_all(&buf).map_err(|e| {
TensorError::new(&format!("cluster_controller: handshake ack write failed: {e}"))
})?;
Ok(())
}
#[cfg(test)]
#[path = "../controller_tests.rs"]
mod tests;