use std::collections::HashMap;
use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream, SocketAddr, ToSocketAddrs};
use std::sync::Arc;
use parking_lot::RwLock;
use crate::ClusterNode;
const MSG_GEMM_REQUEST: u32 = 0x01;
const MSG_GEMM_RESPONSE: u32 = 0x02;
const MSG_PING: u32 = 0xFF;
const MSG_PONG: u32 = 0xFE;
const WORKER_PORT: u16 = 50051;
const RESPONSE_TIMEOUT_MS: u64 = 30000;
pub struct DistributedGemm {
peers: Arc<RwLock<HashMap<String, ClusterNode>>>,
local_id: String,
}
impl DistributedGemm {
pub fn new(peers: Arc<RwLock<HashMap<String, ClusterNode>>>, local_id: String) -> Self {
Self { peers, local_id }
}
pub fn gemm(&self, m: usize, n: usize, k: usize, a: &[f32], b: &[f32]) -> Option<Vec<f32>> {
let peers = self.peers.read().clone();
let peer_list: Vec<&ClusterNode> = peers.values().collect();
if peer_list.is_empty() {
return None; }
let total_nodes = peer_list.len() + 1;
let rows_per_node = m / total_nodes;
let remainder = m % total_nodes;
let mut results: Vec<Vec<f32>> = Vec::with_capacity(total_nodes);
for (i, peer) in peer_list.iter().enumerate() {
let row_start = i * rows_per_node;
let row_end = if i == total_nodes - 1 {
row_start + rows_per_node + remainder
} else {
row_start + rows_per_node
};
let slice_m = row_end - row_start;
if slice_m == 0 { continue; }
let a_slice: Vec<f32> = (row_start..row_end)
.flat_map(|row| (0..k).map(move |col| a[row * k + col]))
.collect();
let peer_result = self.dispatch_gemm(peer, slice_m, n, k, &a_slice, b);
match peer_result {
Some(c_slice) => results.push(c_slice),
None => {
log::warn!("[distributed] Peer {} failed, falling back to local", peer.name);
let c_slice = local_gemm(slice_m, n, k, &a_slice, b);
results.push(c_slice);
}
}
}
let self_idx = peer_list.len();
let row_start = self_idx * rows_per_node;
let row_end = row_start + rows_per_node + remainder;
let slice_m = row_end - row_start;
if slice_m > 0 {
let a_slice: Vec<f32> = (row_start..row_end)
.flat_map(|row| (0..k).map(move |col| a[row * k + col]))
.collect();
let c_slice = local_gemm(slice_m, n, k, &a_slice, b);
results.push(c_slice);
}
let mut c = vec![0.0f32; m * n];
let mut offset = 0;
for slice in &results {
let len = slice.len();
c[offset..offset + len].copy_from_slice(slice);
offset += len;
}
Some(c)
}
fn dispatch_gemm(&self, peer: &ClusterNode, m: usize, n: usize, k: usize, a: &[f32], b: &[f32]) -> Option<Vec<f32>> {
let addr = format!("{}:{}", peer.addr, WORKER_PORT);
let addr = addr.to_socket_addrs().ok()?.next()?;
let mut stream = TcpStream::connect_timeout(
&addr,
std::time::Duration::from_millis(5000),
).ok()?;
stream.set_read_timeout(Some(std::time::Duration::from_millis(RESPONSE_TIMEOUT_MS))).ok()?;
stream.set_write_timeout(Some(std::time::Duration::from_millis(5000))).ok()?;
let mut payload = Vec::new();
payload.extend_from_slice(&(m as u32).to_le_bytes());
payload.extend_from_slice(&(n as u32).to_le_bytes());
payload.extend_from_slice(&(k as u32).to_le_bytes());
for &v in a { payload.extend_from_slice(&v.to_le_bytes()); }
for &v in b { payload.extend_from_slice(&v.to_le_bytes()); }
send_message(&mut stream, MSG_GEMM_REQUEST, &payload).ok()?;
let (msg_type, response) = read_message(&mut stream).ok()?;
if msg_type != MSG_GEMM_RESPONSE { return None; }
let c: Vec<f32> = response.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
Some(c)
}
pub fn start_worker(&self) {
let listener = match TcpListener::bind(format!("0.0.0.0:{}", WORKER_PORT)) {
Ok(l) => l,
Err(e) => {
log::warn!("[distributed] Failed to bind worker on port {}: {}", WORKER_PORT, e);
return;
}
};
log::info!("[distributed] Worker listening on port {}", WORKER_PORT);
for stream in listener.incoming() {
if let Ok(mut stream) = stream {
std::thread::spawn(move || {
handle_worker_request(&mut stream);
});
}
}
}
}
fn handle_worker_request(stream: &mut TcpStream) {
let (msg_type, payload) = match read_message(stream) {
Ok(v) => v,
Err(_) => return,
};
match msg_type {
MSG_GEMM_REQUEST => {
if payload.len() < 12 { return; }
let m = u32::from_le_bytes([payload[0], payload[1], payload[2], payload[3]]) as usize;
let n = u32::from_le_bytes([payload[4], payload[5], payload[6], payload[7]]) as usize;
let k = u32::from_le_bytes([payload[8], payload[9], payload[10], payload[11]]) as usize;
let a_len = m * k * 4;
let b_len = k * n * 4;
if payload.len() < 12 + a_len + b_len { return; }
let a: Vec<f32> = payload[12..12 + a_len].chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
let b: Vec<f32> = payload[12 + a_len..12 + a_len + b_len].chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
let c = local_gemm(m, n, k, &a, &b);
let mut response = Vec::new();
for &v in &c { response.extend_from_slice(&v.to_le_bytes()); }
let _ = send_message(stream, MSG_GEMM_RESPONSE, &response);
}
MSG_PING => {
let _ = send_message(stream, MSG_PONG, &[]);
}
_ => {}
}
}
fn local_gemm(m: usize, n: usize, k: usize, a: &[f32], b: &[f32]) -> Vec<f32> {
let mut c = vec![0.0f32; m * n];
for i in 0..m {
for j in 0..n {
let mut sum = 0.0f32;
for p in 0..k {
sum += a[i * k + p] * b[p * n + j];
}
c[i * n + j] = sum;
}
}
c
}
fn send_message(stream: &mut TcpStream, msg_type: u32, payload: &[u8]) -> std::io::Result<()> {
stream.write_all(&msg_type.to_le_bytes())?;
stream.write_all(&(payload.len() as u32).to_le_bytes())?;
stream.write_all(payload)?;
stream.flush()?;
Ok(())
}
fn read_message(stream: &mut TcpStream) -> std::io::Result<(u32, Vec<u8>)> {
let mut buf = [0u8; 4];
stream.read_exact(&mut buf)?;
let msg_type = u32::from_le_bytes(buf);
stream.read_exact(&mut buf)?;
let len = u32::from_le_bytes(buf) as usize;
let mut payload = vec![0u8; len];
stream.read_exact(&mut payload)?;
Ok((msg_type, payload))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_local_gemm() {
let a = vec![1.0, 2.0, 3.0, 4.0]; let b = vec![5.0, 6.0, 7.0, 8.0]; let c = local_gemm(2, 2, 2, &a, &b);
assert_eq!(c, vec![19.0, 22.0, 43.0, 50.0]);
}
#[test]
fn test_send_read_message() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
std::thread::spawn(move || {
let mut stream = listener.incoming().next().unwrap().unwrap();
let (msg_type, payload) = read_message(&mut stream).unwrap();
assert_eq!(msg_type, 0x01);
assert_eq!(payload, vec![1, 2, 3]);
send_message(&mut stream, 0x02, &vec![4, 5, 6]).unwrap();
});
let mut stream = TcpStream::connect(format!("127.0.0.1:{}", port)).unwrap();
send_message(&mut stream, 0x01, &vec![1, 2, 3]).unwrap();
let (msg_type, payload) = read_message(&mut stream).unwrap();
assert_eq!(msg_type, 0x02);
assert_eq!(payload, vec![4, 5, 6]);
}
}