use std::collections::BTreeSet;
use async_trait::async_trait;
use openraft::Raft;
use openraft::RaftMetrics;
use openraft::ReadPolicy;
use openraft::ServerState;
use openraft::async_runtime::watch::WatchReceiver;
use openraft::error::{ClientWriteError, LinearizableReadError, RaftError};
use openraft::type_config::alias::WatchReceiverOf;
use tsoracle_consensus::ConsensusError;
use crate::capabilities::{
CapabilitySource, FormatActivationError, NodeCapabilities, all_members_can_read, gather_with,
target_in_local_readable_range,
};
use crate::host::OpenraftHighWaterHost;
use crate::log_entry::{HighWaterCommand, SetFormatVersionPayload};
use crate::state_machine::HighWaterStateMachine;
use crate::type_config::{ApplyOutcome, OpenraftPeer, TypeConfig};
use tsoracle_consensus::AdvancePayload;
pub struct StandaloneHost {
raft: Raft<TypeConfig, HighWaterStateMachine>,
state_machine: HighWaterStateMachine,
}
impl StandaloneHost {
pub fn new(
raft: Raft<TypeConfig, HighWaterStateMachine>,
state_machine: HighWaterStateMachine,
) -> Self {
Self {
raft,
state_machine,
}
}
pub fn active_write_version(&self) -> u8 {
self.state_machine.active_write_version()
}
pub fn local_capabilities(&self) -> NodeCapabilities {
NodeCapabilities::local(self.active_write_version())
}
fn membership_snapshot(&self) -> (u64, bool, Vec<(u64, OpenraftPeer)>) {
let metrics = self.raft.metrics();
let snapshot = metrics.borrow_watched();
let local_node_id = snapshot.id;
let is_leader = snapshot.state == ServerState::Leader;
let members = snapshot
.membership_config
.nodes()
.map(|(node_id, node)| (*node_id, node.clone()))
.collect();
(local_node_id, is_leader, members)
}
pub async fn gather_member_capabilities<S>(
&self,
source: &S,
) -> Result<Vec<(u64, NodeCapabilities)>, FormatActivationError>
where
S: CapabilitySource<Node = OpenraftPeer>,
{
let (local_node_id, _is_leader, members) = self.membership_snapshot();
gather_with(local_node_id, self.local_capabilities(), &members, source).await
}
pub async fn run_activation_gate<S>(
&self,
target: u8,
source: &S,
) -> Result<BTreeSet<u64>, FormatActivationError>
where
S: CapabilitySource<Node = OpenraftPeer>,
{
let (_local_node_id, is_leader, _members) = self.membership_snapshot();
if !is_leader {
return Err(FormatActivationError::NotLeader);
}
if let Err(err) = target_in_local_readable_range(target) {
tracing::warn!(
target = target,
?err,
"format activation rejected: target outside local readable range"
);
crate::observability::record_rejected_by_gate();
return Err(err);
}
let gathered = self.gather_member_capabilities(source).await?;
let min_member_max_readable = gathered
.iter()
.map(|(_node_id, capability)| capability.max_readable_version)
.min()
.unwrap_or(tsoracle_openraft_toolkit::MIN_READABLE_VERSION);
crate::observability::record_min_member_read_capability(min_member_max_readable);
for (node_id, capability) in &gathered {
tracing::debug!(
node_id = node_id,
min_readable_version = capability.min_readable_version,
max_readable_version = capability.max_readable_version,
active_write_version = capability.active_write_version,
"format activation: member capability"
);
}
match all_members_can_read(target, &gathered) {
Some(gated_members) => {
tracing::info!(
target = target,
gated_members = ?gated_members,
"format activation gate passed"
);
Ok(gated_members)
}
None => {
let incapable: Vec<(u64, u8)> = gathered
.iter()
.filter(|(_, capabilities)| capabilities.max_readable_version < target)
.map(|(node_id, capabilities)| (*node_id, capabilities.max_readable_version))
.collect();
tracing::warn!(
target = target,
min_member_read_capability = min_member_max_readable,
incapable = ?incapable,
"format activation rejected by all-members gate"
);
crate::observability::record_rejected_by_gate();
Err(FormatActivationError::MembersBelowTarget { target, incapable })
}
}
}
async fn submit_set_format_version(
&self,
target: u8,
gated_members: BTreeSet<u64>,
) -> Result<ApplyOutcome, FormatActivationError> {
match self
.raft
.client_write(HighWaterCommand::SetFormatVersion(
SetFormatVersionPayload {
target,
gated_members,
},
))
.await
{
Ok(resp) => {
crate::observability::record_committed();
Ok(resp.data.outcome)
}
Err(err) => Err(FormatActivationError::ProposalFailed(err.to_string())),
}
}
pub async fn initiate_format_activation<S>(
&self,
target: u8,
source: &S,
) -> Result<(), FormatActivationError>
where
S: CapabilitySource<Node = OpenraftPeer>,
{
let gated_members = self.run_activation_gate(target, source).await?;
crate::observability::record_proposed();
let outcome = self
.submit_set_format_version(target, gated_members)
.await?;
classify_activation_outcome(outcome, target)
}
}
fn classify_activation_outcome(
outcome: ApplyOutcome,
target: u8,
) -> Result<(), FormatActivationError> {
match outcome {
ApplyOutcome::FormatActivated { target: applied } if applied == target => Ok(()),
ApplyOutcome::FormatActivationTargetOutOfRange { target: applied } => {
Err(FormatActivationError::TargetOutOfRange {
target: applied,
min: tsoracle_openraft_toolkit::MIN_READABLE_VERSION,
max: tsoracle_openraft_toolkit::MAX_READABLE_VERSION,
})
}
ApplyOutcome::FormatActivated { .. }
| ApplyOutcome::FormatActivationNoop { .. }
| ApplyOutcome::Advanced => {
Err(FormatActivationError::MembershipChangedSinceGate { target })
}
other @ (ApplyOutcome::DenseAdvanced { .. }
| ApplyOutcome::DenseCardinalityExceeded { .. }
| ApplyOutcome::DenseOverflow) => {
unreachable!("dense ApplyOutcome {other:?} returned from a SetFormatVersion entry")
}
}
}
#[async_trait]
impl OpenraftHighWaterHost for StandaloneHost {
type Config = TypeConfig;
fn metrics(&self) -> WatchReceiverOf<Self::Config, RaftMetrics<Self::Config>> {
self.raft.metrics()
}
async fn current_high_water(&self) -> Result<u64, ConsensusError> {
if let Err(e) = self.raft.ensure_linearizable(ReadPolicy::ReadIndex).await {
return Err(classify_read_error(e));
}
Ok(self.state_machine.current_value().await)
}
async fn submit_advance(&self, at_least: u64) -> Result<u64, ConsensusError> {
match self
.raft
.client_write(HighWaterCommand::Advance(AdvancePayload { at_least }))
.await
{
Ok(resp) => Ok(resp.data.value),
Err(e) => Err(classify_client_write_error(e)),
}
}
fn active_write_version(&self) -> u8 {
StandaloneHost::active_write_version(self)
}
async fn submit_advance_dense(
&self,
key: &tsoracle_core::SeqKey,
count: u32,
) -> Result<u64, ConsensusError> {
match self
.raft
.client_write(HighWaterCommand::AdvanceDense {
key: key.clone(),
count,
})
.await
{
Ok(resp) => match resp.data.outcome {
crate::type_config::ApplyOutcome::DenseAdvanced { start } => Ok(start),
crate::type_config::ApplyOutcome::DenseCardinalityExceeded { cap } => {
Err(ConsensusError::SeqKeyCardinalityExceeded { cap })
}
crate::type_config::ApplyOutcome::DenseOverflow => Err(ConsensusError::SeqOverflow),
other => Err(ConsensusError::PermanentDriver(
format!("unexpected ApplyOutcome for AdvanceDense: {other:?}").into(),
)),
},
Err(e) => Err(classify_client_write_error(e)),
}
}
async fn current_dense_seq(&self, key: &tsoracle_core::SeqKey) -> Result<u64, ConsensusError> {
if let Err(e) = self.raft.ensure_linearizable(ReadPolicy::ReadIndex).await {
return Err(classify_read_error(e));
}
Ok(self.state_machine.dense_value(key.as_str()))
}
}
fn classify_read_error(
err: RaftError<TypeConfig, LinearizableReadError<TypeConfig>>,
) -> ConsensusError {
match err {
RaftError::APIError(LinearizableReadError::ForwardToLeader(_)) => {
ConsensusError::NotLeader { observed: None }
}
RaftError::Fatal(_) => ConsensusError::PermanentDriver(Box::new(err)),
_ => ConsensusError::TransientDriver(Box::new(err)),
}
}
fn classify_client_write_error(
err: RaftError<TypeConfig, ClientWriteError<TypeConfig>>,
) -> ConsensusError {
match err {
RaftError::APIError(ClientWriteError::ForwardToLeader(_)) => {
ConsensusError::NotLeader { observed: None }
}
RaftError::Fatal(_) => ConsensusError::PermanentDriver(Box::new(err)),
_ => ConsensusError::TransientDriver(Box::new(err)),
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeSet;
use openraft::error::{
ChangeMembershipError, EmptyMembership, Fatal, ForwardToLeader, QuorumNotEnough,
};
use super::*;
#[test]
fn fatal_read_error_classifies_as_permanent_driver() {
let err =
RaftError::<TypeConfig, LinearizableReadError<TypeConfig>>::Fatal(Fatal::Panicked);
assert!(matches!(
classify_read_error(err),
ConsensusError::PermanentDriver(_)
));
}
#[test]
fn fatal_client_write_error_classifies_as_permanent_driver() {
let err = RaftError::<TypeConfig, ClientWriteError<TypeConfig>>::Fatal(Fatal::Stopped);
assert!(matches!(
classify_client_write_error(err),
ConsensusError::PermanentDriver(_)
));
}
#[test]
fn forward_to_leader_read_error_classifies_as_not_leader() {
let err = RaftError::<TypeConfig, LinearizableReadError<TypeConfig>>::APIError(
LinearizableReadError::ForwardToLeader(ForwardToLeader::empty()),
);
assert!(matches!(
classify_read_error(err),
ConsensusError::NotLeader { observed: None }
));
}
#[test]
fn quorum_not_enough_read_error_classifies_as_transient_driver() {
let err = RaftError::<TypeConfig, LinearizableReadError<TypeConfig>>::APIError(
LinearizableReadError::QuorumNotEnough(QuorumNotEnough {
cluster: String::new(),
got: BTreeSet::new(),
}),
);
assert!(matches!(
classify_read_error(err),
ConsensusError::TransientDriver(_)
));
}
#[test]
fn change_membership_client_write_error_classifies_as_transient_driver() {
let err = RaftError::<TypeConfig, ClientWriteError<TypeConfig>>::APIError(
ClientWriteError::ChangeMembershipError(ChangeMembershipError::EmptyMembership(
EmptyMembership {},
)),
);
assert!(matches!(
classify_client_write_error(err),
ConsensusError::TransientDriver(_)
));
}
#[tokio::test]
async fn host_active_write_version_delegates_to_the_shared_cell() {
let cell = tsoracle_openraft_toolkit::ActiveWriteVersion::default();
let store: std::sync::Arc<dyn crate::snapshot_store::SnapshotStore> =
std::sync::Arc::new(crate::snapshot_store::InMemorySnapshotStore::new());
let state_machine =
HighWaterStateMachine::with_store_and_active_version(store, cell.clone()).expect("sm");
assert_eq!(
state_machine.active_write_version(),
tsoracle_openraft_toolkit::BASELINE_WRITE_VERSION
);
cell.set(7);
assert_eq!(state_machine.active_write_version(), 7);
let _assert_signature: fn(&StandaloneHost) -> u8 = StandaloneHost::active_write_version;
}
#[test]
fn classify_activation_outcome_maps_success_to_ok() {
let result = classify_activation_outcome(ApplyOutcome::FormatActivated { target: 7 }, 7);
assert!(matches!(result, Ok(())));
}
#[test]
fn classify_activation_outcome_maps_noop_to_membership_changed() {
let result =
classify_activation_outcome(ApplyOutcome::FormatActivationNoop { target: 7 }, 7);
assert!(matches!(
result,
Err(FormatActivationError::MembershipChangedSinceGate { target: 7 })
));
}
#[test]
fn classify_activation_outcome_maps_advanced_to_membership_changed() {
let result = classify_activation_outcome(ApplyOutcome::Advanced, 7);
assert!(matches!(
result,
Err(FormatActivationError::MembershipChangedSinceGate { target: 7 })
));
}
#[test]
fn classify_activation_outcome_rejects_target_mismatch() {
let result = classify_activation_outcome(ApplyOutcome::FormatActivated { target: 8 }, 7);
assert!(matches!(
result,
Err(FormatActivationError::MembershipChangedSinceGate { target: 7 })
));
}
#[test]
fn membership_changed_since_gate_is_a_distinct_variant() {
let err = FormatActivationError::MembershipChangedSinceGate { target: 4 };
assert!(matches!(
err,
FormatActivationError::MembershipChangedSinceGate { target: 4 }
));
}
#[test]
fn classify_activation_outcome_maps_target_out_of_range_to_local_range_error() {
let outcome = ApplyOutcome::FormatActivationTargetOutOfRange { target: 1 };
let result = classify_activation_outcome(outcome, 1);
match result {
Err(FormatActivationError::TargetOutOfRange { target, min, max }) => {
assert_eq!(target, 1);
assert_eq!(min, tsoracle_openraft_toolkit::MIN_READABLE_VERSION);
assert_eq!(max, tsoracle_openraft_toolkit::MAX_READABLE_VERSION);
}
other => panic!("expected TargetOutOfRange, got: {other:?}"),
}
}
}