use super::*;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum RoundKind {
#[default]
Model,
Control,
}
impl RoundKind {
fn to_wire(self) -> u8 {
match self {
RoundKind::Model => 0,
RoundKind::Control => 1,
}
}
fn from_wire(b: u8) -> Result<Self> {
match b {
0 => Ok(RoundKind::Model),
1 => Ok(RoundKind::Control),
other => Err(TensorError::new(&format!(
"cluster_controller: unknown RoundKind wire byte {other}"
))),
}
}
}
#[derive(Debug, Clone, PartialEq, Default)]
pub struct RoundFrame {
pub tensors: Vec<TensorPayload>,
pub kind: RoundKind,
pub weight: f64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct TensorPayload {
pub dtype: u8,
pub shape: Vec<u32>,
pub bytes: Vec<u8>,
}
impl TensorPayload {
pub fn numel(&self) -> usize {
self.shape.iter().map(|d| *d as usize).product()
}
}
pub(crate) fn read_round_frame<R: Read>(
stream: &mut R,
salt: &SessionSalt,
) -> Result<Option<RoundFrame>> {
let mut tensors = Vec::new();
match read_round_frame_streamed(stream, salt, &mut |_, payload| {
tensors.push(payload);
Ok(())
})? {
Some((kind, weight)) => Ok(Some(RoundFrame {
tensors,
kind,
weight,
})),
None => Ok(None),
}
}
pub(crate) fn read_round_frame_streamed<R: Read>(
stream: &mut R,
salt: &SessionSalt,
sink: &mut dyn FnMut(usize, TensorPayload) -> Result<()>,
) -> Result<Option<(RoundKind, f64)>> {
let mut mac = HMAC::new(salt.as_slice());
let mut hdr = [0u8; 8];
match stream.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!(
"cluster_controller: frame header read failed: {e}"
)));
}
}
mac.update(hdr);
let magic = u32::from_le_bytes(hdr[0..4].try_into().unwrap());
if magic != ROUND_FRAME_MAGIC {
return Err(TensorError::new(&format!(
"cluster_controller: frame magic 0x{magic:08x} != 0x{ROUND_FRAME_MAGIC:08x}"
)));
}
let num_tensors = u32::from_le_bytes(hdr[4..8].try_into().unwrap()) as usize;
if num_tensors > MAX_ROUND_FRAME_TENSORS {
return Err(TensorError::new(&format!(
"cluster_controller: frame claims {num_tensors} tensors \
(> {MAX_ROUND_FRAME_TENSORS}); corrupt or hostile peer"
)));
}
let mut kind_byte = [0u8; 1];
stream.read_exact(&mut kind_byte).map_err(|e| {
TensorError::new(&format!("cluster_controller: frame kind read failed: {e}"))
})?;
mac.update(kind_byte);
let kind = RoundKind::from_wire(kind_byte[0])?;
let mut weight_bytes = [0u8; 8];
stream.read_exact(&mut weight_bytes).map_err(|e| {
TensorError::new(&format!("cluster_controller: frame weight read failed: {e}"))
})?;
mac.update(weight_bytes);
let weight = f64::from_le_bytes(weight_bytes);
if !weight.is_finite() || weight < 0.0 {
return Err(TensorError::new(&format!(
"cluster_controller: frame weight {weight} is not a finite non-negative \
realized-work mass"
)));
}
let mut total_bytes: usize = 0;
for ti in 0..num_tensors {
let mut meta = [0u8; 2];
stream.read_exact(&mut meta).map_err(|e| {
TensorError::new(&format!(
"cluster_controller: tensor[{ti}] meta read failed: {e}"
))
})?;
mac.update(meta);
let dtype = meta[0];
let ndim = meta[1] as usize;
let mut shape = Vec::with_capacity(ndim);
for _ in 0..ndim {
let mut d = [0u8; 4];
stream.read_exact(&mut d).map_err(|e| {
TensorError::new(&format!(
"cluster_controller: tensor[{ti}] shape read failed: {e}"
))
})?;
mac.update(d);
shape.push(u32::from_le_bytes(d));
}
let mut nb = [0u8; 8];
stream.read_exact(&mut nb).map_err(|e| {
TensorError::new(&format!(
"cluster_controller: tensor[{ti}] nbytes read failed: {e}"
))
})?;
mac.update(nb);
let nbytes = u64::from_le_bytes(nb) as usize;
total_bytes = total_bytes.saturating_add(nbytes);
if total_bytes > crate::distributed::wire::frame_ceiling() {
return Err(TensorError::new(&format!(
"cluster_controller: tensor[{ti}] pushes frame past the \
{} byte ceiling; corrupt or hostile peer, or a model that \
has outgrown the frame ceiling",
crate::distributed::wire::frame_ceiling()
)));
}
let bytes = crate::distributed::wire::read_exact_incremental(stream, nbytes)
.map_err(|e| {
TensorError::new(&format!(
"cluster_controller: tensor[{ti}] data read failed: {e}"
))
})?;
mac.update(&bytes);
sink(
ti,
TensorPayload {
dtype,
shape,
bytes,
},
)?;
}
let mut footer = [0u8; 8];
stream.read_exact(&mut footer).map_err(|e| {
TensorError::new(&format!(
"cluster_controller: frame HMAC footer read failed: {e} \
(sender at PROTOCOL_VERSION < 2, or stream truncated mid-frame)"
))
})?;
let received = u64::from_le_bytes(footer);
let computed_full: [u8; 32] = mac.finalize();
let computed = u64::from_le_bytes(computed_full[0..8].try_into().unwrap());
if computed != received {
return Err(TensorError::new(&format!(
"cluster_controller: RoundFrame HMAC verification failed (computed \
0x{computed:016x}, wire carried 0x{received:016x}); session salt \
disagreement, tampered frame, or payload corruption"
)));
}
Ok(Some((kind, weight)))
}
pub(crate) struct PayloadPart<'a> {
pub dtype: u8,
pub shape: &'a [u32],
pub nbytes: u64,
}
pub(crate) fn round_frame_wire_len(parts: &[PayloadPart<'_>]) -> u64 {
let mut len: u64 = 8 + 1 + 8; for p in parts {
len += 2 + 4 * p.shape.len() as u64 + 8 + p.nbytes;
}
len + 8 }
pub(crate) struct MacTee<'a, W: Write> {
inner: &'a mut W,
mac: &'a mut HMAC,
written: u64,
}
impl<W: Write> Write for MacTee<'_, W> {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
let n = self.inner.write(buf)?;
self.mac.update(&buf[..n]);
self.written += n as u64;
Ok(n)
}
fn flush(&mut self) -> std::io::Result<()> {
self.inner.flush()
}
}
pub(crate) fn write_round_frame_streamed<W: Write>(
stream: &mut W,
kind: RoundKind,
weight: f64,
parts: &[PayloadPart<'_>],
salt: &SessionSalt,
emit: &mut dyn FnMut(usize, &mut MacTee<'_, W>) -> Result<()>,
) -> Result<()> {
let mut mac = HMAC::new(salt.as_slice());
let mut hdr = [0u8; 8];
hdr[0..4].copy_from_slice(&ROUND_FRAME_MAGIC.to_le_bytes());
hdr[4..8].copy_from_slice(&(parts.len() as u32).to_le_bytes());
stream.write_all(&hdr).map_err(|e| {
TensorError::new(&format!("cluster_controller: frame header write failed: {e}"))
})?;
mac.update(hdr);
let kind_byte = [kind.to_wire()];
stream.write_all(&kind_byte).map_err(|e| {
TensorError::new(&format!("cluster_controller: frame kind write failed: {e}"))
})?;
mac.update(kind_byte);
let weight_bytes = weight.to_le_bytes();
stream.write_all(&weight_bytes).map_err(|e| {
TensorError::new(&format!("cluster_controller: frame weight write failed: {e}"))
})?;
mac.update(weight_bytes);
for (ti, p) in parts.iter().enumerate() {
let meta = [p.dtype, p.shape.len() as u8];
stream.write_all(&meta).map_err(|e| {
TensorError::new(&format!(
"cluster_controller: tensor[{ti}] meta write failed: {e}"
))
})?;
mac.update(meta);
for d in p.shape {
let d_bytes = d.to_le_bytes();
stream.write_all(&d_bytes).map_err(|e| {
TensorError::new(&format!(
"cluster_controller: tensor[{ti}] shape write failed: {e}"
))
})?;
mac.update(d_bytes);
}
let nb_bytes = p.nbytes.to_le_bytes();
stream.write_all(&nb_bytes).map_err(|e| {
TensorError::new(&format!(
"cluster_controller: tensor[{ti}] nbytes write failed: {e}"
))
})?;
mac.update(nb_bytes);
let mut tee = MacTee {
inner: stream,
mac: &mut mac,
written: 0,
};
emit(ti, &mut tee).map_err(|e| {
TensorError::new(&format!(
"cluster_controller: tensor[{ti}] data write failed: {e}"
))
})?;
if tee.written != p.nbytes {
return Err(TensorError::new(&format!(
"cluster_controller: tensor[{ti}] emitter wrote {} bytes, \
declared {} — frame length already committed, stream is torn",
tee.written, p.nbytes
)));
}
}
let computed_full: [u8; 32] = mac.finalize();
let mut footer = [0u8; 8];
footer.copy_from_slice(&computed_full[0..8]);
stream.write_all(&footer).map_err(|e| {
TensorError::new(&format!("cluster_controller: frame HMAC footer write failed: {e}"))
})?;
stream
.flush()
.map_err(|e| TensorError::new(&format!("cluster_controller: frame flush failed: {e}")))?;
Ok(())
}
pub(crate) fn write_round_frame<W: Write>(
stream: &mut W,
frame: &RoundFrame,
salt: &SessionSalt,
) -> Result<()> {
let parts: Vec<PayloadPart<'_>> = frame
.tensors
.iter()
.map(|t| PayloadPart {
dtype: t.dtype,
shape: &t.shape,
nbytes: t.bytes.len() as u64,
})
.collect();
write_round_frame_streamed(
stream,
frame.kind,
frame.weight,
&parts,
salt,
&mut |ti, tee| {
tee.write_all(&frame.tensors[ti].bytes)
.map_err(|e| TensorError::new(&e.to_string()))
},
)
}
pub(crate) fn sum_frames(frames: &[&RoundFrame]) -> Result<RoundFrame> {
let Some(ref_frame) = frames.first() else {
return Err(TensorError::new(
"cluster_controller: sum_frames called with no frames",
));
};
let w_sum: f64 = frames.iter().map(|f| f.weight).sum();
for (i, f) in frames.iter().enumerate().skip(1) {
if f.kind != ref_frame.kind {
return Err(TensorError::new(&format!(
"cluster_controller: frame {i} kind {:?} != frame 0 kind {:?} \
(desynced reduce rounds)",
f.kind, ref_frame.kind
)));
}
if f.tensors.len() != ref_frame.tensors.len() {
return Err(TensorError::new(&format!(
"cluster_controller: frame {i} carries {} tensors; frame 0 carries {}",
f.tensors.len(),
ref_frame.tensors.len()
)));
}
for (ti, (a, b)) in ref_frame.tensors.iter().zip(f.tensors.iter()).enumerate() {
if a.dtype != b.dtype {
return Err(TensorError::new(&format!(
"cluster_controller: frame {i} tensor[{ti}] dtype {} != frame 0 dtype {}",
b.dtype, a.dtype
)));
}
if a.shape != b.shape {
return Err(TensorError::new(&format!(
"cluster_controller: frame {i} tensor[{ti}] shape {:?} != frame 0 shape {:?}",
b.shape, a.shape
)));
}
if a.bytes.len() != b.bytes.len() {
return Err(TensorError::new(&format!(
"cluster_controller: frame {i} tensor[{ti}] nbytes {} != frame 0 nbytes {}",
b.bytes.len(),
a.bytes.len()
)));
}
}
}
let mut out_tensors = Vec::with_capacity(ref_frame.tensors.len());
for ti in 0..ref_frame.tensors.len() {
let dtype = ref_frame.tensors[ti].dtype;
let elem = payload_element_size(dtype).map_err(|e| {
TensorError::new(&format!("cluster_controller: tensor[{ti}]: {e}"))
})?;
let shape = ref_frame.tensors[ti].shape.clone();
let numel = ref_frame.tensors[ti].numel();
if numel * elem != ref_frame.tensors[ti].bytes.len() {
return Err(TensorError::new(&format!(
"cluster_controller: tensor[{ti}] shape {shape:?} numel*element_size {} != nbytes {}",
numel * elem,
ref_frame.tensors[ti].bytes.len()
)));
}
let mut accum: Vec<f32> = vec![0.0; numel];
for f in frames.iter() {
accumulate_payload_into(&f.tensors[ti], &mut accum)?;
}
out_tensors.push(TensorPayload {
dtype,
shape,
bytes: f32_slice_to_payload_bytes(&accum, dtype)?,
});
}
Ok(RoundFrame {
tensors: out_tensors,
kind: ref_frame.kind,
weight: w_sum,
})
}
pub(super) fn reduce_realized_work(frames: &[Option<RoundFrame>]) -> Result<RoundFrame> {
let accepted: Vec<&RoundFrame> = frames.iter().filter_map(|f| f.as_ref()).collect();
if accepted.is_empty() {
return Err(TensorError::new(
"cluster_controller: reduce_realized_work called with no accepted frames \
(all participants dead — caller should not have reached this point)",
));
}
let mut summed = sum_frames(&accepted)?;
if matches!(summed.kind, RoundKind::Model)
&& crate::distributed::realized_work::is_realized(summed.weight)
{
let inv = (1.0 / summed.weight) as f32;
for payload in &mut summed.tensors {
scale_payload(payload, inv)?;
}
}
Ok(summed)
}
pub(crate) fn f32_to_bf16_bits(x: f32) -> u16 {
let bits = x.to_bits();
if x.is_nan() {
return ((bits >> 16) as u16) | 0x0040;
}
let round_bit = (bits >> 16) & 1;
((bits + 0x7FFF + round_bit) >> 16) as u16
}
pub(crate) fn bf16_bits_to_f32(b: u16) -> f32 {
f32::from_bits((b as u32) << 16)
}
pub(crate) fn payload_element_size(dtype: u8) -> Result<usize> {
match dtype {
DTYPE_F32 => Ok(4),
DTYPE_BF16 => Ok(2),
other => Err(TensorError::new(&format!(
"unsupported wire dtype tag {other} (0 = f32, 1 = bf16); extend \
round_frame.rs::payload_element_size and the codec helpers together"
))),
}
}
pub(crate) fn accumulate_payload_into(
payload: &TensorPayload,
accum: &mut [f32],
) -> Result<()> {
let elem = payload_element_size(payload.dtype)?;
if payload.bytes.len() != accum.len() * elem {
return Err(TensorError::new(&format!(
"payload byte count {} != accumulator numel {} x element size {elem}",
payload.bytes.len(),
accum.len(),
)));
}
match payload.dtype {
DTYPE_F32 => {
for (a, c) in accum.iter_mut().zip(payload.bytes.chunks_exact(4)) {
*a += f32::from_le_bytes([c[0], c[1], c[2], c[3]]);
}
}
DTYPE_BF16 => {
for (a, c) in accum.iter_mut().zip(payload.bytes.chunks_exact(2)) {
*a += bf16_bits_to_f32(u16::from_le_bytes([c[0], c[1]]));
}
}
_ => unreachable!("payload_element_size validated the tag"),
}
Ok(())
}
pub(crate) fn payload_to_f32(payload: &TensorPayload) -> Result<Vec<f32>> {
let elem = payload_element_size(payload.dtype)?;
if payload.bytes.len() % elem != 0 {
return Err(TensorError::new(&format!(
"payload byte count {} not divisible by element size {elem}",
payload.bytes.len(),
)));
}
let mut out = vec![0.0f32; payload.bytes.len() / elem];
accumulate_payload_into(payload, &mut out)?;
Ok(out)
}
pub(crate) fn f32_slice_to_payload_bytes(data: &[f32], dtype: u8) -> Result<Vec<u8>> {
let elem = payload_element_size(dtype)?;
let mut out = Vec::with_capacity(data.len() * elem);
match dtype {
DTYPE_F32 => {
for x in data {
out.extend_from_slice(&x.to_le_bytes());
}
}
DTYPE_BF16 => {
for x in data {
out.extend_from_slice(&f32_to_bf16_bits(*x).to_le_bytes());
}
}
_ => unreachable!("payload_element_size validated the tag"),
}
Ok(out)
}
pub(crate) fn scale_payload(payload: &mut TensorPayload, factor: f32) -> Result<()> {
scale_payload_bytes(&mut payload.bytes, payload.dtype, factor)
}
pub(crate) fn scale_payload_bytes(bytes: &mut [u8], dtype: u8, factor: f32) -> Result<()> {
match dtype {
DTYPE_F32 => {
for c in bytes.chunks_exact_mut(4) {
let v = f32::from_le_bytes([c[0], c[1], c[2], c[3]]) * factor;
c.copy_from_slice(&v.to_le_bytes());
}
Ok(())
}
DTYPE_BF16 => {
for c in bytes.chunks_exact_mut(2) {
let v = bf16_bits_to_f32(u16::from_le_bytes([c[0], c[1]])) * factor;
c.copy_from_slice(&f32_to_bf16_bits(v).to_le_bytes());
}
Ok(())
}
other => Err(TensorError::new(&format!(
"scale_payload: unsupported wire dtype tag {other}"
))),
}
}
#[cfg(test)]
pub(super) fn bytes_as_f32(bytes: &[u8]) -> Result<Vec<f32>> {
if bytes.len() % 4 != 0 {
return Err(TensorError::new(&format!(
"cluster_controller: f32 byte count {} not divisible by 4",
bytes.len()
)));
}
let n = bytes.len() / 4;
let mut out = Vec::with_capacity(n);
for i in 0..n {
let mut b = [0u8; 4];
b.copy_from_slice(&bytes[i * 4..(i + 1) * 4]);
out.push(f32::from_le_bytes(b));
}
Ok(out)
}
#[cfg(test)]
pub(super) fn f32_to_bytes(data: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(data.len() * 4);
for x in data {
out.extend_from_slice(&x.to_le_bytes());
}
out
}