use std::io::{ErrorKind, Read, Write};
use serde::{Deserialize, Serialize};
use crate::distributed::wire::{hmac_sha256_64, SessionSalt};
use crate::tensor::{Result, TensorError};
pub const MUX_RECORD_MAGIC: u32 = 0xF10D_17D0;
pub const MUX_PROTOCOL_VERSION: u32 = 1;
const MUX_HEADER_LEN: usize = 25;
const REC_DATA: u8 = 0x01;
const REC_CONTROL: u8 = 0x02;
const REC_HOST_FRAME: u8 = 0x03;
const REC_BROADCAST: u8 = 0x04;
fn bincode_config() -> impl bincode::config::Config {
bincode::config::standard()
}
fn encode<T: Serialize>(value: &T) -> Result<Vec<u8>> {
bincode::serde::encode_to_vec(value, bincode_config())
.map_err(|e| TensorError::new(&format!("relay_mux: bincode encode failed: {e}")))
}
fn decode<T: for<'de> Deserialize<'de>>(bytes: &[u8]) -> Result<T> {
let (v, _used) = bincode::serde::decode_from_slice(bytes, bincode_config())
.map_err(|e| TensorError::new(&format!("relay_mux: bincode decode failed: {e}")))?;
Ok(v)
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum RelayControlMsg {
Hello {
host: String,
ranks: Vec<u32>,
},
HelloAck,
RankExit { rank: u32 },
DeclareDead { rank: u32 },
}
#[derive(Debug, Clone, PartialEq)]
pub enum MuxRecord {
Data { rank: u32, payload: Vec<u8> },
Control(RelayControlMsg),
HostFrame { payload: Vec<u8> },
Broadcast { payload: Vec<u8> },
}
use crate::distributed::wire::frame_ceiling;
impl MuxRecord {
pub fn data(rank: u32, payload: Vec<u8>) -> Self {
MuxRecord::Data { rank, payload }
}
pub fn control(msg: RelayControlMsg) -> Self {
MuxRecord::Control(msg)
}
pub fn host_frame(payload: Vec<u8>) -> Self {
MuxRecord::HostFrame { payload }
}
pub fn broadcast(payload: Vec<u8>) -> Self {
MuxRecord::Broadcast { payload }
}
fn parts(&self) -> Result<(u8, u32, std::borrow::Cow<'_, [u8]>)> {
match self {
MuxRecord::Data { rank, payload } => {
Ok((REC_DATA, *rank, std::borrow::Cow::Borrowed(payload)))
}
MuxRecord::Control(msg) => {
Ok((REC_CONTROL, 0, std::borrow::Cow::Owned(encode(msg)?)))
}
MuxRecord::HostFrame { payload } => {
Ok((REC_HOST_FRAME, 0, std::borrow::Cow::Borrowed(payload)))
}
MuxRecord::Broadcast { payload } => {
Ok((REC_BROADCAST, 0, std::borrow::Cow::Borrowed(payload)))
}
}
}
pub fn write_to<W: Write>(&self, w: &mut W, salt: &SessionSalt) -> Result<()> {
let (kind, rank, payload) = self.parts()?;
let payload_len = u32::try_from(payload.len()).map_err(|_| {
TensorError::new(&format!(
"relay_mux: payload too large: {} bytes (max {})",
payload.len(),
u32::MAX
))
})?;
let mut hdr = [0u8; MUX_HEADER_LEN];
hdr[0..4].copy_from_slice(&MUX_RECORD_MAGIC.to_le_bytes());
hdr[4..8].copy_from_slice(&MUX_PROTOCOL_VERSION.to_le_bytes());
hdr[8] = kind;
hdr[9..13].copy_from_slice(&rank.to_le_bytes());
hdr[13..17].copy_from_slice(&payload_len.to_le_bytes());
let auth_tag = hmac_sha256_64_2(salt, &hdr[0..17], &payload);
hdr[17..25].copy_from_slice(&auth_tag.to_le_bytes());
match self {
MuxRecord::HostFrame { .. } | MuxRecord::Broadcast { .. } => {
w.write_all(&hdr).map_err(|e| {
TensorError::new(&format!("relay_mux: record header write failed: {e}"))
})?;
w.write_all(&payload).map_err(|e| {
TensorError::new(&format!("relay_mux: record write failed: {e}"))
})?;
}
MuxRecord::Data { .. } | MuxRecord::Control(_) => {
let mut frame = Vec::with_capacity(MUX_HEADER_LEN + payload.len());
frame.extend_from_slice(&hdr);
frame.extend_from_slice(&payload);
w.write_all(&frame).map_err(|e| {
TensorError::new(&format!("relay_mux: record write failed: {e}"))
})?;
}
}
Ok(())
}
pub fn read_from<R: Read>(r: &mut R, salt: &SessionSalt) -> Result<Option<Self>> {
let mut hdr = [0u8; MUX_HEADER_LEN];
match r.read_exact(&mut hdr) {
Ok(()) => {}
Err(e)
if matches!(
e.kind(),
ErrorKind::UnexpectedEof | ErrorKind::ConnectionReset
) =>
{
return Ok(None);
}
Err(e) => {
return Err(TensorError::new(&format!(
"relay_mux: header read failed: {e}"
)));
}
}
Self::finish_read(hdr, r, salt).map(Some)
}
pub fn try_read_from<R: Read>(r: &mut R, salt: &SessionSalt) -> Result<MuxRead> {
let mut hdr = [0u8; MUX_HEADER_LEN];
match read_idle_gate(r)? {
IdleGate::Idle => return Ok(MuxRead::WouldBlock),
IdleGate::Eof => return Ok(MuxRead::Eof),
IdleGate::Byte(b) => hdr[0] = b,
}
fill_committed(r, &mut hdr[1..])?;
Self::finish_read(hdr, r, salt).map(MuxRead::Record)
}
fn finish_read<R: Read>(hdr: [u8; MUX_HEADER_LEN], r: &mut R, salt: &SessionSalt) -> Result<Self> {
let magic = u32::from_le_bytes(hdr[0..4].try_into().unwrap());
if magic != MUX_RECORD_MAGIC {
return Err(TensorError::new(&format!(
"relay_mux: record magic 0x{magic:08x} != 0x{MUX_RECORD_MAGIC:08x}"
)));
}
let version = u32::from_le_bytes(hdr[4..8].try_into().unwrap());
if version != MUX_PROTOCOL_VERSION {
return Err(TensorError::new(&format!(
"relay_mux: record version {version} != {MUX_PROTOCOL_VERSION}"
)));
}
let kind = hdr[8];
let rank = u32::from_le_bytes(hdr[9..13].try_into().unwrap());
let payload_len = u32::from_le_bytes(hdr[13..17].try_into().unwrap()) as usize;
let auth_tag = u64::from_le_bytes(hdr[17..25].try_into().unwrap());
let ceiling = frame_ceiling();
if payload_len > ceiling {
return Err(TensorError::new(&format!(
"relay_mux: record payload_len {payload_len} exceeds the frame \
ceiling {ceiling} (kind=0x{kind:02x}, rank={rank}); corrupt or \
hostile peer, or a model that has outgrown the frame ceiling"
)));
}
let payload = fill_committed_incremental(r, payload_len)?;
let actual = hmac_sha256_64_2(salt, &hdr[0..17], &payload);
if actual != auth_tag {
return Err(TensorError::new(&format!(
"relay_mux: HMAC verification failed (computed 0x{actual:016x}, \
wire carried 0x{auth_tag:016x}); session salt disagreement, \
tampered record, or corruption (kind=0x{kind:02x}, rank={rank}, \
len={payload_len})"
)));
}
match kind {
REC_DATA => Ok(MuxRecord::Data { rank, payload }),
REC_CONTROL => Ok(MuxRecord::Control(decode(&payload)?)),
REC_HOST_FRAME => Ok(MuxRecord::HostFrame { payload }),
REC_BROADCAST => Ok(MuxRecord::Broadcast { payload }),
other => Err(TensorError::new(&format!(
"relay_mux: unknown record kind 0x{other:02x}"
))),
}
}
}
#[derive(Debug)]
pub enum MuxRead {
Record(MuxRecord),
WouldBlock,
Eof,
}
pub fn write_len_framed<W: Write>(w: &mut W, bytes: &[u8]) -> Result<()> {
let len = u32::try_from(bytes.len()).map_err(|_| {
TensorError::new(&format!(
"relay_mux: len-framed blob too large: {} bytes (max {})",
bytes.len(),
u32::MAX
))
})?;
let mut framed = Vec::with_capacity(4 + bytes.len());
framed.extend_from_slice(&len.to_le_bytes());
framed.extend_from_slice(bytes);
w.write_all(&framed)
.map_err(|e| TensorError::new(&format!("relay_mux: len-framed write failed: {e}")))?;
Ok(())
}
pub fn write_len_prefix<W: Write>(w: &mut W, len: u64) -> Result<()> {
let ceiling = frame_ceiling();
if len > ceiling as u64 {
return Err(TensorError::new(&format!(
"relay_mux: len-framed blob length {len} exceeds the frame ceiling \
{ceiling}; model has outgrown the frame ceiling (see \
wire::frame_ceiling)"
)));
}
let len = u32::try_from(len).map_err(|_| {
TensorError::new(&format!(
"relay_mux: len-framed blob too large: {len} bytes (max {})",
u32::MAX
))
})?;
w.write_all(&len.to_le_bytes())
.map_err(|e| TensorError::new(&format!("relay_mux: len prefix write failed: {e}")))
}
pub fn read_len_prefix<R: Read>(r: &mut R) -> Result<Option<usize>> {
let mut len_buf = [0u8; 4];
match r.read_exact(&mut len_buf) {
Ok(()) => {}
Err(e)
if matches!(
e.kind(),
ErrorKind::UnexpectedEof | ErrorKind::ConnectionReset
) =>
{
return Ok(None);
}
Err(e) => {
return Err(TensorError::new(&format!(
"relay_mux: len prefix read failed: {e}"
)));
}
}
let len = u32::from_le_bytes(len_buf) as usize;
let ceiling = frame_ceiling();
if len > ceiling {
return Err(TensorError::new(&format!(
"relay_mux: len-framed blob length {len} exceeds the frame ceiling \
{ceiling}; corrupt or hostile peer"
)));
}
Ok(Some(len))
}
#[cfg_attr(not(test), allow(dead_code))]
pub fn read_len_framed<R: Read>(r: &mut R) -> Result<Option<Vec<u8>>> {
let Some(len) = read_len_prefix(r)? else {
return Ok(None);
};
let body = crate::distributed::wire::read_exact_incremental(r, len)
.map_err(|e| TensorError::new(&format!("relay_mux: len-framed body read failed: {e}")))?;
Ok(Some(body))
}
pub fn try_read_len_framed<R: Read>(r: &mut R) -> Result<LenFramedRead> {
let mut len_buf = [0u8; 4];
match read_idle_gate(r)? {
IdleGate::Idle => return Ok(LenFramedRead::WouldBlock),
IdleGate::Eof => return Ok(LenFramedRead::Eof),
IdleGate::Byte(b) => len_buf[0] = b,
}
fill_committed(r, &mut len_buf[1..])?;
let len = u32::from_le_bytes(len_buf) as usize;
let ceiling = frame_ceiling();
if len > ceiling {
return Err(TensorError::new(&format!(
"relay_mux: len-framed blob length {len} exceeds the frame ceiling \
{ceiling}; corrupt or hostile peer"
)));
}
let body = fill_committed_incremental(r, len)?;
Ok(LenFramedRead::Blob(body))
}
#[derive(Debug)]
pub enum LenFramedRead {
Blob(Vec<u8>),
WouldBlock,
Eof,
}
enum IdleGate {
Idle,
Eof,
Byte(u8),
}
fn read_idle_gate<R: Read>(r: &mut R) -> Result<IdleGate> {
let mut b = [0u8; 1];
loop {
match r.read(&mut b) {
Ok(0) => return Ok(IdleGate::Eof),
Ok(_) => return Ok(IdleGate::Byte(b[0])),
Err(e) if e.kind() == ErrorKind::Interrupted => continue,
Err(e)
if matches!(e.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut) =>
{
return Ok(IdleGate::Idle);
}
Err(e)
if matches!(
e.kind(),
ErrorKind::UnexpectedEof | ErrorKind::ConnectionReset
) =>
{
return Ok(IdleGate::Eof);
}
Err(e) => {
return Err(TensorError::new(&format!(
"relay_mux: idle-gate read failed: {e}"
)));
}
}
}
}
const COMMITTED_READ_STARVATION_SECS: u64 = 60;
fn fill_committed<R: Read>(r: &mut R, buf: &mut [u8]) -> Result<()> {
let mut filled = 0;
let mut last_progress = std::time::Instant::now();
while filled < buf.len() {
match r.read(&mut buf[filled..]) {
Ok(0) => {
return Err(TensorError::new(
"relay_mux: peer closed mid-frame (committed read hit EOF)",
));
}
Ok(n) => {
filled += n;
last_progress = std::time::Instant::now();
}
Err(e)
if matches!(
e.kind(),
ErrorKind::WouldBlock | ErrorKind::TimedOut | ErrorKind::Interrupted
) =>
{
if last_progress.elapsed().as_secs() >= COMMITTED_READ_STARVATION_SECS {
return Err(TensorError::new(&format!(
"relay_mux: committed read starved mid-frame for \
{COMMITTED_READ_STARVATION_SECS}s ({filled}/{} bytes); \
peer presumed gone",
buf.len(),
)));
}
continue;
}
Err(e) => {
return Err(TensorError::new(&format!(
"relay_mux: committed read failed: {e}"
)));
}
}
}
Ok(())
}
fn fill_committed_incremental<R: Read>(r: &mut R, len: usize) -> Result<Vec<u8>> {
let mut buf: Vec<u8> = Vec::new();
while buf.len() < len {
let chunk = (len - buf.len())
.min(crate::distributed::wire::READ_CHUNK);
let old_len = buf.len();
buf.resize(old_len + chunk, 0);
fill_committed(r, &mut buf[old_len..])?;
}
Ok(buf)
}
fn hmac_sha256_64_2(salt: &SessionSalt, a: &[u8], b: &[u8]) -> u64 {
let mut buf = Vec::with_capacity(a.len() + b.len());
buf.extend_from_slice(a);
buf.extend_from_slice(b);
hmac_sha256_64(salt, &buf)
}
#[cfg(test)]
#[path = "mux_tests.rs"]
mod tests;