diff --git a/crates/ember-cluster/src/gossip.rs b/crates/ember-cluster/src/gossip.rs new file mode 100644 index 00000000..029acefd --- /dev/null +++ b/crates/ember-cluster/src/gossip.rs @@ -0,0 +1,670 @@ +//! SWIM gossip protocol implementation. +//! +//! Implements the Scalable Weakly-consistent Infection-style Membership +//! protocol for failure detection and cluster membership management. +//! +//! # Protocol Overview +//! +//! Each protocol period: +//! 1. Pick a random node to probe with PING +//! 2. If no ACK within timeout, send PING-REQ to k random nodes +//! 3. If still no ACK, mark node as SUSPECT +//! 4. After suspicion timeout, mark as DEAD +//! 5. Piggyback state updates on all messages + +use std::collections::HashMap; +use std::net::SocketAddr; +use std::time::{Duration, Instant}; + +use rand::prelude::IndexedRandom; +use tokio::sync::mpsc; +use tracing::{debug, info, trace, warn}; + +use crate::message::{GossipMessage, MemberInfo, NodeUpdate}; +use crate::{NodeId, SlotRange}; + +/// Configuration for the gossip protocol. +#[derive(Debug, Clone)] +pub struct GossipConfig { + /// How often to run the protocol period (probe a random node). + pub protocol_period: Duration, + /// How long to wait for a direct probe response. + pub probe_timeout: Duration, + /// Multiplier for suspicion timeout (protocol_period * suspicion_mult). + pub suspicion_mult: u32, + /// Number of nodes to ask for indirect probes. + pub indirect_probes: usize, + /// Maximum number of updates to piggyback per message. + pub max_piggyback: usize, + /// Port offset for gossip (data_port + gossip_port_offset). + pub gossip_port_offset: u16, +} + +impl Default for GossipConfig { + fn default() -> Self { + Self { + protocol_period: Duration::from_secs(1), + probe_timeout: Duration::from_millis(500), + suspicion_mult: 5, + indirect_probes: 3, + max_piggyback: 10, + gossip_port_offset: 10000, + } + } +} + +/// Internal state of a cluster member as tracked by gossip. +#[derive(Debug, Clone)] +pub struct MemberState { + pub id: NodeId, + pub addr: SocketAddr, + pub incarnation: u64, + pub state: MemberStatus, + pub state_change: Instant, + pub is_primary: bool, + pub slots: Vec, +} + +/// Health status of a member. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum MemberStatus { + Alive, + Suspect, + Dead, + Left, +} + +/// Events emitted by the gossip engine. +#[derive(Debug, Clone)] +pub enum GossipEvent { + /// A new node joined the cluster. + MemberJoined(NodeId, SocketAddr), + /// A node is suspected to be failing. + MemberSuspected(NodeId), + /// A node has been confirmed dead. + MemberFailed(NodeId), + /// A node left gracefully. + MemberLeft(NodeId), + /// A node that was suspected is now alive. + MemberAlive(NodeId), +} + +/// The gossip engine manages cluster membership and failure detection. +pub struct GossipEngine { + /// Our node's identity. + local_id: NodeId, + /// Our advertised address. + local_addr: SocketAddr, + /// Our incarnation number (incremented to refute suspicion). + incarnation: u64, + /// Protocol configuration. + config: GossipConfig, + /// Known cluster members. + members: HashMap, + /// Pending updates to piggyback on outgoing messages. + pending_updates: Vec, + /// Sequence number for protocol messages. + next_seq: u64, + /// Pending probes awaiting acknowledgment. + pending_probes: HashMap, + /// Channel for emitting events. + event_tx: mpsc::Sender, +} + +struct PendingProbe { + target: NodeId, + sent_at: Instant, + indirect: bool, +} + +impl GossipEngine { + /// Creates a new gossip engine. + pub fn new( + local_id: NodeId, + local_addr: SocketAddr, + config: GossipConfig, + event_tx: mpsc::Sender, + ) -> Self { + Self { + local_id, + local_addr, + incarnation: 1, + config, + members: HashMap::new(), + pending_updates: Vec::new(), + next_seq: 1, + pending_probes: HashMap::new(), + event_tx, + } + } + + /// Returns the local node ID. + pub fn local_id(&self) -> NodeId { + self.local_id + } + + /// Returns all known members. + pub fn members(&self) -> impl Iterator { + self.members.values() + } + + /// Returns the number of alive members (excluding self). + pub fn alive_count(&self) -> usize { + self.members + .values() + .filter(|m| m.state == MemberStatus::Alive) + .count() + } + + /// Adds a seed node to bootstrap cluster discovery. + pub fn add_seed(&mut self, id: NodeId, addr: SocketAddr) { + if id == self.local_id { + return; + } + self.members.entry(id).or_insert_with(|| MemberState { + id, + addr, + incarnation: 0, + state: MemberStatus::Alive, + state_change: Instant::now(), + is_primary: false, + slots: Vec::new(), + }); + } + + /// Handles an incoming gossip message. + pub async fn handle_message( + &mut self, + msg: GossipMessage, + from: SocketAddr, + ) -> Option { + match msg { + GossipMessage::Ping { + seq, + sender, + updates, + } => { + trace!("received ping seq={} from {}", seq, sender); + self.apply_updates(&updates).await; + self.ensure_member(sender, from); + + // Reply with ACK + let response_updates = self.collect_updates(); + Some(GossipMessage::Ack { + seq, + sender: self.local_id, + updates: response_updates, + }) + } + + GossipMessage::PingReq { + seq, + sender, + target, + target_addr: _, + } => { + trace!( + "received ping-req seq={} from {} for {}", + seq, + sender, + target + ); + self.ensure_member(sender, from); + + // Forward ping to target (handled externally) + // For now, we just record that we might need to relay + None + } + + GossipMessage::Ack { + seq, + sender, + updates, + } => { + trace!("received ack seq={} from {}", seq, sender); + self.apply_updates(&updates).await; + self.ensure_member(sender, from); + + // Clear pending probe + if let Some(probe) = self.pending_probes.remove(&seq) { + if self.members.get(&probe.target).map(|m| m.state) + == Some(MemberStatus::Suspect) + { + // Node recovered from suspicion + self.mark_alive(probe.target).await; + } + } + None + } + + GossipMessage::Join { + sender, + sender_addr, + } => { + info!("node {} joining from {}", sender, sender_addr); + self.ensure_member(sender, sender_addr); + + // Broadcast alive update + self.queue_update(NodeUpdate::Alive { + node: sender, + addr: sender_addr, + incarnation: 1, + }); + + // Send welcome with current members + let members: Vec = self + .members + .values() + .filter(|m| m.state == MemberStatus::Alive) + .map(|m| MemberInfo { + id: m.id, + addr: m.addr, + incarnation: m.incarnation, + is_primary: m.is_primary, + slots: m.slots.clone(), + }) + .collect(); + + Some(GossipMessage::Welcome { + sender: self.local_id, + members, + }) + } + + GossipMessage::Welcome { sender, members } => { + info!( + "received welcome from {} with {} members", + sender, + members.len() + ); + self.ensure_member(sender, from); + + for member in members { + if member.id != self.local_id { + self.members + .entry(member.id) + .or_insert_with(|| MemberState { + id: member.id, + addr: member.addr, + incarnation: member.incarnation, + state: MemberStatus::Alive, + state_change: Instant::now(), + is_primary: member.is_primary, + slots: member.slots, + }); + } + } + None + } + } + } + + /// Runs one protocol period: probe a random node. + pub fn tick(&mut self) -> Option<(SocketAddr, GossipMessage)> { + // Check for timed-out probes + self.check_probe_timeouts(); + + // Check for expired suspicions + self.check_suspicion_timeouts(); + + // Select a random alive member to probe + let target_info = { + let alive_members: Vec<_> = self + .members + .values() + .filter(|m| m.state == MemberStatus::Alive || m.state == MemberStatus::Suspect) + .map(|m| (m.id, m.addr)) + .collect(); + + if alive_members.is_empty() { + return None; + } + + *alive_members.choose(&mut rand::rng())? + }; + + let (target_id, target_addr) = target_info; + let seq = self.next_seq; + self.next_seq += 1; + + let updates = self.collect_updates(); + let msg = GossipMessage::Ping { + seq, + sender: self.local_id, + updates, + }; + + self.pending_probes.insert( + seq, + PendingProbe { + target: target_id, + sent_at: Instant::now(), + indirect: false, + }, + ); + + Some((target_addr, msg)) + } + + /// Creates a join message to send to a seed node. + pub fn create_join_message(&self) -> GossipMessage { + GossipMessage::Join { + sender: self.local_id, + sender_addr: self.local_addr, + } + } + + fn ensure_member(&mut self, id: NodeId, addr: SocketAddr) { + if id == self.local_id { + return; + } + self.members.entry(id).or_insert_with(|| MemberState { + id, + addr, + incarnation: 0, + state: MemberStatus::Alive, + state_change: Instant::now(), + is_primary: false, + slots: Vec::new(), + }); + } + + async fn apply_updates(&mut self, updates: &[NodeUpdate]) { + for update in updates { + match update { + NodeUpdate::Alive { + node, + addr, + incarnation, + } => { + if *node == self.local_id { + // Someone thinks we're alive, good + continue; + } + if let Some(member) = self.members.get_mut(node) { + if *incarnation > member.incarnation { + member.incarnation = *incarnation; + member.addr = *addr; + if member.state != MemberStatus::Alive { + member.state = MemberStatus::Alive; + member.state_change = Instant::now(); + let _ = self.event_tx.send(GossipEvent::MemberAlive(*node)).await; + } + } + } else { + self.members.insert( + *node, + MemberState { + id: *node, + addr: *addr, + incarnation: *incarnation, + state: MemberStatus::Alive, + state_change: Instant::now(), + is_primary: false, + slots: Vec::new(), + }, + ); + let _ = self + .event_tx + .send(GossipEvent::MemberJoined(*node, *addr)) + .await; + } + } + + NodeUpdate::Suspect { node, incarnation } => { + if *node == self.local_id { + // Refute suspicion by incrementing our incarnation + if *incarnation >= self.incarnation { + self.incarnation = incarnation + 1; + self.queue_update(NodeUpdate::Alive { + node: self.local_id, + addr: self.local_addr, + incarnation: self.incarnation, + }); + } + continue; + } + if let Some(member) = self.members.get_mut(node) { + if *incarnation >= member.incarnation && member.state == MemberStatus::Alive + { + member.state = MemberStatus::Suspect; + member.state_change = Instant::now(); + let _ = self + .event_tx + .send(GossipEvent::MemberSuspected(*node)) + .await; + } + } + } + + NodeUpdate::Dead { node, incarnation } => { + if *node == self.local_id { + // Refute death claim + self.incarnation = incarnation + 1; + self.queue_update(NodeUpdate::Alive { + node: self.local_id, + addr: self.local_addr, + incarnation: self.incarnation, + }); + continue; + } + if let Some(member) = self.members.get_mut(node) { + if *incarnation >= member.incarnation && member.state != MemberStatus::Dead + { + member.state = MemberStatus::Dead; + member.state_change = Instant::now(); + let _ = self.event_tx.send(GossipEvent::MemberFailed(*node)).await; + } + } + } + + NodeUpdate::Left { node } => { + if *node == self.local_id { + continue; + } + if let Some(member) = self.members.get_mut(node) { + if member.state != MemberStatus::Left { + member.state = MemberStatus::Left; + member.state_change = Instant::now(); + let _ = self.event_tx.send(GossipEvent::MemberLeft(*node)).await; + } + } + } + } + } + } + + async fn mark_alive(&mut self, node: NodeId) { + if let Some(member) = self.members.get_mut(&node) { + if member.state == MemberStatus::Suspect { + member.state = MemberStatus::Alive; + member.state_change = Instant::now(); + let _ = self.event_tx.send(GossipEvent::MemberAlive(node)).await; + } + } + } + + fn check_probe_timeouts(&mut self) { + let timeout = self.config.probe_timeout; + let now = Instant::now(); + + // Collect timed out probes first + let timed_out: Vec<_> = self + .pending_probes + .iter() + .filter(|(_, probe)| now.duration_since(probe.sent_at) > timeout && !probe.indirect) + .map(|(seq, probe)| (*seq, probe.target)) + .collect(); + + // Now process them + for (seq, target) in timed_out { + self.pending_probes.remove(&seq); + + // Get incarnation before mutating + let incarnation = self + .members + .get(&target) + .filter(|m| m.state == MemberStatus::Alive) + .map(|m| m.incarnation); + + if let Some(inc) = incarnation { + if let Some(member) = self.members.get_mut(&target) { + debug!("node {} failed to respond, marking suspect", target); + member.state = MemberStatus::Suspect; + member.state_change = Instant::now(); + } + self.queue_update(NodeUpdate::Suspect { + node: target, + incarnation: inc, + }); + } + } + } + + fn check_suspicion_timeouts(&mut self) { + let suspicion_timeout = self.config.protocol_period * self.config.suspicion_mult; + let now = Instant::now(); + let mut to_mark_dead = Vec::new(); + + for member in self.members.values() { + if member.state == MemberStatus::Suspect + && now.duration_since(member.state_change) > suspicion_timeout + { + to_mark_dead.push((member.id, member.incarnation)); + } + } + + for (id, incarnation) in to_mark_dead { + if let Some(member) = self.members.get_mut(&id) { + warn!("node {} confirmed dead after suspicion timeout", id); + member.state = MemberStatus::Dead; + member.state_change = Instant::now(); + self.queue_update(NodeUpdate::Dead { + node: id, + incarnation, + }); + } + } + } + + fn queue_update(&mut self, update: NodeUpdate) { + self.pending_updates.push(update); + // Keep bounded + if self.pending_updates.len() > self.config.max_piggyback * 2 { + self.pending_updates.drain(0..self.config.max_piggyback); + } + } + + fn collect_updates(&mut self) -> Vec { + let count = self.pending_updates.len().min(self.config.max_piggyback); + self.pending_updates.drain(0..count).collect() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::net::Ipv4Addr; + + fn test_addr(port: u16) -> SocketAddr { + SocketAddr::from((Ipv4Addr::new(127, 0, 0, 1), port)) + } + + #[tokio::test] + async fn engine_creation() { + let (tx, _rx) = mpsc::channel(16); + let engine = GossipEngine::new(NodeId::new(), test_addr(6379), GossipConfig::default(), tx); + assert_eq!(engine.alive_count(), 0); + } + + #[tokio::test] + async fn add_seed() { + let (tx, _rx) = mpsc::channel(16); + let mut engine = + GossipEngine::new(NodeId::new(), test_addr(6379), GossipConfig::default(), tx); + + let seed_id = NodeId::new(); + engine.add_seed(seed_id, test_addr(6380)); + assert_eq!(engine.alive_count(), 1); + } + + #[tokio::test] + async fn handle_ping() { + let (tx, _rx) = mpsc::channel(16); + let mut engine = + GossipEngine::new(NodeId::new(), test_addr(6379), GossipConfig::default(), tx); + + let sender = NodeId::new(); + let msg = GossipMessage::Ping { + seq: 1, + sender, + updates: vec![], + }; + + let response = engine.handle_message(msg, test_addr(6380)).await; + assert!(matches!(response, Some(GossipMessage::Ack { .. }))); + assert_eq!(engine.alive_count(), 1); + } + + #[tokio::test] + async fn handle_join() { + let (tx, _rx) = mpsc::channel(16); + let mut engine = + GossipEngine::new(NodeId::new(), test_addr(6379), GossipConfig::default(), tx); + + let joiner = NodeId::new(); + let msg = GossipMessage::Join { + sender: joiner, + sender_addr: test_addr(6380), + }; + + let response = engine.handle_message(msg, test_addr(6380)).await; + assert!(matches!(response, Some(GossipMessage::Welcome { .. }))); + assert_eq!(engine.alive_count(), 1); + } + + #[tokio::test] + async fn tick_with_no_members() { + let (tx, _rx) = mpsc::channel(16); + let mut engine = + GossipEngine::new(NodeId::new(), test_addr(6379), GossipConfig::default(), tx); + + let probe = engine.tick(); + assert!(probe.is_none()); + } + + #[tokio::test] + async fn tick_with_members() { + let (tx, _rx) = mpsc::channel(16); + let mut engine = + GossipEngine::new(NodeId::new(), test_addr(6379), GossipConfig::default(), tx); + + engine.add_seed(NodeId::new(), test_addr(6380)); + let probe = engine.tick(); + assert!(probe.is_some()); + + let (addr, msg) = probe.unwrap(); + assert_eq!(addr.port(), 6380); + assert!(matches!(msg, GossipMessage::Ping { .. })); + } + + #[tokio::test] + async fn create_join_message() { + let (tx, _rx) = mpsc::channel(16); + let id = NodeId::new(); + let addr = test_addr(6379); + let engine = GossipEngine::new(id, addr, GossipConfig::default(), tx); + + let msg = engine.create_join_message(); + match msg { + GossipMessage::Join { + sender, + sender_addr, + } => { + assert_eq!(sender, id); + assert_eq!(sender_addr, addr); + } + _ => panic!("expected Join message"), + } + } +} diff --git a/crates/ember-cluster/src/lib.rs b/crates/ember-cluster/src/lib.rs index 5d6dc2c9..f4c13ea9 100644 --- a/crates/ember-cluster/src/lib.rs +++ b/crates/ember-cluster/src/lib.rs @@ -30,9 +30,13 @@ //! ``` mod error; +mod gossip; +mod message; mod slots; mod topology; pub use error::ClusterError; +pub use gossip::{GossipConfig, GossipEngine, GossipEvent, MemberState, MemberStatus}; +pub use message::{GossipMessage, MemberInfo, NodeUpdate}; pub use slots::{key_slot, SlotMap, SlotRange, SLOT_COUNT}; pub use topology::{ClusterHealth, ClusterNode, ClusterState, NodeFlags, NodeId, NodeRole}; diff --git a/crates/ember-cluster/src/message.rs b/crates/ember-cluster/src/message.rs new file mode 100644 index 00000000..8e39bd33 --- /dev/null +++ b/crates/ember-cluster/src/message.rs @@ -0,0 +1,582 @@ +//! Binary wire format for cluster gossip messages. +//! +//! Uses a compact binary encoding for efficiency over the network. +//! All multi-byte integers are little-endian. + +use std::io::{self, Read}; +use std::net::SocketAddr; + +use bytes::{Buf, BufMut, Bytes, BytesMut}; + +use crate::{NodeId, SlotRange}; + +/// Message types for the SWIM gossip protocol. +#[derive(Debug, Clone, PartialEq)] +pub enum GossipMessage { + /// Direct probe to check if a node is alive. + Ping { + seq: u64, + sender: NodeId, + /// Piggybacked state updates. + updates: Vec, + }, + + /// Request another node to probe a target on our behalf. + PingReq { + seq: u64, + sender: NodeId, + target: NodeId, + target_addr: SocketAddr, + }, + + /// Response to a Ping or forwarded PingReq. + Ack { + seq: u64, + sender: NodeId, + /// Piggybacked state updates. + updates: Vec, + }, + + /// Join request from a new node. + Join { + sender: NodeId, + sender_addr: SocketAddr, + }, + + /// Welcome response with current cluster state. + Welcome { + sender: NodeId, + members: Vec, + }, +} + +/// A state update about a node, piggybacked on protocol messages. +#[derive(Debug, Clone, PartialEq)] +pub enum NodeUpdate { + /// Node is alive with given incarnation number. + Alive { + node: NodeId, + addr: SocketAddr, + incarnation: u64, + }, + /// Node is suspected to be failing. + Suspect { node: NodeId, incarnation: u64 }, + /// Node has been confirmed dead. + Dead { node: NodeId, incarnation: u64 }, + /// Node left the cluster gracefully. + Left { node: NodeId }, +} + +/// Information about a cluster member. +#[derive(Debug, Clone, PartialEq)] +pub struct MemberInfo { + pub id: NodeId, + pub addr: SocketAddr, + pub incarnation: u64, + pub is_primary: bool, + pub slots: Vec, +} + +// Wire format constants +const MSG_PING: u8 = 1; +const MSG_PING_REQ: u8 = 2; +const MSG_ACK: u8 = 3; +const MSG_JOIN: u8 = 4; +const MSG_WELCOME: u8 = 5; + +const UPDATE_ALIVE: u8 = 1; +const UPDATE_SUSPECT: u8 = 2; +const UPDATE_DEAD: u8 = 3; +const UPDATE_LEFT: u8 = 4; + +impl GossipMessage { + /// Serializes the message to bytes. + pub fn encode(&self) -> Bytes { + let mut buf = BytesMut::with_capacity(256); + self.encode_into(&mut buf); + buf.freeze() + } + + /// Serializes the message into the given buffer. + pub fn encode_into(&self, buf: &mut BytesMut) { + match self { + GossipMessage::Ping { + seq, + sender, + updates, + } => { + buf.put_u8(MSG_PING); + buf.put_u64_le(*seq); + encode_node_id(buf, sender); + encode_updates(buf, updates); + } + GossipMessage::PingReq { + seq, + sender, + target, + target_addr, + } => { + buf.put_u8(MSG_PING_REQ); + buf.put_u64_le(*seq); + encode_node_id(buf, sender); + encode_node_id(buf, target); + encode_socket_addr(buf, target_addr); + } + GossipMessage::Ack { + seq, + sender, + updates, + } => { + buf.put_u8(MSG_ACK); + buf.put_u64_le(*seq); + encode_node_id(buf, sender); + encode_updates(buf, updates); + } + GossipMessage::Join { + sender, + sender_addr, + } => { + buf.put_u8(MSG_JOIN); + encode_node_id(buf, sender); + encode_socket_addr(buf, sender_addr); + } + GossipMessage::Welcome { sender, members } => { + buf.put_u8(MSG_WELCOME); + encode_node_id(buf, sender); + buf.put_u16_le(members.len() as u16); + for member in members { + encode_member_info(buf, member); + } + } + } + } + + /// Deserializes a message from bytes. + pub fn decode(mut buf: &[u8]) -> io::Result { + if buf.is_empty() { + return Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "empty message", + )); + } + + let msg_type = buf.get_u8(); + match msg_type { + MSG_PING => { + let seq = buf.get_u64_le(); + let sender = decode_node_id(&mut buf)?; + let updates = decode_updates(&mut buf)?; + Ok(GossipMessage::Ping { + seq, + sender, + updates, + }) + } + MSG_PING_REQ => { + let seq = buf.get_u64_le(); + let sender = decode_node_id(&mut buf)?; + let target = decode_node_id(&mut buf)?; + let target_addr = decode_socket_addr(&mut buf)?; + Ok(GossipMessage::PingReq { + seq, + sender, + target, + target_addr, + }) + } + MSG_ACK => { + let seq = buf.get_u64_le(); + let sender = decode_node_id(&mut buf)?; + let updates = decode_updates(&mut buf)?; + Ok(GossipMessage::Ack { + seq, + sender, + updates, + }) + } + MSG_JOIN => { + let sender = decode_node_id(&mut buf)?; + let sender_addr = decode_socket_addr(&mut buf)?; + Ok(GossipMessage::Join { + sender, + sender_addr, + }) + } + MSG_WELCOME => { + let sender = decode_node_id(&mut buf)?; + let count = buf.get_u16_le() as usize; + let mut members = Vec::with_capacity(count); + for _ in 0..count { + members.push(decode_member_info(&mut buf)?); + } + Ok(GossipMessage::Welcome { sender, members }) + } + other => Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("unknown message type: {other}"), + )), + } + } +} + +fn encode_node_id(buf: &mut BytesMut, id: &NodeId) { + buf.put_slice(id.0.as_bytes()); +} + +fn decode_node_id(buf: &mut &[u8]) -> io::Result { + if buf.len() < 16 { + return Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "not enough bytes for node id", + )); + } + let mut bytes = [0u8; 16]; + buf.read_exact(&mut bytes)?; + Ok(NodeId(uuid::Uuid::from_bytes(bytes))) +} + +fn encode_socket_addr(buf: &mut BytesMut, addr: &SocketAddr) { + match addr { + SocketAddr::V4(v4) => { + buf.put_u8(4); + buf.put_slice(&v4.ip().octets()); + buf.put_u16_le(v4.port()); + } + SocketAddr::V6(v6) => { + buf.put_u8(6); + buf.put_slice(&v6.ip().octets()); + buf.put_u16_le(v6.port()); + } + } +} + +fn decode_socket_addr(buf: &mut &[u8]) -> io::Result { + if buf.is_empty() { + return Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "not enough bytes for address type", + )); + } + let addr_type = buf.get_u8(); + match addr_type { + 4 => { + if buf.len() < 6 { + return Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "not enough bytes for ipv4 address", + )); + } + let mut octets = [0u8; 4]; + buf.read_exact(&mut octets)?; + let port = buf.get_u16_le(); + Ok(SocketAddr::from((octets, port))) + } + 6 => { + if buf.len() < 18 { + return Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "not enough bytes for ipv6 address", + )); + } + let mut octets = [0u8; 16]; + buf.read_exact(&mut octets)?; + let port = buf.get_u16_le(); + Ok(SocketAddr::from((octets, port))) + } + other => Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("unknown address type: {other}"), + )), + } +} + +fn encode_updates(buf: &mut BytesMut, updates: &[NodeUpdate]) { + buf.put_u16_le(updates.len() as u16); + for update in updates { + encode_update(buf, update); + } +} + +fn encode_update(buf: &mut BytesMut, update: &NodeUpdate) { + match update { + NodeUpdate::Alive { + node, + addr, + incarnation, + } => { + buf.put_u8(UPDATE_ALIVE); + encode_node_id(buf, node); + encode_socket_addr(buf, addr); + buf.put_u64_le(*incarnation); + } + NodeUpdate::Suspect { node, incarnation } => { + buf.put_u8(UPDATE_SUSPECT); + encode_node_id(buf, node); + buf.put_u64_le(*incarnation); + } + NodeUpdate::Dead { node, incarnation } => { + buf.put_u8(UPDATE_DEAD); + encode_node_id(buf, node); + buf.put_u64_le(*incarnation); + } + NodeUpdate::Left { node } => { + buf.put_u8(UPDATE_LEFT); + encode_node_id(buf, node); + } + } +} + +fn decode_updates(buf: &mut &[u8]) -> io::Result> { + if buf.len() < 2 { + return Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "not enough bytes for update count", + )); + } + let count = buf.get_u16_le() as usize; + let mut updates = Vec::with_capacity(count); + for _ in 0..count { + updates.push(decode_update(buf)?); + } + Ok(updates) +} + +fn decode_update(buf: &mut &[u8]) -> io::Result { + if buf.is_empty() { + return Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "not enough bytes for update type", + )); + } + let update_type = buf.get_u8(); + match update_type { + UPDATE_ALIVE => { + let node = decode_node_id(buf)?; + let addr = decode_socket_addr(buf)?; + let incarnation = buf.get_u64_le(); + Ok(NodeUpdate::Alive { + node, + addr, + incarnation, + }) + } + UPDATE_SUSPECT => { + let node = decode_node_id(buf)?; + let incarnation = buf.get_u64_le(); + Ok(NodeUpdate::Suspect { node, incarnation }) + } + UPDATE_DEAD => { + let node = decode_node_id(buf)?; + let incarnation = buf.get_u64_le(); + Ok(NodeUpdate::Dead { node, incarnation }) + } + UPDATE_LEFT => { + let node = decode_node_id(buf)?; + Ok(NodeUpdate::Left { node }) + } + other => Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("unknown update type: {other}"), + )), + } +} + +fn encode_member_info(buf: &mut BytesMut, member: &MemberInfo) { + encode_node_id(buf, &member.id); + encode_socket_addr(buf, &member.addr); + buf.put_u64_le(member.incarnation); + buf.put_u8(if member.is_primary { 1 } else { 0 }); + buf.put_u16_le(member.slots.len() as u16); + for slot in &member.slots { + buf.put_u16_le(slot.start); + buf.put_u16_le(slot.end); + } +} + +fn decode_member_info(buf: &mut &[u8]) -> io::Result { + let id = decode_node_id(buf)?; + let addr = decode_socket_addr(buf)?; + let incarnation = buf.get_u64_le(); + let is_primary = buf.get_u8() != 0; + let slot_count = buf.get_u16_le() as usize; + let mut slots = Vec::with_capacity(slot_count); + for _ in 0..slot_count { + let start = buf.get_u16_le(); + let end = buf.get_u16_le(); + slots.push(SlotRange::new(start, end)); + } + Ok(MemberInfo { + id, + addr, + incarnation, + is_primary, + slots, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::net::{Ipv4Addr, Ipv6Addr}; + + fn test_addr() -> SocketAddr { + SocketAddr::from((Ipv4Addr::new(127, 0, 0, 1), 6379)) + } + + fn test_addr_v6() -> SocketAddr { + SocketAddr::from((Ipv6Addr::LOCALHOST, 6379)) + } + + #[test] + fn ping_roundtrip() { + let msg = GossipMessage::Ping { + seq: 42, + sender: NodeId::new(), + updates: vec![], + }; + let encoded = msg.encode(); + let decoded = GossipMessage::decode(&encoded).unwrap(); + assert_eq!(msg, decoded); + } + + #[test] + fn ping_with_updates() { + let node1 = NodeId::new(); + let node2 = NodeId::new(); + let msg = GossipMessage::Ping { + seq: 100, + sender: node1, + updates: vec![ + NodeUpdate::Alive { + node: node2, + addr: test_addr(), + incarnation: 5, + }, + NodeUpdate::Suspect { + node: node1, + incarnation: 3, + }, + ], + }; + let encoded = msg.encode(); + let decoded = GossipMessage::decode(&encoded).unwrap(); + assert_eq!(msg, decoded); + } + + #[test] + fn ping_req_roundtrip() { + let msg = GossipMessage::PingReq { + seq: 99, + sender: NodeId::new(), + target: NodeId::new(), + target_addr: test_addr(), + }; + let encoded = msg.encode(); + let decoded = GossipMessage::decode(&encoded).unwrap(); + assert_eq!(msg, decoded); + } + + #[test] + fn ack_roundtrip() { + let msg = GossipMessage::Ack { + seq: 42, + sender: NodeId::new(), + updates: vec![NodeUpdate::Dead { + node: NodeId::new(), + incarnation: 10, + }], + }; + let encoded = msg.encode(); + let decoded = GossipMessage::decode(&encoded).unwrap(); + assert_eq!(msg, decoded); + } + + #[test] + fn join_roundtrip() { + let msg = GossipMessage::Join { + sender: NodeId::new(), + sender_addr: test_addr(), + }; + let encoded = msg.encode(); + let decoded = GossipMessage::decode(&encoded).unwrap(); + assert_eq!(msg, decoded); + } + + #[test] + fn welcome_roundtrip() { + let msg = GossipMessage::Welcome { + sender: NodeId::new(), + members: vec![ + MemberInfo { + id: NodeId::new(), + addr: test_addr(), + incarnation: 1, + is_primary: true, + slots: vec![SlotRange::new(0, 5460)], + }, + MemberInfo { + id: NodeId::new(), + addr: test_addr(), + incarnation: 2, + is_primary: false, + slots: vec![], + }, + ], + }; + let encoded = msg.encode(); + let decoded = GossipMessage::decode(&encoded).unwrap(); + assert_eq!(msg, decoded); + } + + #[test] + fn ipv6_address() { + let msg = GossipMessage::Join { + sender: NodeId::new(), + sender_addr: test_addr_v6(), + }; + let encoded = msg.encode(); + let decoded = GossipMessage::decode(&encoded).unwrap(); + assert_eq!(msg, decoded); + } + + #[test] + fn all_update_types() { + let node = NodeId::new(); + let updates = vec![ + NodeUpdate::Alive { + node, + addr: test_addr(), + incarnation: 1, + }, + NodeUpdate::Suspect { + node, + incarnation: 2, + }, + NodeUpdate::Dead { + node, + incarnation: 3, + }, + NodeUpdate::Left { node }, + ]; + let msg = GossipMessage::Ping { + seq: 1, + sender: node, + updates, + }; + let encoded = msg.encode(); + let decoded = GossipMessage::decode(&encoded).unwrap(); + assert_eq!(msg, decoded); + } + + #[test] + fn empty_message_error() { + let result = GossipMessage::decode(&[]); + assert!(result.is_err()); + } + + #[test] + fn unknown_message_type_error() { + let result = GossipMessage::decode(&[255]); + assert!(result.is_err()); + } +}