use async_trait::async_trait;
use openraft::Raft;
use openraft::RaftMetrics;
use openraft::ReadPolicy;
use openraft::error::{ClientWriteError, LinearizableReadError, RaftError};
use openraft::type_config::alias::WatchReceiverOf;
use tsoracle_consensus::ConsensusError;
use crate::host::OpenraftHighWaterHost;
use crate::log_entry::HighWaterCommand;
use crate::state_machine::HighWaterStateMachine;
use crate::type_config::TypeConfig;
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,
}
}
}
#[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::Bump { target: at_least })
.await
{
Ok(resp) => Ok(resp.data.value),
Err(e) => Err(classify_client_write_error(e)),
}
}
}
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(_)
));
}
}