Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 2 additions & 6 deletions crates/ember-cluster/src/election.rs
Original file line number Diff line number Diff line change
Expand Up @@ -54,12 +54,8 @@ impl Election {
return false;
}
self.votes.insert(from);
if self.votes.len() >= Self::quorum(total_primaries) {
self.promoted = true;
true
} else {
false
}
self.promoted = self.votes.len() >= Self::quorum(total_primaries);
self.promoted
}

/// Minimum votes required for a majority.
Expand Down
46 changes: 23 additions & 23 deletions crates/ember-cluster/src/gossip.rs
Original file line number Diff line number Diff line change
Expand Up @@ -478,10 +478,10 @@ impl GossipEngine {
if member.id == self.local_id {
continue;
}
let slots = member.slots.clone();
if let std::collections::hash_map::Entry::Vacant(e) =
self.members.entry(member.id)
{
let slots = member.slots.clone();
e.insert(MemberState {
id: member.id,
addr: member.addr,
Expand Down Expand Up @@ -759,9 +759,11 @@ impl GossipEngine {
if let Some(member) = self.members.get_mut(node) {
// only accept if incarnation is at least as recent
if *incarnation >= member.incarnation {
member.slots = slots.clone();
self.emit(GossipEvent::SlotsChanged(*node, slots.clone()))
.await;
let owned = slots.clone();
member.slots = owned.clone();
// release the mutable borrow before the async call
drop(member);
self.emit(GossipEvent::SlotsChanged(*node, owned)).await;
}
}
}
Expand Down Expand Up @@ -849,17 +851,23 @@ impl GossipEngine {
let now = Instant::now();
let mut outgoing = Vec::new();

// Phase 2: indirect probe timeouts → mark Suspect
let indirect_timed_out: Vec<_> = self
.pending_probes
.iter()
.filter(|(_, probe)| probe.indirect && now.duration_since(probe.sent_at) > timeout)
.map(|(seq, probe)| (*seq, probe.target))
.collect();

for (seq, target) in indirect_timed_out {
self.pending_probes.remove(&seq);
// Single pass: collect timed-out entries by kind and remove them.
let mut timed_out_indirect: Vec<(u64, NodeId)> = Vec::new();
let mut timed_out_direct: Vec<(u64, NodeId)> = Vec::new();
self.pending_probes.retain(|seq, probe| {
if now.duration_since(probe.sent_at) <= timeout {
return true;
}
if probe.indirect {
timed_out_indirect.push((*seq, probe.target));
} else {
timed_out_direct.push((*seq, probe.target));
}
false
});

// Phase 2: indirect probe timeouts → mark Suspect
for (_seq, target) in timed_out_indirect {
let incarnation = self
.members
.get(&target)
Expand All @@ -880,15 +888,7 @@ impl GossipEngine {
}

// Phase 1: direct ping timeouts → send PingReq
let direct_timed_out: Vec<_> = self
.pending_probes
.iter()
.filter(|(_, probe)| !probe.indirect && now.duration_since(probe.sent_at) > timeout)
.map(|(seq, probe)| (*seq, probe.target))
.collect();

for (seq, target) in direct_timed_out {
self.pending_probes.remove(&seq);
for (_seq, target) in timed_out_direct {

let target_addr = match self.members.get(&target) {
Some(m) if m.state == MemberStatus::Alive => m.addr,
Expand Down
2 changes: 1 addition & 1 deletion crates/ember-cluster/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
//!
//! # Quick Start
//!
//! ```rust,ignore
//! ```rust,no_run
//! use ember_cluster::{ClusterState, ClusterNode, NodeId, key_slot};
//!
//! // Create a single-node cluster
Expand Down
101 changes: 55 additions & 46 deletions crates/ember-cluster/src/message.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,11 @@ use crate::{NodeId, SlotRange};
/// Prevents allocation bombs from crafted messages.
const MAX_COLLECTION_COUNT: usize = 1024;

/// Address family discriminant for IPv4 (standard AF_INET byte count: 4).
const ADDR_IPV4: u8 = 4;
/// Address family discriminant for IPv6 (standard AF_INET6 byte count: 6 hex groups, 16 bytes).
const ADDR_IPV6: u8 = 6;

// Safe read helpers that return io::Error instead of panicking on truncated input.

fn safe_get_u8(buf: &mut &[u8]) -> io::Result<u8> {
Expand Down Expand Up @@ -163,6 +168,22 @@ const UPDATE_ROLE_CHANGED: u8 = 6;
const UPDATE_VOTE_REQUEST: u8 = 7;
const UPDATE_VOTE_GRANTED: u8 = 8;

/// Writes a count-capped slice into `buf` using a per-item encoder.
///
/// Writes `min(items.len(), MAX_COLLECTION_COUNT)` as a u16, then calls
/// `encode_one` for each item in the truncated slice. Used everywhere a
/// repeated list is encoded to avoid duplicating the truncation logic.
fn encode_slice<T, F>(buf: &mut BytesMut, items: &[T], encode_one: F)
where
F: Fn(&mut BytesMut, &T),
{
let count = items.len().min(MAX_COLLECTION_COUNT);
buf.put_u16_le(count as u16);
for item in &items[..count] {
encode_one(buf, item);
}
}

impl GossipMessage {
/// Serializes the message to bytes.
pub fn encode(&self) -> Bytes {
Expand Down Expand Up @@ -217,11 +238,7 @@ impl GossipMessage {
GossipMessage::Welcome { sender, members } => {
buf.put_u8(MSG_WELCOME);
encode_node_id(buf, sender);
let count = members.len().min(MAX_COLLECTION_COUNT);
buf.put_u16_le(count as u16);
for member in &members[..count] {
encode_member_info(buf, member);
}
encode_slice(buf, members, encode_member_info);
}
GossipMessage::SlotsAnnounce {
sender,
Expand All @@ -231,12 +248,10 @@ impl GossipMessage {
buf.put_u8(MSG_SLOTS_ANNOUNCE);
encode_node_id(buf, sender);
buf.put_u64_le(*incarnation);
let count = slots.len().min(MAX_COLLECTION_COUNT);
buf.put_u16_le(count as u16);
for slot in &slots[..count] {
buf.put_u16_le(slot.start);
buf.put_u16_le(slot.end);
}
encode_slice(buf, slots, |b, slot| {
b.put_u16_le(slot.start);
b.put_u16_le(slot.end);
});
}
}
}
Expand Down Expand Up @@ -356,12 +371,12 @@ fn decode_node_id(buf: &mut &[u8]) -> io::Result<NodeId> {
fn encode_socket_addr(buf: &mut BytesMut, addr: &SocketAddr) {
match addr {
SocketAddr::V4(v4) => {
buf.put_u8(4);
buf.put_u8(ADDR_IPV4);
buf.put_slice(&v4.ip().octets());
buf.put_u16_le(v4.port());
}
SocketAddr::V6(v6) => {
buf.put_u8(6);
buf.put_u8(ADDR_IPV6);
buf.put_slice(&v6.ip().octets());
buf.put_u16_le(v6.port());
}
Expand All @@ -377,7 +392,8 @@ fn decode_socket_addr(buf: &mut &[u8]) -> io::Result<SocketAddr> {
}
let addr_type = buf.get_u8();
match addr_type {
4 => {
ADDR_IPV4 => {
// 4 octets + 2-byte port = 6 bytes
if buf.len() < 6 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
Expand All @@ -389,7 +405,8 @@ fn decode_socket_addr(buf: &mut &[u8]) -> io::Result<SocketAddr> {
let port = buf.get_u16_le();
Ok(SocketAddr::from((octets, port)))
}
6 => {
ADDR_IPV6 => {
// 16 octets + 2-byte port = 18 bytes
if buf.len() < 18 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
Expand All @@ -409,11 +426,7 @@ fn decode_socket_addr(buf: &mut &[u8]) -> io::Result<SocketAddr> {
}

fn encode_updates(buf: &mut BytesMut, updates: &[NodeUpdate]) {
let count = updates.len().min(MAX_COLLECTION_COUNT);
buf.put_u16_le(count as u16);
for update in &updates[..count] {
encode_update(buf, update);
}
encode_slice(buf, updates, encode_update);
}

fn encode_update(buf: &mut BytesMut, update: &NodeUpdate) {
Expand Down Expand Up @@ -450,12 +463,10 @@ fn encode_update(buf: &mut BytesMut, update: &NodeUpdate) {
buf.put_u8(UPDATE_SLOTS_CHANGED);
encode_node_id(buf, node);
buf.put_u64_le(*incarnation);
let count = slots.len().min(MAX_COLLECTION_COUNT);
buf.put_u16_le(count as u16);
for slot in &slots[..count] {
buf.put_u16_le(slot.start);
buf.put_u16_le(slot.end);
}
encode_slice(buf, slots, |b, slot| {
b.put_u16_le(slot.start);
b.put_u16_le(slot.end);
});
}
NodeUpdate::RoleChanged {
node,
Expand Down Expand Up @@ -613,12 +624,10 @@ fn encode_member_info(buf: &mut BytesMut, member: &MemberInfo) {
encode_socket_addr(buf, &member.addr);
buf.put_u64_le(member.incarnation);
buf.put_u8(if member.is_primary { 1 } else { 0 });
let slot_count = member.slots.len().min(MAX_COLLECTION_COUNT);
buf.put_u16_le(slot_count as u16);
for slot in &member.slots[..slot_count] {
buf.put_u16_le(slot.start);
buf.put_u16_le(slot.end);
}
encode_slice(buf, &member.slots, |b, slot| {
b.put_u16_le(slot.start);
b.put_u16_le(slot.end);
});
}

fn decode_member_info(buf: &mut &[u8]) -> io::Result<MemberInfo> {
Expand Down Expand Up @@ -669,7 +678,7 @@ mod tests {
updates: vec![],
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand All @@ -693,7 +702,7 @@ mod tests {
],
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand All @@ -706,7 +715,7 @@ mod tests {
target_addr: test_addr(),
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand All @@ -721,7 +730,7 @@ mod tests {
}],
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand All @@ -732,7 +741,7 @@ mod tests {
sender_addr: test_addr(),
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand All @@ -758,7 +767,7 @@ mod tests {
],
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand All @@ -769,7 +778,7 @@ mod tests {
sender_addr: test_addr_v6(),
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand All @@ -790,7 +799,7 @@ mod tests {
}],
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);

// primary variant (no replicates field)
Expand All @@ -805,7 +814,7 @@ mod tests {
}],
};
let encoded2 = msg2.encode();
let decoded2 = GossipMessage::decode(&encoded2).unwrap();
let decoded2 = GossipMessage::decode(&encoded2).expect("roundtrip should succeed");
assert_eq!(msg2, decoded2);
}

Expand Down Expand Up @@ -839,7 +848,7 @@ mod tests {
updates,
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand All @@ -856,7 +865,7 @@ mod tests {
}],
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand All @@ -877,7 +886,7 @@ mod tests {
}],
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand Down Expand Up @@ -925,7 +934,7 @@ mod tests {
}],
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand All @@ -943,7 +952,7 @@ mod tests {
}],
};
let encoded = msg.encode();
let decoded = GossipMessage::decode(&encoded).unwrap();
let decoded = GossipMessage::decode(&encoded).expect("roundtrip should succeed");
assert_eq!(msg, decoded);
}

Expand Down
Loading
Loading