use crate::error::OpenRaft;
use crate::error::Result;
use crate::error::RockRaftError;
use crate::node::RaftNode;
use crate::raft::types::AppliedState;
use crate::raft::types::BatchWriteReq;
use crate::raft::types::Cmd;
use crate::raft::types::ForwardRequestBody;
use crate::raft::types::ForwardResponse;
use crate::raft::types::GetKVReq;
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::ScanPrefixReq;
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 std::time::SystemTime;
use std::time::UNIX_EPOCH;
use tonic::Status;
use tracing::debug;
use tracing::error;
use tracing::info;
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 handle(&self, body: ForwardRequestBody) -> Result<ForwardResponse> {
match body {
ForwardRequestBody::Join(req) => self.handle_join(req).await,
ForwardRequestBody::Leave(req) => self.handle_leave(req).await,
ForwardRequestBody::GetMembers(req) => self.handle_get_members(req).await,
ForwardRequestBody::Write(entry) => self.handle_write(entry).await,
ForwardRequestBody::GetKV(req) => self.handle_get_kv(req).await,
ForwardRequestBody::ScanPrefix(req) => self.handle_scan_prefix(req).await,
ForwardRequestBody::BatchWrite(req) => self.handle_batch_write(req).await,
ForwardRequestBody::Txn(req) => self.handle_txn(req).await,
}
}
async fn handle_join(&self, req: JoinRequest) -> Result<ForwardResponse> {
let node_id = req.node_id;
info!("Handling join request for node {}", node_id);
let metrics = self.raft().metrics().borrow_watched().clone();
let membership = metrics.membership_config.membership();
let voters: BTreeSet<u64> = membership.voter_ids().collect();
if voters.contains(&node_id) {
info!("Node {} already in membership, skipping join", node_id);
return Ok(ForwardResponse::Join(()));
}
info!("Syncing existing nodes to state machine");
for (id, node) in membership.nodes() {
info!("Syncing node {} to state machine", id);
let entry = LogEntry::new(Cmd::AddNode {
node: node.clone(),
overriding: true,
});
if let Err(e) = self.write(entry).await {
error!("Failed to sync node {}: {:?}", id, e);
return Err(RockRaftError::TonicStatus(Status::internal(format!(
"Failed to sync node {}: {}",
id, e
))));
}
}
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,
});
if let Err(e) = self.write(entry).await {
error!("Failed to write AddNode entry: {:?}", e);
return Err(RockRaftError::TonicStatus(Status::internal(format!(
"Failed to join node: {}",
e
))));
}
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);
if let Err(e) = self.raft().change_membership(msg, false).await {
error!("Failed to change membership: {:?}", e);
return Err(RockRaftError::TonicStatus(Status::internal(format!(
"Failed to join node: {}",
e
))));
}
info!("Node {} joined successfully", node_id);
Ok(ForwardResponse::Join(()))
}
async fn handle_leave(&self, req: LeaveRequest) -> Result<ForwardResponse> {
let node_id = req.node_id;
let metrics = self.raft().metrics().borrow_watched().clone();
let membership = metrics.membership_config.membership();
let voters: BTreeSet<u64> = membership.voter_ids().collect();
if !voters.contains(&node_id) {
return Ok(ForwardResponse::Leave(()));
}
let entry = LogEntry::new(Cmd::RemoveNode { node_id });
if let Err(e) = self.write(entry).await {
error!("Failed to leave node: {:?}", e);
return Err(RockRaftError::TonicStatus(Status::internal(format!(
"Failed to leave node: {}",
e
))));
}
let mut remove_voters: BTreeSet<u64> = BTreeSet::new();
remove_voters.insert(node_id);
if let Err(e) = self.raft().change_membership(remove_voters, true).await {
error!("Failed to leave node: {:?}", e);
return Err(RockRaftError::TonicStatus(Status::internal(format!(
"Failed to leave node: {}",
e
))));
}
Ok(ForwardResponse::Leave(()))
}
async fn handle_write(&self, entry: LogEntry) -> Result<ForwardResponse> {
match self.write(entry).await {
Ok(applied_state) => Ok(ForwardResponse::Write(applied_state)),
Err(e) => {
error!("Failed to write log entry: {:?}", e);
Err(RockRaftError::TonicStatus(Status::internal(format!(
"Failed to write log entry: {}",
e
))))
}
}
}
async fn handle_batch_write(&self, req: BatchWriteReq) -> Result<ForwardResponse> {
if req.entries.is_empty() {
return Ok(ForwardResponse::BatchWrite(AppliedState::None));
}
let entry = LogEntry::new(Cmd::BatchUpsertKV {
entries: req.entries,
});
match self.write(entry).await {
Ok(applied_state) => Ok(ForwardResponse::BatchWrite(applied_state)),
Err(e) => {
error!("Failed to write batch log entry: {:?}", e);
Err(RockRaftError::TonicStatus(Status::internal(format!(
"Failed to write batch log entry: {}",
e
))))
}
}
}
async fn handle_get_kv(&self, req: GetKVReq) -> Result<ForwardResponse> {
match self.node.state_machine().get_kv(&req.key) {
Ok(value) => Ok(ForwardResponse::GetKV(value)),
Err(e) => {
error!("Failed to get kv: {:?}", e);
Err(RockRaftError::TonicStatus(Status::internal(format!(
"Failed to get kv: {}",
e
))))
}
}
}
async fn handle_scan_prefix(&self, req: ScanPrefixReq) -> Result<ForwardResponse> {
match self.node.state_machine().scan_prefix(&req.prefix) {
Ok(results) => Ok(ForwardResponse::ScanPrefix(results)),
Err(e) => {
error!("Failed to scan prefix: {:?}", e);
Err(RockRaftError::TonicStatus(Status::internal(format!(
"Failed to scan prefix: {}",
e
))))
}
}
}
async fn handle_get_members(&self, _req: GetMembersReq) -> Result<ForwardResponse> {
match self.node.state_machine().get_nodes() {
Ok(nodes) => Ok(ForwardResponse::GetMembers(nodes)),
Err(e) => {
error!("Failed to get members: {:?}", e);
Err(RockRaftError::TonicStatus(Status::internal(format!(
"Failed to get members: {}",
e
))))
}
}
}
async fn handle_txn(&self, req: TxnReq) -> Result<ForwardResponse> {
let entry = LogEntry::new(Cmd::Txn { req, result: None });
match self.write(entry).await {
Ok(AppliedState::Txn(reply)) => {
Ok(ForwardResponse::Txn(reply))
}
Ok(_) => {
error!("Unexpected AppliedState from transaction");
Err(RockRaftError::TonicStatus(Status::internal(
"Unexpected response from transaction".to_string(),
)))
}
Err(e) => {
error!("Failed to execute transaction: {:?}", e);
Err(RockRaftError::TonicStatus(Status::internal(format!(
"Failed to execute transaction: {}",
e
))))
}
}
}
pub async fn write(&self, mut entry: LogEntry) -> Result<AppliedState> {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64;
entry.time_ms = Some(now);
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"
);
match e {
RaftError::APIError(api_err) => {
Err(RockRaftError::OpenRaft(OpenRaft::ClientWrite(api_err)))
}
RaftError::Fatal(fatal_err) => Err(RockRaftError::OpenRaft(OpenRaft::Fatal(fatal_err))),
}
}
}
}
}