aetheric-gpu 0.1.0-alpha

Aetheric Silicon: turn this host's RAM into a Digital GPU endpoint
//! Distributed GEMM — splits matrix multiplication across cluster nodes.
//!
//! When multiple Aetheric nodes are discovered on the LAN, GEMM work is
//! partitioned horizontally (row-slices of C) and dispatched to peers via TCP.
//! Each peer computes its slice independently and returns the result.
//!
//! Wire protocol: simple length-prefixed binary messages over TCP:
//!   [4 bytes: msg_type] [4 bytes: payload_len] [payload_len bytes: payload]
//!
//! Message types:
//!   0x01 = GEMMRequest  (m, n, k, A_slice, B)
//!   0x02 = GEMMResponse (C_slice)
//!   0xFF = Ping/Pong (liveness check)

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;

/// Distributed GEMM coordinator — splits work across cluster peers.
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 }
    }

    /// Execute GEMM distributed across cluster: C[m×n] = A[m×k] × B[k×n]
    /// Returns Some(C_flat) on success, None on failure.
    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; // No peers — use local GEMM
        }

        // Include self as a compute node
        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);

        // Dispatch to each peer (including self for local slices)
        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; }

            // Extract A slice for this peer's rows
            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);
                    // Fall back to local compute for this slice
                    let c_slice = local_gemm(slice_m, n, k, &a_slice, b);
                    results.push(c_slice);
                }
            }
        }

        // Compute local slice (self)
        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);
        }

        // Assemble C from all slices
        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)
    }

    /// Dispatch a GEMM to a remote peer via TCP.
    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()?;

        // Build GEMM request payload
        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
        send_message(&mut stream, MSG_GEMM_REQUEST, &payload).ok()?;

        // Receive response
        let (msg_type, response) = read_message(&mut stream).ok()?;
        if msg_type != MSG_GEMM_RESPONSE { return None; }

        // Parse C slice
        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)
    }

    /// Start the worker TCP server (runs in a background thread).
    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 => {
            // Parse 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, &[]);
        }
        _ => {}
    }
}

/// Local GEMM compute (reference implementation).
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]; // 2×2
        let b = vec![5.0, 6.0, 7.0, 8.0]; // 2×2
        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]);
    }
}