use crate::error::{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::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))
}
}
})
};
}
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,
});
map_err_log!(
self.write(entry).await,
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,
});
map_err_log!(self.write(entry).await, "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);
map_raft_err_log!(
self.raft().change_membership(msg, false).await,
"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 });
map_err_log!(self.write(entry).await, "Failed to leave node")?;
let mut remove_voters: BTreeSet<u64> = BTreeSet::new();
remove_voters.insert(node_id);
map_raft_err_log!(
self.raft().change_membership(remove_voters, true).await,
"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,
});
map_err_log!(
self.do_write(entry).await,
"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 map_err_log!(self.do_write(entry).await, "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 map_raft_err_log!(self.raft().client_write(entry).await, "client write") {
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(e)
}
}
}
}