aetheric-gpu 0.1.0-alpha

Aetheric Silicon: turn this host's RAM into a Digital GPU endpoint
//! LAN cluster discovery — UDP broadcast for finding peer Aetheric nodes.
//!
//! Wire-protocol: every [DISCOVERY_PORT_MS] ms, broadcast a "Hello" packet
//! containing JSON describing this node. Each peer responds with its own
//! hello packet on the same port. We maintain a live peer table.
//!
//! "Hi" packet format: Aetheric Hello Protocol v1 (compact JSON):
//!   {
//!     "magic":   "AETH-v1",
//!     "id":      "<uuid>",
//!     "name":    "<host>",
//!     "ram_gb":  128,
//!     "pledged": 1,
//!     "comp":    0  // CompressionMode index
//!   }

use std::net::{UdpSocket, SocketAddr};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use parking_lot::Mutex;
use std::collections::HashMap;
use serde::{Deserialize, Serialize};
use crate::{ClusterNode, CompressionMode};

const DISCOVERY_PORT: u16 = 50050;
const HELLO_MAGIC: &[u8; 7] = b"AETH-v1";
const BROADCAST_ADDR: &str = "255.255.255.255";
const LOCALHOST_ADDR: &str = "127.0.0.1";
const DISCOVERY_PERIOD_MS: u64 = 2000;
const PEER_EXPIRY_MS: u64 = 6000;

/// Default discovery port (can be overridden for multi-node on same machine)
pub const DEFAULT_DISCOVERY_PORT: u16 = DISCOVERY_PORT;

#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HelloPacket {
    pub magic: String,
    pub id: String,
    pub name: String,
    pub ram_gb: u64,
    pub pledged: u64,
    pub compression: u8,  // matches CompressionMode variants
}

/// Discovers and tracks peers on LAN.
pub struct ClusterDiscovery {
    pub node_id: String,
    pub node_name: String,
    pub physical_ram_gb: u64,
    pub pledged_gb: u64,
    pub compression: CompressionMode,

    pub peers: Arc<Mutex<HashMap<String, ClusterNode>>>,

    running: AtomicBool,
    last_broadcast_ms: AtomicU64,
    discovery_port: u16,
    /// If true, use localhost instead of broadcast (for local multi-node testing)
    use_localhost: bool,
}

impl ClusterDiscovery {
    pub fn new(ram_gb: u64, pledged_gb: u64, compression: CompressionMode) -> Self {
        Self {
            node_id: uuid::Uuid::new_v4().to_string(),
            node_name: hostname_or_default(),
            physical_ram_gb: ram_gb,
            pledged_gb: pledged_gb,
            compression,
            peers: Arc::new(Mutex::new(HashMap::new())),
            running: AtomicBool::new(false),
            last_broadcast_ms: AtomicU64::new(0),
            discovery_port: DISCOVERY_PORT,
            use_localhost: false,
        }
    }

    /// Create with custom discovery port (for running multiple nodes on same host).
    pub fn with_port(
        ram_gb: u64,
        pledged_gb: u64,
        compression: CompressionMode,
        port: u16,
    ) -> Self {
        let mut s = Self::new(ram_gb, pledged_gb, compression);
        s.discovery_port = port;
        s
    }

    /// Enable localhost mode (sends to 127.0.0.1 instead of broadcast).
    /// Useful for local multi-node testing where broadcast doesn't reach loopback.
    pub fn with_localhost(mut self) -> Self {
        self.use_localhost = true;
        self
    }

    /// Set a custom node name.
    pub fn with_name(mut self, name: String) -> Self {
        self.node_name = name;
        self
    }

    /// Read a peer hello UDP packet, decode as ClusterNode, store it.
    pub fn ingest_packet(&self, packet: HelloPacket, source: SocketAddr, now_ms: u64) -> Option<ClusterNode> {
        if packet.magic != std::str::from_utf8(HELLO_MAGIC).ok()? {
            return None;
        }
        if packet.id == self.node_id {
            return None;
        }

        let compression = match packet.compression {
            0 => CompressionMode::BinaryGemm,
            1 => CompressionMode::Int4,
            2 => CompressionMode::Int8,
            _ => CompressionMode::None,
        };

        let node = ClusterNode {
            id: packet.id.clone(),
            addr: source.to_string(),
            name: packet.name.clone(),
            physical_ram_gb: packet.ram_gb,
            pledged_gb: packet.pledged,
            compression,
            last_seen: now_ms,
        };

        self.peers.lock().insert(packet.id, node.clone());
        Some(node)
    }

