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, RoundFrame};
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 {
frames: HashMap<u32, RoundFrame>,
dead: HashSet<u32>,
}
impl FoldCtx {
fn new(local_ranks: Vec<u32>, salt: SessionSalt) -> Self {
FoldCtx {
inner: Mutex::new(FoldInner {
frames: HashMap::with_capacity(local_ranks.len()),
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,
frame: RoundFrame,
frame_bytes: u64,
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.frames.insert(rank, frame).is_some() {
return Err(TensorError::new(&format!(
"relay fold: rank {rank} deposited twice in one round \
(rank↔relay ping-pong violated; stream desynced)"
)));
}
self.bytes_in.fetch_add(frame_bytes, 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<Vec<RoundFrame>> {
let complete = self
.local_ranks
.iter()
.all(|r| inner.dead.contains(r) || inner.frames.contains_key(r));
if !complete || inner.frames.is_empty() {
return None;
}
Some(inner.frames.drain().map(|(_, f)| f).collect())
}
fn fold_and_ship(
&self,
taken: Option<Vec<RoundFrame>>,
tx: &mpsc::SyncSender<MuxRecord>,
) -> Result<()> {
let Some(frames) = taken else {
return Ok(());
};
let refs: Vec<&RoundFrame> = frames.iter().collect();
let folded = controller::sum_frames(&refs)?;
let mut buf = Vec::new();
controller::write_round_frame(&mut buf, &folded, &self.salt)?;
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) => {
let parsed = controller::read_round_frame(
&mut blob.as_slice(),
&ctx.salt,
);
let deposit = match parsed {
Ok(Some(frame)) => {
ctx.deposit(rank, frame, blob.len() as u64, &tx)
}
Ok(None) => Err(TensorError::new(&format!(
"relay fold: truncated RoundFrame from rank {rank}"
))),
Err(e) => Err(e),
};
if let Err(e) = deposit {
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;