use std::collections::{HashMap, HashSet};
use std::net::{SocketAddr, TcpListener, TcpStream};
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::mpsc::{self, RecvTimeoutError};
use std::sync::{Arc, Mutex};
use std::thread::{self, JoinHandle};
use std::time::Duration;
use crate::distributed::controller::{self, RoundKind, TensorPayload};
use crate::distributed::wire::SessionSalt;
use crate::tensor::{Result, TensorError};
use super::mux::{
try_read_len_framed, write_len_framed, LenFramedRead, MuxRead, MuxRecord, RelayControlMsg,
};
const POLL_TIMEOUT: Duration = Duration::from_millis(100);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ChannelKind {
Data,
Control,
}
impl ChannelKind {
fn channel_magic(self) -> u32 {
match self {
ChannelKind::Data => crate::distributed::wire::CHANNEL_MAGIC_DATA,
ChannelKind::Control => crate::distributed::wire::CHANNEL_MAGIC_CONTROL,
}
}
fn terminate_handshake(
self,
stream: &mut TcpStream,
world_size: usize,
salt: &SessionSalt,
) -> Result<u32> {
match self {
ChannelKind::Data => {
let rank =
crate::distributed::controller::read_handshake(stream, world_size)?;
crate::distributed::controller::write_handshake_ack(stream)?;
Ok(rank as u32)
}
ChannelKind::Control => {
let rank = crate::distributed::wire::read_handshake_rank(
stream,
world_size as u32,
salt,
)?;
crate::distributed::wire::write_handshake_ack(stream, salt)?;
Ok(rank)
}
}
}
}
pub struct RelayChannel {
shutdown: Arc<AtomicBool>,
threads: Vec<JoinHandle<()>>,
}
impl RelayChannel {
pub fn bind(loopback_bind: SocketAddr) -> Result<(TcpListener, u16)> {
let listener = TcpListener::bind(loopback_bind).map_err(|e| {
TensorError::new(&format!("relay: loopback bind {loopback_bind} failed: {e}"))
})?;
let port = listener
.local_addr()
.map_err(|e| TensorError::new(&format!("relay: local_addr failed: {e}")))?
.port();
Ok((listener, port))
}
pub fn start(
listener: TcpListener,
kind: ChannelKind,
upstream_addr: SocketAddr,
host: String,
mut ranks: Vec<u32>,
world_size: usize,
salt: SessionSalt,
) -> Result<Self> {
ranks.sort_unstable();
let expected: HashSet<u32> = ranks.iter().copied().collect();
let mut rank_streams: Vec<(u32, TcpStream)> = Vec::with_capacity(ranks.len());
while rank_streams.len() < ranks.len() {
let (mut stream, _peer) = listener.accept().map_err(|e| {
TensorError::new(&format!("relay: loopback accept failed: {e}"))
})?;
let _ = stream.set_nodelay(true);
let _ = stream
.set_write_timeout(Some(crate::distributed::wire::write_stall_timeout()));
let rank = kind.terminate_handshake(&mut stream, world_size, &salt)?;
if !expected.contains(&rank) {
return Err(TensorError::new(&format!(
"relay: rank {rank} connected but is not in this host's rank set {ranks:?}"
)));
}
if rank_streams.iter().any(|(r, _)| *r == rank) {
return Err(TensorError::new(&format!(
"relay: duplicate rank {rank} connected on loopback"
)));
}
rank_streams.push((rank, stream));
}
let mut upstream = crate::distributed::wire::connect_with_retry(
upstream_addr,
"relay upstream",
)?;
let _ = upstream.set_nodelay(true);
if let Ok(peer) = upstream.peer_addr() {
crate::distributed::wire::warn_cleartext_public_peer("relay upstream", peer);
}
let _ = upstream
.set_write_timeout(Some(crate::distributed::wire::write_stall_timeout()));
crate::distributed::wire::write_channel_magic(
&mut upstream,
kind.channel_magic(),
)?;
MuxRecord::control(RelayControlMsg::Hello {
host,
ranks: ranks.clone(),
})
.write_to(&mut upstream, &salt)?;
match MuxRecord::read_from(&mut upstream, &salt)? {
Some(MuxRecord::Control(RelayControlMsg::HelloAck)) => {}
Some(other) => {
return Err(TensorError::new(&format!(
"relay: expected HelloAck from controller, got {other:?}"
)));
}
None => {
return Err(TensorError::new(
"relay: controller closed connection before HelloAck",
));
}
}
let shutdown = Arc::new(AtomicBool::new(false));
let threads = spawn_mux(kind, rank_streams, upstream, salt, Arc::clone(&shutdown))?;
Ok(RelayChannel { shutdown, threads })
}
#[allow(dead_code)]
pub fn shutdown(mut self) -> Result<()> {
self.shutdown.store(true, Ordering::SeqCst);
for t in self.threads.drain(..) {
let _ = t.join();
}
Ok(())
}
pub fn join(mut self) -> Result<()> {
for t in self.threads.drain(..) {
let _ = t.join();
}
Ok(())
}
}
impl Drop for RelayChannel {
fn drop(&mut self) {
self.shutdown.store(true, Ordering::SeqCst);
for t in self.threads.drain(..) {
let _ = t.join();
}
}
}
struct FoldCtx {
inner: Mutex<FoldInner>,
local_ranks: Vec<u32>,
salt: SessionSalt,
rounds: AtomicU64,
bytes_in: AtomicU64,
bytes_up: AtomicU64,
}
struct FoldInner {
fold: Option<HostFold>,
dead: HashSet<u32>,
}
struct HostFold {
kind: RoundKind,
weight: f64,
payloads: FoldPayloads,
deposited: HashSet<u32>,
}
enum FoldPayloads {
Seed(Vec<TensorPayload>),
Sums {
schema: Vec<FoldSchema>,
sums: Vec<Vec<f32>>,
},
}
struct FoldSchema {
dtype: u8,
shape: Vec<u32>,
}
fn fold_schema_check(
rank: u32,
ti: usize,
got_dtype: u8,
got_shape: &[u32],
want_dtype: u8,
want_shape: &[u32],
) -> Result<()> {
if got_dtype != want_dtype {
return Err(TensorError::new(&format!(
"relay fold: rank {rank} tensor[{ti}] dtype {got_dtype} != the round's \
dtype {want_dtype} (cohort mixing bf16_wire settings, or desynced \
rounds)"
)));
}
if got_shape != want_shape {
return Err(TensorError::new(&format!(
"relay fold: rank {rank} tensor[{ti}] shape {got_shape:?} != the \
round's shape {want_shape:?} (desynced rounds)"
)));
}
Ok(())
}
impl HostFold {
fn accumulate(
fold: &mut Option<HostFold>,
rank: u32,
blob: &[u8],
salt: &SessionSalt,
) -> Result<()> {
let Some(f) = fold else {
let mut payloads: Vec<TensorPayload> = Vec::new();
let hdr = controller::read_round_frame_streamed(
&mut &blob[..],
salt,
&mut |_, p| {
payloads.push(p);
Ok(())
},
)?;
let Some((kind, weight)) = hdr else {
return Err(TensorError::new(&format!(
"relay fold: truncated RoundFrame from rank {rank}"
)));
};
*fold = Some(HostFold {
kind,
weight,
payloads: FoldPayloads::Seed(payloads),
deposited: HashSet::from([rank]),
});
return Ok(());
};
if let FoldPayloads::Seed(seed) = &mut f.payloads {
let mut schema: Vec<FoldSchema> = Vec::with_capacity(seed.len());
let mut sums: Vec<Vec<f32>> = Vec::with_capacity(seed.len());
for p in seed.iter_mut() {
sums.push(controller::payload_to_f32(p)?);
schema.push(FoldSchema {
dtype: p.dtype,
shape: std::mem::take(&mut p.shape),
});
p.bytes = Vec::new();
}
f.payloads = FoldPayloads::Sums { schema, sums };
}
let FoldPayloads::Sums { schema, sums } = &mut f.payloads else {
unreachable!("seed promoted just above");
};
let expected = schema.len();
let mut seen = 0usize;
let hdr = controller::read_round_frame_streamed(
&mut &blob[..],
salt,
&mut |ti, p| {
let (Some(sc), Some(sum)) = (schema.get(ti), sums.get_mut(ti)) else {
return Err(TensorError::new(&format!(
"relay fold: rank {rank} frame carries more than {expected} \
tensors (schema pinned by the round's first deposit)"
)));
};
fold_schema_check(rank, ti, p.dtype, &p.shape, sc.dtype, &sc.shape)?;
controller::accumulate_payload_into(&p, sum).map_err(|e| {
TensorError::new(&format!("relay fold: rank {rank} tensor[{ti}]: {e}"))
})?;
seen = ti + 1;
Ok(())
},
)?;
let Some((kind, weight)) = hdr else {
return Err(TensorError::new(&format!(
"relay fold: truncated RoundFrame from rank {rank}"
)));
};
if kind != f.kind {
return Err(TensorError::new(&format!(
"relay fold: rank {rank} frame kind {kind:?} != the round's kind \
{:?} (desynced reduce rounds)",
f.kind
)));
}
if seen != expected {
return Err(TensorError::new(&format!(
"relay fold: rank {rank} frame carries {seen} tensors; the round's \
first deposit carried {expected}"
)));
}
f.weight += weight;
f.deposited.insert(rank);
Ok(())
}
}
impl FoldCtx {
fn new(local_ranks: Vec<u32>, salt: SessionSalt) -> Self {
FoldCtx {
inner: Mutex::new(FoldInner {
fold: None,
dead: HashSet::new(),
}),
local_ranks,
salt,
rounds: AtomicU64::new(0),
bytes_in: AtomicU64::new(0),
bytes_up: AtomicU64::new(0),
}
}
fn deposit(
&self,
rank: u32,
blob: &[u8],
tx: &mpsc::SyncSender<MuxRecord>,
) -> Result<()> {
let taken = {
let mut inner = self.inner.lock().expect("relay fold lock poisoned");
if inner.dead.contains(&rank) {
return Ok(());
}
if inner
.fold
.as_ref()
.is_some_and(|f| f.deposited.contains(&rank))
{
return Err(TensorError::new(&format!(
"relay fold: rank {rank} deposited twice in one round \
(rank↔relay ping-pong violated; stream desynced)"
)));
}
HostFold::accumulate(&mut inner.fold, rank, blob, &self.salt)?;
self.bytes_in.fetch_add(blob.len() as u64, Ordering::Relaxed);
self.take_if_complete(&mut inner)
};
self.fold_and_ship(taken, tx)
}
fn mark_dead(&self, rank: u32, tx: &mpsc::SyncSender<MuxRecord>) -> Result<()> {
let taken = {
let mut inner = self.inner.lock().expect("relay fold lock poisoned");
inner.dead.insert(rank);
self.take_if_complete(&mut inner)
};
self.fold_and_ship(taken, tx)
}
fn take_if_complete(&self, inner: &mut FoldInner) -> Option<HostFold> {
let complete = self.local_ranks.iter().all(|r| {
inner.dead.contains(r)
|| inner
.fold
.as_ref()
.is_some_and(|f| f.deposited.contains(r))
});
if !complete {
return None;
}
inner.fold.take()
}
fn fold_and_ship(
&self,
taken: Option<HostFold>,
tx: &mpsc::SyncSender<MuxRecord>,
) -> Result<()> {
let Some(fold) = taken else {
return Ok(());
};
let HostFold {
kind,
weight,
payloads,
deposited: _,
} = fold;
let mut buf: Vec<u8>;
match payloads {
FoldPayloads::Seed(seed) => {
let parts: Vec<controller::PayloadPart<'_>> = seed
.iter()
.map(|p| controller::PayloadPart {
dtype: p.dtype,
shape: &p.shape,
nbytes: p.bytes.len() as u64,
})
.collect();
buf = Vec::with_capacity(
controller::round_frame_wire_len(&parts) as usize,
);
controller::write_round_frame_streamed(
&mut buf,
kind,
weight,
&parts,
&self.salt,
&mut |ti, tee| {
use std::io::Write;
tee.write_all(&seed[ti].bytes)
.map_err(|e| TensorError::new(&e.to_string()))
},
)?;
drop(seed);
}
FoldPayloads::Sums { schema, mut sums } => {
let parts: Vec<controller::PayloadPart<'_>> = schema
.iter()
.zip(sums.iter())
.map(|(s, sum)| {
Ok(controller::PayloadPart {
dtype: s.dtype,
shape: &s.shape,
nbytes: (sum.len()
* controller::payload_element_size(s.dtype)?)
as u64,
})
})
.collect::<Result<_>>()?;
buf = Vec::with_capacity(
controller::round_frame_wire_len(&parts) as usize,
);
controller::write_round_frame_streamed(
&mut buf,
kind,
weight,
&parts,
&self.salt,
&mut |ti, tee| {
use std::io::Write;
let sum = std::mem::take(&mut sums[ti]);
let bytes = controller::f32_slice_to_payload_bytes(
&sum,
schema[ti].dtype,
)?;
drop(sum);
tee.write_all(&bytes)
.map_err(|e| TensorError::new(&e.to_string()))
},
)?;
}
}
self.rounds.fetch_add(1, Ordering::Relaxed);
self.bytes_up.fetch_add(buf.len() as u64, Ordering::Relaxed);
tx.send(MuxRecord::host_frame(buf)).map_err(|_| {
TensorError::new("relay fold: outbound writer gone before fold shipped")
})
}
fn fan_out(&self, payload: &[u8], rank_writes: &mut HashMap<u32, TcpStream>) {
let dead: Vec<u32> = {
let inner = self.inner.lock().expect("relay fold lock poisoned");
inner.dead.iter().copied().collect()
};
for rank in &self.local_ranks {
if dead.contains(rank) {
continue;
}
if let Some(w) = rank_writes.get_mut(rank) {
let _ = write_len_framed(w, payload);
}
}
}
}
impl Drop for FoldCtx {
fn drop(&mut self) {
let rounds = self.rounds.load(Ordering::Relaxed);
if rounds == 0 {
return;
}
let bytes_in = self.bytes_in.load(Ordering::Relaxed) as f64 / 1e6;
let bytes_up = self.bytes_up.load(Ordering::Relaxed) as f64 / 1e6;
let ratio = if bytes_up > 0.0 { bytes_in / bytes_up } else { 0.0 };
eprintln!(
"[relay-fold-prof] ranks={} rounds={rounds} | local={bytes_in:.2}MB \
uplink={bytes_up:.2}MB fold-ratio={ratio:.2}x",
self.local_ranks.len(),
);
}
}
fn spawn_mux(
kind: ChannelKind,
rank_streams: Vec<(u32, TcpStream)>,
upstream: TcpStream,
salt: SessionSalt,
shutdown: Arc<AtomicBool>,
) -> Result<Vec<JoinHandle<()>>> {
let fold: Option<Arc<FoldCtx>> = match kind {
ChannelKind::Data => Some(Arc::new(FoldCtx::new(
rank_streams.iter().map(|(r, _)| *r).collect(),
salt,
))),
ChannelKind::Control => None,
};
let (tx, rx) = mpsc::sync_channel::<MuxRecord>(64);
let active = Arc::new(AtomicUsize::new(rank_streams.len()));
let up_read = upstream
.try_clone()
.map_err(|e| TensorError::new(&format!("relay: upstream try_clone failed: {e}")))?;
up_read
.set_read_timeout(Some(POLL_TIMEOUT))
.map_err(|e| TensorError::new(&format!("relay: upstream set_read_timeout: {e}")))?;
let up_write = upstream;
let mut rank_writes: HashMap<u32, TcpStream> = HashMap::with_capacity(rank_streams.len());
let mut rank_reads: Vec<(u32, TcpStream)> = Vec::with_capacity(rank_streams.len());
for (rank, stream) in rank_streams {
let write_half = stream
.try_clone()
.map_err(|e| TensorError::new(&format!("relay: rank {rank} try_clone failed: {e}")))?;
stream
.set_read_timeout(Some(POLL_TIMEOUT))
.map_err(|e| TensorError::new(&format!("relay: rank {rank} set_read_timeout: {e}")))?;
rank_writes.insert(rank, write_half);
rank_reads.push((rank, stream));
}
let mut threads: Vec<JoinHandle<()>> = Vec::with_capacity(rank_reads.len() + 2);
{
let shutdown = Arc::clone(&shutdown);
threads.push(
thread::Builder::new()
.name("flodl-relay-out".into())
.spawn(move || outbound_writer(up_write, rx, salt, shutdown))
.map_err(|e| TensorError::new(&format!("relay: spawn outbound writer: {e}")))?,
);
}
{
let shutdown = Arc::clone(&shutdown);
let fold = fold.clone();
let tx = tx.clone();
threads.push(
thread::Builder::new()
.name("flodl-relay-up".into())
.spawn(move || upstream_reader(up_read, rank_writes, salt, shutdown, fold, tx))
.map_err(|e| TensorError::new(&format!("relay: spawn upstream reader: {e}")))?,
);
}
for (rank, stream) in rank_reads {
let tx = tx.clone();
let shutdown = Arc::clone(&shutdown);
let active = Arc::clone(&active);
let fold = fold.clone();
threads.push(
thread::Builder::new()
.name(format!("flodl-relay-r{rank}"))
.spawn(move || rank_reader(rank, stream, tx, shutdown, active, fold))
.map_err(|e| TensorError::new(&format!("relay: spawn rank {rank} reader: {e}")))?,
);
}
drop(tx);
Ok(threads)
}
fn rank_reader(
rank: u32,
mut stream: TcpStream,
tx: mpsc::SyncSender<MuxRecord>,
shutdown: Arc<AtomicBool>,
active: Arc<AtomicUsize>,
fold: Option<Arc<FoldCtx>>,
) {
loop {
if shutdown.load(Ordering::SeqCst) {
break;
}
match try_read_len_framed(&mut stream) {
Ok(LenFramedRead::Blob(blob)) => match &fold {
None => {
if tx.send(MuxRecord::data(rank, blob)).is_err() {
break; }
}
Some(ctx) => {
if let Err(e) = ctx.deposit(rank, &blob, &tx) {
eprintln!("relay fold: rank {rank}: {e}");
shutdown.store(true, Ordering::SeqCst);
break;
}
}
},
Ok(LenFramedRead::WouldBlock) => {
}
Ok(LenFramedRead::Eof) => {
if let Some(ctx) = &fold {
if let Err(e) = ctx.mark_dead(rank, &tx) {
eprintln!("relay fold: rank {rank} exit-fold: {e}");
shutdown.store(true, Ordering::SeqCst);
}
}
let _ = tx.send(MuxRecord::control(RelayControlMsg::RankExit { rank }));
break;
}
Err(_) => {
if let Some(ctx) = &fold {
if let Err(e) = ctx.mark_dead(rank, &tx) {
eprintln!("relay fold: rank {rank} error-fold: {e}");
shutdown.store(true, Ordering::SeqCst);
}
}
let _ = tx.send(MuxRecord::control(RelayControlMsg::RankExit { rank }));
break;
}
}
}
if active.fetch_sub(1, Ordering::SeqCst) == 1 {
shutdown.store(true, Ordering::SeqCst);
}
}
fn outbound_writer(
mut up_write: TcpStream,
rx: mpsc::Receiver<MuxRecord>,
salt: SessionSalt,
shutdown: Arc<AtomicBool>,
) {
loop {
match rx.recv_timeout(POLL_TIMEOUT) {
Ok(rec) => {
if rec.write_to(&mut up_write, &salt).is_err() {
shutdown.store(true, Ordering::SeqCst);
break;
}
}
Err(RecvTimeoutError::Timeout) => {
if shutdown.load(Ordering::SeqCst) {
break;
}
}
Err(RecvTimeoutError::Disconnected) => break,
}
}
}
fn upstream_reader(
mut up_read: TcpStream,
mut rank_writes: HashMap<u32, TcpStream>,
salt: SessionSalt,
shutdown: Arc<AtomicBool>,
fold: Option<Arc<FoldCtx>>,
tx: mpsc::SyncSender<MuxRecord>,
) {
loop {
if shutdown.load(Ordering::SeqCst) {
break;
}
match MuxRecord::try_read_from(&mut up_read, &salt) {
Ok(MuxRead::Record(MuxRecord::Data { rank, payload })) => {
if let Some(w) = rank_writes.get_mut(&rank) {
let _ = write_len_framed(w, &payload);
}
}
Ok(MuxRead::Record(MuxRecord::Broadcast { payload })) => {
match &fold {
Some(ctx) => ctx.fan_out(&payload, &mut rank_writes),
None => {
eprintln!(
"relay: Broadcast record on the control channel; \
dropping (desynced controller?)"
);
}
}
}
Ok(MuxRead::Record(MuxRecord::Control(
RelayControlMsg::DeclareDead { rank },
))) => {
if let Some(ctx) = &fold {
if let Err(e) = ctx.mark_dead(rank, &tx) {
eprintln!("relay fold: DeclareDead({rank}): {e}");
shutdown.store(true, Ordering::SeqCst);
break;
}
}
}
Ok(MuxRead::Record(MuxRecord::Control(_))) => {
}
Ok(MuxRead::Record(MuxRecord::HostFrame { .. })) => {
eprintln!("relay: unexpected HostFrame from controller; dropping");
}
Ok(MuxRead::WouldBlock) => {
}
Ok(MuxRead::Eof) => {
shutdown.store(true, Ordering::SeqCst);
break;
}
Err(_) => {
shutdown.store(true, Ordering::SeqCst);
break;
}
}
}
}
#[cfg(test)]
#[path = "agent_tests.rs"]
mod tests;