    /// Build the hello packet for broadcast.
    pub fn hello_packet(&self) -> HelloPacket {
        HelloPacket {
            magic: String::from(std::str::from_utf8(HELLO_MAGIC).unwrap()),
            id: self.node_id.clone(),
            name: self.node_name.clone(),
            ram_gb: self.physical_ram_gb,
            pledged: self.pledged_gb,
            compression: match self.compression {
                CompressionMode::BinaryGemm => 0,
                CompressionMode::Int4 => 1,
                CompressionMode::Int8 => 2,
                CompressionMode::None => 3,
            },
        }
    }

    /// Drop peers older than the heartbeat window.
    pub fn evict_stale(&self, now_ms: u64) -> Vec<String> {
        let mut peers = self.peers.lock();
        let mut evicted = Vec::new();
        peers.retain(|id, node| {
            let alive = now_ms.saturating_sub(node.last_seen) < PEER_EXPIRY_MS;
            if !alive {
                evicted.push(id.clone());
            }
            alive
        });
        evicted
    }

    /// Live peer count.
    pub fn peer_count(&self) -> usize {
        self.peers.lock().len()
    }

    /// All known peers.
    pub fn peer_list(&self) -> Vec<ClusterNode> {
        self.peers.lock().values().cloned().collect()
    }
}

/// Spawn a listener thread that:
/// - Broadcasts a hello packet every DISCOVERY_PERIOD_MS
/// - Receives and ingests peer hello packets
///
/// Returns a JoinHandle.
pub fn spawn_discovery(discovery: Arc<ClusterDiscovery>) -> std::thread::JoinHandle<()> {
    spawn_discovery_on_port(discovery, DISCOVERY_PORT)
}

/// Same as `spawn_discovery` but binds to a custom port (for local multi-node testing).
pub fn spawn_discovery_on_port(
    discovery: Arc<ClusterDiscovery>, port: u16
) -> std::thread::JoinHandle<()> {
    std::thread::spawn(move || {
        let socket = match UdpSocket::bind(("0.0.0.0", port)) {
            Ok(s) => s,
            Err(e) => {
                log::warn!("cluster discovery: bind failed: {}", e);
                return;
            }
        };
        if let Err(e) = socket.set_broadcast(true) {
            log::warn!("cluster discovery: set_broadcast failed: {}", e);
        }
        socket.set_read_timeout(Some(Duration::from_millis(500))).ok();

        let mut buf = [0u8; 2048];
        let mut next_broadcast = std::time::Instant::now();

        loop {
            match socket.recv_from(&mut buf) {
                Ok((n, src)) => {
                    if let Ok(packet) = serde_json::from_slice::<HelloPacket>(&buf[..n]) {
                        let now_ms = unix_ms_now();
                        let _ = discovery.ingest_packet(packet, src, now_ms);
                    }
                }
                Err(_) => {}
            }

            if next_broadcast.elapsed() > Duration::from_millis(DISCOVERY_PERIOD_MS) {
                let packet = discovery.hello_packet();
                let bytes = match serde_json::to_vec(&packet) {
                    Ok(b) => b,
                    Err(_) => continue,
                };
                // Send to broadcast address (LAN)
                if let Ok(broadcast_addr) = format!("{}:{}", BROADCAST_ADDR, port).parse::<SocketAddr>() {
                    let _ = socket.send_to(&bytes, broadcast_addr);
                }
                // Also send to localhost for local multi-process testing
                if let Ok(local_addr) = format!("{}:{}", LOCALHOST_ADDR, port).parse::<SocketAddr>() {
                    let _ = socket.send_to(&bytes, local_addr);
                }
                next_broadcast = std::time::Instant::now();

                let now_ms = unix_ms_now();
                let evicted = discovery.evict_stale(now_ms);
                for e in evicted {
                    log::debug!("peer {} evicted (heartbeat timeout)", e);
                }
            }
        }
    })
}

