use crate::error::ManagementError;
use crate::error::Result;
use crate::error::RockRaftError;
use crate::error::StartupError;
use crate::grpc::JoinConnectionFactory;
use crate::raft::protobuf as pb;
use crate::raft::protobuf::raft_service_client::RaftServiceClient;
use crate::raft::types::ForwardRequestBody;
use crate::raft::types::JoinRequest;
use anyerror::AnyError;
use openraft::error::{InitializeError, RaftError};
use std::collections::BTreeMap;
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::Mutex;
use tokio::net::TcpListener;
use tokio::sync::{broadcast, oneshot};
use tokio::time::{Duration, timeout};
use tracing::debug;
use tracing::error;
use tracing::info;
use std::result::Result as StdResult;
use tokio::time::sleep;
use openraft::Config as OpenRaftConfig;
use openraft::Raft;
use openraft::async_runtime::watch::WatchReceiver;
use tonic::Status;
use tonic::transport::Server;
use super::LeaderHandler;
use super::parsed_config::ParsedConfig;
use crate::config::Config;
use crate::engine::RocksDBEngine;
use crate::raft::grpc_client::ClientPool;
use crate::raft::network::NetworkFactory;
use crate::raft::protobuf::raft_service_server::RaftServiceServer;
use crate::raft::store::RocksLogStore;
use crate::raft::store::RocksStateMachine;
use crate::raft::store::column_family_list;
use crate::raft::types::{
AppliedState, BatchWriteReply, BatchWriteReq, ForwardRequest, ForwardResponse, LogEntry, Node,
TxnReply, TxnReq, TypeConfig, decode,
};
use crate::raft::types::{
ForwardToLeader, GetKVReply, GetKVReq, GetMembersReply, GetMembersReq, LeaveRequest, NodeId,
ScanPrefixReply, ScanPrefixReq,
};
use crate::service::RaftServiceImpl;
pub struct RaftNode {
#[allow(dead_code)]
engine: Arc<RocksDBEngine>,
raft: Arc<Raft<TypeConfig>>,
config: ParsedConfig,
#[allow(dead_code)]
factory: NetworkFactory,
state_machine: Arc<RocksStateMachine>,
shutdown_tx: broadcast::Sender<()>,
_shutdown_rx: broadcast::Receiver<()>,
service_handle: Mutex<Option<tokio::task::JoinHandle<()>>>,
}
impl RaftNode {
pub fn raft(&self) -> &Arc<Raft<TypeConfig>> {
&self.raft
}
pub fn state_machine(&self) -> &Arc<RocksStateMachine> {
&self.state_machine
}
pub async fn shutdown(&self) -> Result<()> {
let _ = self.shutdown_tx.send(());
let handle = self.service_handle.lock().unwrap().take();
if let Some(h) = handle {
h.await.ok();
}
Ok(())
}
pub async fn create(config: &Config) -> Result<Arc<Self>> {
let engine = Arc::new(RocksDBEngine::new(
&config.rocksdb.data_path,
config.rocksdb.max_open_files,
column_family_list(),
));
let node_id = config.node_id;
let log_store = RocksLogStore::create(engine.db.clone())?;
let data_dir = PathBuf::from(&config.rocksdb.data_path);
let state_machine = RocksStateMachine::new(engine.db.clone(), data_dir).await?;
let client_pool = Arc::new(ClientPool::new(10));
let factory = NetworkFactory::new(client_pool);
let raft_config = OpenRaftConfig::default();
let raft = Arc::new(
Raft::new(
node_id,
Arc::new(raft_config),
factory.clone(),
log_store,
state_machine.clone(),
)
.await
.map_err(crate::error::OpenRaft::Fatal)?,
);
let (shutdown_tx, shutdown_rx_for_struct) = broadcast::channel(1);
Ok(Arc::new(Self {
engine,
raft,
config: ParsedConfig::from(config)?,
factory,
state_machine: Arc::new(state_machine),
shutdown_tx,
_shutdown_rx: shutdown_rx_for_struct,
service_handle: Mutex::new(None),
}))
}
pub async fn start(raft_node: Arc<Self>) -> Result<()> {
let config = &raft_node.config;
Self::start_raft_service(raft_node.clone()).await?;
if config.raft_single {
let node = Node {
node_id: config.node_id,
endpoint: config.raft_endpoint.clone(),
};
raft_node.init_cluster(node).await?;
} else {
raft_node.join_cluster().await?;
}
Ok(())
}
async fn start_raft_service(raft_node: Arc<Self>) -> Result<()> {
let raft_endpoint = raft_node.config.raft_endpoint.clone();
let mut shutdown_rx = raft_node.shutdown_tx.subscribe();
let raft_node_for_service = raft_node.clone();
let (startup_tx, startup_rx) = oneshot::channel::<StdResult<(), String>>();
let handle = tokio::task::spawn(async move {
tracing::info!("Starting Raft gRPC service on {}", raft_endpoint);
let raft_service = RaftServiceImpl::new(raft_node_for_service);
let listener = match TcpListener::bind(&raft_endpoint.to_string()).await {
Ok(l) => l,
Err(e) => {
let err_msg = format!("Failed to bind gRPC server to {}: {}", raft_endpoint, e);
tracing::error!("{}", err_msg);
let _ = startup_tx.send(Err(err_msg));
return;
}
};
if startup_tx.send(Ok(())).is_err() {
error!("Failed to signal startup completion");
return;
}
info!("Raft gRPC service listening on {}", raft_endpoint);
let server_future = Server::builder()
.add_service(RaftServiceServer::new(raft_service))
.serve_with_incoming(tokio_stream::wrappers::TcpListenerStream::new(listener));
tokio::select! {
_ = shutdown_rx.recv() => {
info!("Raft gRPC service received shutdown signal, shutting down...");
}
result = server_future => {
match result {
Ok(_) => info!("Raft gRPC service finished normally"),
Err(e) => error!("Raft gRPC service error: {}", e),
}
}
}
info!("Raft gRPC service stopped");
});
match startup_rx.await {
Ok(Ok(())) => {
*raft_node.service_handle.lock().unwrap() = Some(handle);
info!("Raft gRPC service started successfully");
Ok(())
}
Ok(Err(err_msg)) => {
let _ = handle.await;
Err(RockRaftError::Startup(StartupError::OtherError(err_msg)))
}
Err(_) => {
let _ = handle.await;
Err(RockRaftError::Startup(StartupError::OtherError(
"gRPC service startup task failed unexpectedly".to_string(),
)))
}
}
}
async fn get_leader(&self) -> Result<Option<NodeId>> {
let deadline = Duration::from_millis(2000);
let mut metrics_rx = self.raft.metrics();
let result = timeout(deadline, async {
loop {
if let Some(leader) = metrics_rx.borrow_watched().current_leader {
return Ok(Some(leader));
}
if let Err(e) = WatchReceiver::changed(&mut metrics_rx).await {
let error_msg = format!("Metrics watch error: {:?}", e);
tracing::debug!("{}", error_msg);
return Err(AnyError::error(error_msg).into());
}
}
})
.await;
match result {
Ok(inner_result) => inner_result,
Err(_) => {
Ok(None)
}
}
}
async fn assume_leader(&self) -> StdResult<LeaderHandler<'_>, ForwardToLeader> {
let current_node_id = *self.raft.node_id();
match self.get_leader().await {
Ok(Some(leader_id)) => {
if leader_id == current_node_id {
Ok(LeaderHandler::new(self))
} else {
Err(ForwardToLeader {
leader_id: Some(leader_id),
leader_node: None,
})
}
}
Ok(None) => {
Err(ForwardToLeader {
leader_id: None,
leader_node: None,
})
}
Err(_) => {
Err(ForwardToLeader {
leader_id: None,
leader_node: None,
})
}
}
}
fn is_in_cluster(&self) -> Result<bool> {
let last_membership = self
.state_machine
.get_last_membership()
.map_err(|e| AnyError::error(format!("get_last_membership error: {}", e)))?;
let node_id = *self.raft.node_id();
let is_voter = last_membership
.membership()
.voter_ids()
.any(|id| id == node_id);
Ok(is_voter)
}
async fn init_cluster(&self, node: Node) -> Result<()> {
if node.node_id != *self.raft.node_id() {
let err = StartupError::invalid_config(format!(
"Node ID {} does not match current node ID {}",
node.node_id,
self.raft.node_id()
));
return Err(crate::error::RockRaftError::from(err));
}
if node.endpoint.addr().is_empty() {
let err = StartupError::invalid_config("Node endpoint address cannot be empty");
return Err(crate::error::RockRaftError::from(err));
}
info!("Adding node {} to state machine", node.node_id);
self.state_machine.add_node(node.clone()).map_err(|e| {
error!("Failed to add node: {}", e);
StartupError::OtherError(format!("Failed to add node: {}", e))
})?;
info!("Node {} added to state machine successfully", node.node_id);
let node_id = node.node_id;
let mut nodes = BTreeMap::new();
nodes.insert(node_id, node);
if let Err(e) = self.raft.initialize(nodes).await {
match e {
RaftError::APIError(e) => match e {
InitializeError::NotAllowed(e) => {
info!("Already initialized: {}", e);
}
InitializeError::NotInMembers(e) => {
let err = StartupError::InvalidConfig(e.to_string());
return Err(err.into());
}
},
RaftError::Fatal(e) => {
let err = StartupError::OtherError(e.to_string());
return Err(err.into());
}
}
}
Ok(())
}
pub async fn join_cluster(&self) -> Result<()> {
let config = &self.config;
if config.raft_join.is_empty() {
info!("'--join' is empty, do not need joining cluster");
return Ok(());
}
if self.is_in_cluster()? {
info!("node has already in cluster, do not need joining cluster");
return Ok(());
}
self.do_join_cluster().await?;
Ok(())
}
async fn do_join_cluster(&self) -> StdResult<(), ManagementError> {
let config = &self.config;
let addrs = &config.raft_join;
let mut errors = vec![];
let raft_address = config.raft_endpoint.to_string();
let raft_advertise_address = config.raft_advertise_endpoint.to_string();
for addr in addrs {
if addr == &raft_address || addr == &raft_advertise_address {
debug!("ignore join cluster via self node address {}", addr);
continue;
}
for _i in 0..3 {
let result = self.join_via(addr).await;
info!("join cluster via {} result: {:?}", addr, result);
match result {
Ok(x) => return Ok(x),
Err(api_error) => {
let can_retry = api_error.is_retryable();
if can_retry {
debug!("try to connect to addr {} again", addr);
sleep(Duration::from_millis(1_000)).await;
continue;
} else {
errors.push(api_error);
break;
}
}
}
}
}
Err(ManagementError::Join(AnyError::error(format!(
"fail to join node-{} to cluster via {:?}, errors: {}",
self.raft.node_id(),
addrs,
errors
.into_iter()
.map(|e| e.to_string())
.collect::<Vec<_>>()
.join(", ")
))))
}
async fn send_forward_request(
&self,
addr: &String,
request: ForwardRequest,
) -> Result<pb::RaftReply> {
let timeout = Some(Duration::from_millis(10_000));
let chan_result = JoinConnectionFactory::create_rpc_channel(addr, timeout, None).await;
let channel = match chan_result {
Ok(channel) => channel,
Err(e) => {
error!("Failed to connect to {}: {:?}", addr, e);
return Err(e);
}
};
let mut raft_client = RaftServiceClient::new(channel);
let response = raft_client
.forward(request)
.await
.map_err(|e| Status::internal(format!("Failed to forward request: {}", e)))?;
Ok(response.into_inner())
}
async fn join_via(&self, addr: &String) -> Result<()> {
let config = &self.config;
let join_req = JoinRequest {
node_id: config.node_id,
endpoint: config.raft_endpoint.clone(),
};
let req = ForwardRequest {
forward_to_leader: 1,
body: ForwardRequestBody::Join(join_req),
};
let reply = self.send_forward_request(addr, req).await?;
if reply.error.is_empty() {
Ok(())
} else {
Err(RockRaftError::Raft(format!(
"Join failed: {:?}",
String::from_utf8_lossy(&reply.error)
)))
}
}
pub async fn write(&self, entry: LogEntry) -> StdResult<AppliedState, Status> {
debug!("write log entry: {:?}", entry);
let request = ForwardRequest {
forward_to_leader: 1,
body: ForwardRequestBody::Write(entry),
};
match self.handle_forward_request(request).await {
Ok(ForwardResponse::Write(applied_state)) => Ok(applied_state),
Ok(_) => Err(Status::internal("Unexpected response type from leader")),
Err(e) => Err(Self::error_to_status(e)),
}
}
pub async fn batch_write(&self, req: BatchWriteReq) -> StdResult<BatchWriteReply, Status> {
debug!("batch write: {:?}", req);
let request = ForwardRequest {
forward_to_leader: 1,
body: ForwardRequestBody::BatchWrite(req),
};
match self.handle_forward_request(request).await {
Ok(ForwardResponse::BatchWrite(applied_state)) => Ok(applied_state),
Ok(_) => Err(Status::internal("Unexpected response type from leader")),
Err(e) => Err(Self::error_to_status(e)),
}
}
pub async fn txn(&self, req: TxnReq) -> StdResult<TxnReply, Status> {
debug!("transaction: {:?}", req);
let request = ForwardRequest {
forward_to_leader: 1,
body: ForwardRequestBody::Txn(req),
};
match self.handle_forward_request(request).await {
Ok(ForwardResponse::Txn(reply)) => Ok(reply),
Ok(_) => Err(Status::internal("Unexpected response type from leader")),
Err(e) => Err(Self::error_to_status(e)),
}
}
pub async fn getset(
&self,
key: impl ToString,
value: impl AsRef<[u8]>,
) -> StdResult<Option<Vec<u8>>, Status> {
use crate::raft::types::UpsertKV;
let req = TxnReq::new(vec![]) .if_then(UpsertKV::insert(key, value.as_ref()))
.with_return_previous();
match self.txn(req).await? {
TxnReply::Success { prev_values, .. } => Ok(prev_values.into_iter().next().flatten()),
}
}
pub async fn read(&self, req: GetKVReq) -> StdResult<GetKVReply, Status> {
debug!("read kv: {:?}", req);
let request = ForwardRequest {
forward_to_leader: 1,
body: ForwardRequestBody::GetKV(req),
};
match self.handle_forward_request(request).await {
Ok(ForwardResponse::GetKV(value)) => Ok(value),
Ok(_) => Err(Status::internal("Unexpected response type from leader")),
Err(e) => Err(Self::error_to_status(e)),
}
}
pub async fn scan_prefix(&self, req: ScanPrefixReq) -> StdResult<ScanPrefixReply, Status> {
debug!("scan_prefix: {:?}", req);
let request = ForwardRequest {
forward_to_leader: 1,
body: ForwardRequestBody::ScanPrefix(req),
};
match self.handle_forward_request(request).await {
Ok(ForwardResponse::ScanPrefix(results)) => Ok(results),
Ok(_) => Err(Status::internal("Unexpected response type from leader")),
Err(e) => Err(Self::error_to_status(e)),
}
}
pub async fn join(&self, req: JoinRequest) -> StdResult<(), Status> {
debug!("join node: {:?}", req);
let request = ForwardRequest {
forward_to_leader: 1,
body: ForwardRequestBody::Join(req),
};
match self.handle_forward_request(request).await {
Ok(ForwardResponse::Join(())) => Ok(()),
Ok(_) => Err(Status::internal("Unexpected response type from leader")),
Err(e) => Err(Self::error_to_status(e)),
}
}
pub async fn leave(&self, req: LeaveRequest) -> StdResult<(), Status> {
debug!("leave node: {:?}", req);
let request = ForwardRequest {
forward_to_leader: 1,
body: ForwardRequestBody::Leave(req),
};
match self.handle_forward_request(request).await {
Ok(ForwardResponse::Leave(())) => Ok(()),
Ok(_) => Err(Status::internal("Unexpected response type from leader")),
Err(e) => Err(Self::error_to_status(e)),
}
}
pub async fn get_members(&self, req: GetMembersReq) -> StdResult<GetMembersReply, Status> {
debug!("get members: {:?}", req);
let leader_handler = LeaderHandler::new(self);
match leader_handler
.handle(ForwardRequestBody::GetMembers(req))
.await
{
Ok(ForwardResponse::GetMembers(members)) => Ok(members),
Ok(_) => Err(Status::internal("Unexpected response type")),
Err(e) => Err(Self::error_to_status(e)),
}
}
async fn forward_request_to_leader(
&self,
leader_id: NodeId,
request: ForwardRequest,
) -> Result<ForwardResponse> {
let membership = self.state_machine.get_last_membership().map_err(|e| {
RockRaftError::TonicStatus(Status::internal(format!("Failed to get membership: {}", e)))
})?;
let leader_node = membership
.membership()
.get_node(&leader_id)
.ok_or_else(|| {
RockRaftError::TonicStatus(Status::internal("Leader id not found in membership"))
})?;
let leader_addr = leader_node.endpoint.to_string();
let reply = self.send_forward_request(&leader_addr, request).await?;
if reply.error.is_empty() {
let forward_response: ForwardResponse = decode(&reply.data).map_err(|e| {
RockRaftError::TonicStatus(Status::internal(format!(
"Failed to deserialize response: {}",
e
)))
})?;
Ok(forward_response)
} else {
Err(RockRaftError::TonicStatus(Status::internal(format!(
"Leader returned error: {:?}",
String::from_utf8_lossy(&reply.error)
))))
}
}
fn error_to_status(error: RockRaftError) -> Status {
match error {
RockRaftError::TonicStatus(status) => status,
_ => Status::internal(error.to_string()),
}
}
fn is_retriable_error(error: &RockRaftError) -> bool {
match error {
RockRaftError::TonicStatus(status) => {
matches!(status.code(), tonic::Code::Unavailable)
}
_ => error.is_retryable(),
}
}
pub async fn handle_forward_request(&self, request: ForwardRequest) -> Result<ForwardResponse> {
debug!("recv forward req: {:?}", request);
const MAX_RETRIES: u32 = 20;
const RETRY_INTERVAL: Duration = Duration::from_secs(1);
for attempt in 0..MAX_RETRIES {
match self.assume_leader().await {
Ok(_) => {
let leader_handler = LeaderHandler::new(self);
return leader_handler.handle(request.body.clone()).await;
}
Err(forward_err) => {
let retry_reason = match forward_err.leader_id {
Some(leader_id) => {
match self
.forward_request_to_leader(leader_id, request.clone())
.await
{
Ok(response) => return Ok(response),
Err(e) => {
if Self::is_retriable_error(&e) {
Some(format!("Failed to forward request ({e})"))
} else {
return Err(e);
}
}
}
}
None => {
Some("No leader available to forward request".to_string())
}
};
if let Some(reason) = retry_reason
&& attempt < MAX_RETRIES - 1
{
debug!("{}, retrying {}/{}", reason, attempt + 1, MAX_RETRIES);
sleep(RETRY_INTERVAL).await;
continue;
}
return Err(RockRaftError::TonicStatus(Status::internal(
"No leader available to forward request after max retries",
)));
}
}
}
Err(RockRaftError::TonicStatus(Status::internal(
"No leader available to forward request after max retries",
)))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::RaftConfig;
use crate::config::RocksdbConfig;
use crate::raft::types::Endpoint;
use tempfile::tempdir;
fn create_test_config(data_dir: &str, node_id: u64, addr: &str) -> Config {
Config {
node_id,
raft: RaftConfig {
address: addr.to_string(),
advertise_host: "".to_string(),
single: true,
join: vec![],
},
rocksdb: RocksdbConfig {
data_path: data_dir.to_string(),
max_open_files: 1024,
},
}
}
async fn setup_nodes(raft_node: &RaftNode, node_ids: Vec<u64>) {
let sm = &raft_node.state_machine;
for node_id in node_ids {
let node = Node {
node_id,
endpoint: Endpoint::new("127.0.0.1", 1000 + node_id as u32),
};
sm.add_node(node).unwrap();
}
}
#[tokio::test]
async fn test_is_in_cluster_node_exists() -> Result<()> {
let temp_dir = tempdir().unwrap().keep();
let data_path = temp_dir.into_os_string().into_string().unwrap();
let config = create_test_config(&data_path, 1, "127.0.0.1:5001");
let raft_node = RaftNode::create(&config).await?;
setup_nodes(&raft_node, vec![1, 2, 3]).await;
let result = raft_node.is_in_cluster()?;
assert!(result, "Node 1 should be in the cluster");
Ok(())
}
#[tokio::test]
async fn test_is_in_cluster_node_not_exists() -> Result<()> {
let temp_dir = tempdir().unwrap().keep();
let data_path = temp_dir.into_os_string().into_string().unwrap();
let config = create_test_config(&data_path, 4, "127.0.0.1:5004");
let raft_node = RaftNode::create(&config).await?;
setup_nodes(&raft_node, vec![1, 2, 3]).await;
let result = raft_node.is_in_cluster()?;
assert!(!result, "Node 4 should not be in the cluster");
Ok(())
}
}