use crate::error::Error;
use crate::error::ErrorKind;
use crate::error::Result;
use crate::node::RaftNode;
use crate::raft::types::AppliedState;
use crate::raft::types::BatchWriteReq;
use crate::raft::types::Cmd;
use crate::raft::types::GetKVReq;
use crate::raft::types::GetMembersReply;
use crate::raft::types::GetMembersReq;
use crate::raft::types::JoinRequest;
use crate::raft::types::LeaveRequest;
use crate::raft::types::LogEntry;
use crate::raft::types::Node;
use crate::raft::types::NodeId;
use crate::raft::types::ScanPrefixReply;
use crate::raft::types::ScanPrefixReq;
use crate::raft::types::TxnReply;
use crate::raft::types::TxnReq;
use crate::raft::types::TypeConfig;
use openraft::ChangeMembers;
use openraft::Raft;
use openraft::async_runtime::watch::WatchReceiver;
use openraft::error::ClientWriteError;
use openraft::error::RaftError;
use std::collections::BTreeMap;
use std::collections::BTreeSet;
use std::sync::Arc;
use tracing::debug;
use tracing::error;
use tracing::info;
#[macro_export]
macro_rules! map_err_log {
($result:expr, $context:expr) => {
$result.map_err(|e| {
error!("{}: {:?}", $context, e);
Error::internal(format!("{}: {}", $context, e))
})
};
}
#[macro_export]
macro_rules! map_raft_err_log {
($result:expr, $context:expr) => {
$result.map_err(|e| {
error!("{}: {:?}", $context, e);
match e {
RaftError::APIError(api_err) => Error::internal(format!("{}: {}", $context, api_err)),
RaftError::Fatal(fatal_err) => {
Error::internal(format!("{}: fatal raft error: {}", $context, fatal_err))
}
}
})
};
}
fn preserve_redirect(e: Error, context: &str) -> Error {
if matches!(e.kind(), ErrorKind::ForwardToLeader { .. }) {
debug!("{}: {:?}", context, e);
return e;
}
error!("{}: {:?}", context, e);
Error::internal(format!("{}: {}", context, e))
}
fn map_client_write_err(
e: RaftError<TypeConfig, ClientWriteError<TypeConfig>>,
context: &str,
) -> Error {
if let RaftError::APIError(ClientWriteError::ForwardToLeader(forward)) = &e {
debug!(
"{}: stepped down, forward to leader {:?}",
context, forward.leader_id
);
return Error::forward_to_leader(forward.leader_id);
}
error!("{}: {:?}", context, e);
match e {
RaftError::APIError(api_err) => Error::internal(format!("{}: {}", context, api_err)),
RaftError::Fatal(fatal_err) => {
Error::internal(format!("{}: fatal raft error: {}", context, fatal_err))
}
}
}
pub struct LeaderHandler<'a> {
node: &'a RaftNode,
}
impl<'a> LeaderHandler<'a> {
pub fn new(node: &'a RaftNode) -> Self {
Self { node }
}
pub fn node(&self) -> &RaftNode {
self.node
}
pub fn raft(&self) -> &Arc<Raft<TypeConfig>> {
self.node.raft()
}
pub async fn write(&self, entry: LogEntry) -> Result<AppliedState> {
self.do_write(entry).await
}
fn is_voter(&self, node_id: NodeId) -> bool {
self
.raft()
.metrics()
.borrow_watched()
.membership_config
.membership()
.voter_ids()
.any(|id| id == node_id)
}
pub async fn add_node(&self, req: JoinRequest) -> Result<()> {
let node_id = req.node_id;
info!("Handling join request for node {}", node_id);
if self.is_voter(node_id) {
info!("Node {} already in membership, skipping join", node_id);
return Ok(());
}
let membership = self
.raft()
.metrics()
.borrow_watched()
.membership_config
.membership()
.clone();
for (id, node) in membership.nodes() {
info!("Syncing node {} to state machine", id);
let entry = LogEntry::new(Cmd::AddNode {
node: node.clone(),
overriding: true,
});
self
.write(entry)
.await
.map_err(|e| preserve_redirect(e, &format!("Failed to sync node {}", id)))?;
}
let node = Node {
node_id,
endpoint: req.endpoint.clone(),
};
info!("Writing AddNode command for node {}", node_id);
let entry = LogEntry::new(Cmd::AddNode {
node: node.clone(),
overriding: false,
});
self
.write(entry)
.await
.map_err(|e| preserve_redirect(e, "Failed to join node"))?;
info!("AddNode command written successfully for node {}", node_id);
info!("Changing membership to add node {} as voter", node_id);
let mut add_voters = BTreeMap::new();
add_voters.insert(node_id, node);
let msg = ChangeMembers::AddVoters(add_voters);
self
.raft()
.change_membership(msg, false)
.await
.map_err(|e| map_client_write_err(e, "Failed to join node"))?;
info!("Node {} joined successfully", node_id);
Ok(())
}
pub async fn remove_node(&self, req: LeaveRequest) -> Result<()> {
let node_id = req.node_id;
if !self.is_voter(node_id) {
return Ok(());
}
let entry = LogEntry::new(Cmd::RemoveNode { node_id });
self
.write(entry)
.await
.map_err(|e| preserve_redirect(e, "Failed to leave node"))?;
let mut remove_voters: BTreeSet<u64> = BTreeSet::new();
remove_voters.insert(node_id);
self
.raft()
.change_membership(remove_voters, true)
.await
.map_err(|e| map_client_write_err(e, "Failed to leave node"))?;
Ok(())
}
pub async fn batch_write(&self, req: BatchWriteReq) -> Result<AppliedState> {
if req.entries.is_empty() {
return Ok(AppliedState::None);
}
let entry = LogEntry::new(Cmd::BatchUpsertKV {
entries: req.entries,
});
self
.do_write(entry)
.await
.map_err(|e| preserve_redirect(e, "Failed to write batch log entry"))
}
pub async fn read(&self, req: GetKVReq) -> Result<Option<Vec<u8>>> {
map_err_log!(
self.node.state_machine().get_kv(&req.key),
"Failed to get kv"
)
}
pub async fn scan_prefix(&self, req: ScanPrefixReq) -> Result<ScanPrefixReply> {
map_err_log!(
self.node.state_machine().scan_prefix(&req.prefix),
"Failed to scan prefix"
)
}
pub async fn get_members(&self, _req: GetMembersReq) -> Result<GetMembersReply> {
map_err_log!(
self.node.state_machine().get_nodes(),
"Failed to get members"
)
}
pub async fn txn(&self, req: TxnReq) -> Result<TxnReply> {
let entry = LogEntry::new(Cmd::Txn { req, result: None });
match self
.do_write(entry)
.await
.map_err(|e| preserve_redirect(e, "Failed to execute transaction"))?
{
AppliedState::Txn(reply) => Ok(reply),
_ => {
error!("Unexpected AppliedState from transaction");
Err(Error::internal("Unexpected response from transaction"))
}
}
}
async fn do_write(&self, mut entry: LogEntry) -> Result<AppliedState> {
entry.time_ms = Some(crate::utils::now_millis());
let node_id = self.raft().node_id();
match self.raft().client_write(entry).await {
Ok(response) => {
debug!(
node_id = %node_id,
log_id = %response.log_id,
"Successfully wrote log entry"
);
Ok(response.data)
}
Err(e) => {
error!(
node_id = %node_id,
error = %e,
"Failed to write log entry"
);
Err(map_client_write_err(e, "client write"))
}
}
}
}
#[cfg(test)]
mod tests {
use super::preserve_redirect;
use crate::error::Error;
#[test]
fn test_preserve_redirect_passes_forward_to_leader_through() {
let err = preserve_redirect(Error::forward_to_leader(Some(3)), "test op");
assert_eq!(err.forward_leader_id(), Some(3));
assert!(err.is_retryable());
}
#[test]
fn test_preserve_redirect_wraps_other_errors() {
let err = preserve_redirect(Error::internal("boom"), "test op");
assert!(!err.is_retryable());
assert!(err.to_string().contains("test op"));
}
}