/// Spawn discovery in "local only" mode — only uses 127.0.0.1, no LAN broadcast.
/// Useful for CI and local multi-node testing without network permissions.
pub fn spawn_discovery_local(
    discovery: Arc<ClusterDiscovery>, port: u16
) -> std::thread::JoinHandle<()> {
    std::thread::spawn(move || {
        let socket = match UdpSocket::bind(("127.0.0.1", port)) {
            Ok(s) => s,
            Err(e) => {
                log::warn!("cluster discovery local: bind failed: {}", e);
                return;
            }
        };
        socket.set_read_timeout(Some(Duration::from_millis(500))).ok();

        let mut buf = [0u8; 2048];
        let mut next_broadcast = std::time::Instant::now();

        loop {
            match socket.recv_from(&mut buf) {
                Ok((n, src)) => {
                    if let Ok(packet) = serde_json::from_slice::<HelloPacket>(&buf[..n]) {
                        let now_ms = unix_ms_now();
                        let _ = discovery.ingest_packet(packet, src, now_ms);
                    }
                }
                Err(_) => {}
            }

            if next_broadcast.elapsed() > Duration::from_millis(DISCOVERY_PERIOD_MS) {
                let packet = discovery.hello_packet();
                let bytes = match serde_json::to_vec(&packet) {
                    Ok(b) => b,
                    Err(_) => continue,
                };
                if let Ok(local_addr) = format!("{}:{}", LOCALHOST_ADDR, port).parse::<SocketAddr>() {
                    let _ = socket.send_to(&bytes, local_addr);
                }
                next_broadcast = std::time::Instant::now();

                let now_ms = unix_ms_now();
                let evicted = discovery.evict_stale(now_ms);
                for e in evicted {
                    log::debug!("peer {} evicted (heartbeat timeout)", e);
                }
            }
        }
    })
}

pub fn unix_ms_now() -> u64 {
    std::time::SystemTime::now()
        .duration_since(std::time::UNIX_EPOCH)
        .map(|d| d.as_millis() as u64)
        .unwrap_or(0)
}

fn hostname_or_default() -> String {
    std::env::var("HOSTNAME")
        .or_else(|_| std::env::var("COMPUTERNAME"))
        .unwrap_or_else(|_| "aetheric-node".to_string())
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn hello_packet_round_trip() {
        let d = ClusterDiscovery::new(64, 1, CompressionMode::BinaryGemm);
        let p = d.hello_packet();
        let json = serde_json::to_vec(&p).unwrap();
        let p2: HelloPacket = serde_json::from_slice(&json).unwrap();
        assert_eq!(p.id, p2.id);
        assert_eq!(p.magic, p2.magic);
        assert_eq!(p.compression, 0);
    }

    #[test]
    fn peer_table_lifecycle() {
        let d = ClusterDiscovery::new(64, 1, CompressionMode::BinaryGemm);
        let p = HelloPacket {
            magic: String::from_utf8(HELLO_MAGIC.to_vec()).unwrap(),
            id: "peer-1".into(),
            name: "node-A".into(),
            ram_gb: 32,
            pledged: 1,
            compression: 0,
        };
        let src: SocketAddr = "127.0.0.1:50050".parse().unwrap();
        let node = d.ingest_packet(p, src, unix_ms_now()).unwrap();
        assert_eq!(node.id, "peer-1");
        assert_eq!(d.peer_count(), 1);

        // After 7 seconds (exceeds 6s expiry), should evict
        let stale_time = unix_ms_now() + 7000;
        let evicted = d.evict_stale(stale_time);
        assert_eq!(evicted, vec!["peer-1".to_string()]);
        assert_eq!(d.peer_count(), 0);
    }

    #[test]
    fn ignores_packets_with_wrong_magic() {
        let d = ClusterDiscovery::new(64, 1, CompressionMode::BinaryGemm);
        let p = HelloPacket {
            magic: "INVALID".into(),
            id: "peer-2".into(),
            name: "impostor".into(),
            ram_gb: 32,
            pledged: 1,
            compression: 0,
        };
        let src: SocketAddr = "127.0.0.1:50050".parse().unwrap();
        let none = d.ingest_packet(p, src, unix_ms_now());
        assert!(none.is_none());
    }
}