use std::collections::BTreeSet;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
type NodeId = u64;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct NodeCapabilities {
pub min_readable_version: u8,
pub max_readable_version: u8,
pub active_write_version: u8,
}
impl NodeCapabilities {
pub fn local(active_write_version: u8) -> Self {
Self {
min_readable_version: tsoracle_openraft_toolkit::MIN_READABLE_VERSION,
max_readable_version: tsoracle_openraft_toolkit::MAX_READABLE_VERSION,
active_write_version,
}
}
}
pub fn all_members_can_read(
target: u8,
capabilities: &[(NodeId, NodeCapabilities)],
) -> Option<BTreeSet<NodeId>> {
if capabilities
.iter()
.all(|(_, member)| member.max_readable_version >= target)
{
Some(capabilities.iter().map(|(node_id, _)| *node_id).collect())
} else {
None
}
}
#[derive(Debug, thiserror::Error)]
pub enum FormatActivationError {
#[error("cannot initiate format activation: this node is not the leader")]
NotLeader,
#[error(
"format activation to target {target} blocked: target outside local readable range [{min}, {max}]"
)]
TargetOutOfRange { target: u8, min: u8, max: u8 },
#[error("format activation to target {target} blocked: members below target: {incapable:?}")]
MembersBelowTarget {
target: u8,
incapable: Vec<(NodeId, u8)>,
},
#[error("format activation gate failed: member {node_id} unreachable: {detail}")]
MemberUnreachable { node_id: NodeId, detail: String },
#[error(
"format activation to target {target} applied as a no-op: membership changed since the gate"
)]
MembershipChangedSinceGate { target: u8 },
#[error("format activation proposal failed: {0}")]
ProposalFailed(String),
}
pub fn target_in_local_readable_range(target: u8) -> Result<(), FormatActivationError> {
let min = tsoracle_openraft_toolkit::MIN_READABLE_VERSION;
let max = tsoracle_openraft_toolkit::MAX_READABLE_VERSION;
if !(min..=max).contains(&target) {
return Err(FormatActivationError::TargetOutOfRange { target, min, max });
}
Ok(())
}
#[async_trait]
pub trait CapabilitySource: Send + Sync {
type Node: Send + Sync;
async fn query(&self, node_id: NodeId, member: &Self::Node)
-> Result<NodeCapabilities, String>;
}
pub async fn gather_with<S: CapabilitySource>(
local_node: NodeId,
local_capabilities: NodeCapabilities,
membership: &[(NodeId, S::Node)],
source: &S,
) -> Result<Vec<(NodeId, NodeCapabilities)>, FormatActivationError> {
let mut gathered = Vec::with_capacity(membership.len());
for (node_id, member) in membership {
let capabilities = if *node_id == local_node {
local_capabilities
} else {
source.query(*node_id, member).await.map_err(|detail| {
FormatActivationError::MemberUnreachable {
node_id: *node_id,
detail,
}
})?
};
gathered.push((*node_id, capabilities));
}
Ok(gathered)
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
fn caps(active: u8, max: u8) -> NodeCapabilities {
NodeCapabilities {
min_readable_version: tsoracle_openraft_toolkit::MIN_READABLE_VERSION,
max_readable_version: max,
active_write_version: active,
}
}
#[test]
fn local_capabilities_reports_compile_time_read_range() {
let capabilities = NodeCapabilities::local(7);
assert_eq!(
capabilities.min_readable_version,
tsoracle_openraft_toolkit::MIN_READABLE_VERSION
);
assert_eq!(
capabilities.max_readable_version,
tsoracle_openraft_toolkit::MAX_READABLE_VERSION
);
assert_eq!(capabilities.active_write_version, 7);
}
#[test]
fn node_capabilities_postcard_round_trips() {
let original = NodeCapabilities {
min_readable_version: 3,
max_readable_version: 5,
active_write_version: 4,
};
let bytes = postcard::to_stdvec(&original).expect("encode");
let decoded: NodeCapabilities = postcard::from_bytes(&bytes).expect("decode");
assert_eq!(original, decoded);
}
#[test]
fn gate_passes_when_all_members_can_read_target() {
let reports = vec![(1u64, caps(3, 4)), (2, caps(3, 5)), (3, caps(3, 4))];
let gated = all_members_can_read(4, &reports).expect("all members can read 4");
assert_eq!(gated, BTreeSet::from([1, 2, 3]));
}
#[test]
fn gate_passes_at_exact_equality() {
let reports = vec![(1u64, caps(3, 4)), (2, caps(3, 4))];
assert_eq!(
all_members_can_read(4, &reports),
Some(BTreeSet::from([1, 2]))
);
}
#[test]
fn gate_fails_when_one_member_is_below_target() {
let reports = vec![(1u64, caps(3, 4)), (2, caps(3, 3)), (3, caps(3, 4))];
assert_eq!(all_members_can_read(4, &reports), None);
}
#[test]
fn existing_incompatible_learner_blocks_activation_via_all_members_gate() {
let target = 4u8;
let capable_voter = caps(3, 4);
let incompatible_learner = caps(3, 3); let reports = vec![(1u64, capable_voter), (2u64, incompatible_learner)];
assert_eq!(
all_members_can_read(target, &reports),
None,
"an incompatible already-admitted member (here a learner) must \
fail the all-members gate, blocking activation until remediated"
);
}
#[test]
fn gate_on_empty_membership_is_vacuously_satisfied_with_empty_set() {
assert_eq!(all_members_can_read(4, &[]), Some(BTreeSet::new()));
}
#[test]
fn format_activation_error_below_target_names_members_and_target() {
let err = FormatActivationError::MembersBelowTarget {
target: 4,
incapable: vec![(2, 3), (5, 3)],
};
let rendered = err.to_string();
assert!(rendered.contains("target 4"), "got: {rendered}");
assert!(
rendered.contains('2') && rendered.contains('5'),
"got: {rendered}"
);
}
#[test]
fn target_in_local_readable_range_accepts_min() {
assert!(
target_in_local_readable_range(tsoracle_openraft_toolkit::MIN_READABLE_VERSION).is_ok()
);
}
#[test]
fn target_in_local_readable_range_accepts_max() {
assert!(
target_in_local_readable_range(tsoracle_openraft_toolkit::MAX_READABLE_VERSION).is_ok()
);
}
#[test]
fn target_in_local_readable_range_rejects_zero() {
let err = target_in_local_readable_range(0).expect_err("0 is below MIN");
match err {
FormatActivationError::TargetOutOfRange { target, min, max } => {
assert_eq!(target, 0);
assert_eq!(min, tsoracle_openraft_toolkit::MIN_READABLE_VERSION);
assert_eq!(max, tsoracle_openraft_toolkit::MAX_READABLE_VERSION);
}
other => panic!("expected TargetOutOfRange, got: {other:?}"),
}
}
#[test]
fn target_in_local_readable_range_rejects_below_min() {
let just_below = tsoracle_openraft_toolkit::MIN_READABLE_VERSION - 1;
let err = target_in_local_readable_range(just_below).expect_err("MIN-1 is out of range");
assert!(matches!(
err,
FormatActivationError::TargetOutOfRange { target, .. } if target == just_below
));
}
#[test]
fn target_in_local_readable_range_rejects_above_max() {
let just_above = tsoracle_openraft_toolkit::MAX_READABLE_VERSION.saturating_add(1);
if just_above == tsoracle_openraft_toolkit::MAX_READABLE_VERSION {
return;
}
let err = target_in_local_readable_range(just_above).expect_err("MAX+1 is out of range");
assert!(matches!(
err,
FormatActivationError::TargetOutOfRange { target, .. } if target == just_above
));
}
#[test]
fn target_out_of_range_error_names_target_and_local_range() {
let err = FormatActivationError::TargetOutOfRange {
target: 1,
min: 4,
max: 4,
};
let rendered = err.to_string();
assert!(rendered.contains('1'), "target missing: {rendered}");
assert!(rendered.contains('4'), "range missing: {rendered}");
}
struct FakeSource {
responses: HashMap<NodeId, Result<NodeCapabilities, String>>,
}
#[async_trait]
impl CapabilitySource for FakeSource {
type Node = ();
async fn query(&self, node_id: NodeId, _member: &()) -> Result<NodeCapabilities, String> {
self.responses
.get(&node_id)
.cloned()
.unwrap_or_else(|| Err(format!("no fake response for {node_id}")))
}
}
#[tokio::test]
async fn gather_with_collects_remote_and_local() {
let source = FakeSource {
responses: HashMap::from([(2u64, Ok(caps(3, 5)))]),
};
let membership: Vec<(NodeId, ())> = vec![(1, ()), (2, ())];
let gathered = gather_with(1, caps(3, 4), &membership, &source)
.await
.expect("gather succeeds");
let by_id: HashMap<NodeId, NodeCapabilities> = gathered.into_iter().collect();
assert_eq!(by_id[&1].active_write_version, 3);
assert_eq!(by_id[&2].max_readable_version, 5);
}
#[tokio::test]
async fn gather_with_surfaces_unreachable_member() {
let source = FakeSource {
responses: HashMap::from([(2u64, Err("connection refused".to_string()))]),
};
let membership: Vec<(NodeId, ())> = vec![(1, ()), (2, ())];
let err = gather_with(1, caps(3, 4), &membership, &source)
.await
.expect_err("an unreachable member fails the gather closed");
assert!(matches!(
err,
FormatActivationError::MemberUnreachable { node_id: 2, .. }
));
}
}