use std::result::Result as StdResult;
use std::time::Duration;
use openraft::async_runtime::watch::WatchReceiver;
use tokio::time::{sleep, timeout};
use tracing::debug;
use crate::error::{ApiError, Error, Result, RetryReason};
use crate::network::JoinConnectionFactory;
use crate::raft::protobuf::RaftReply;
use crate::raft::protobuf::raft_service_client::RaftServiceClient;
use crate::raft::types::{
ForwardRequest, ForwardResponse, ForwardToLeader, NodeId, RequestPayload, decode,
};
use super::LeaderHandler;
use super::node::RaftNode;
const MAX_RETRIES: u32 = 20;
const RETRY_INITIAL_INTERVAL: Duration = Duration::from_millis(200);
const RETRY_MAX_INTERVAL: Duration = Duration::from_secs(3);
impl RaftNode {
pub(crate) 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(Error::internal(error_msg));
}
}
})
.await;
match result {
Ok(inner_result) => inner_result,
Err(_) => Ok(None),
}
}
pub(crate) 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,
}),
}
}
pub(crate) async fn execute_or_forward(
&self,
payload: RequestPayload,
) -> Result<ForwardResponse> {
for attempt in 0..MAX_RETRIES {
let result = match self.assume_leader().await {
Ok(leader) => Self::dispatch_leader_handler(leader, payload.clone()).await,
Err(forward_err) => match forward_err.leader_id {
Some(leader_id) => {
self
.forward_request_to_leader(
leader_id,
ForwardRequest {
body: payload.clone(),
},
)
.await
}
None => Err(Error::retryable_with_reason(RetryReason::NoLeader)),
},
};
let err = match result {
Ok(response) => return Ok(response),
Err(err) => err,
};
let result = match err.forward_leader_id() {
Some(leader_id) => {
self
.forward_request_to_leader(
leader_id,
ForwardRequest {
body: payload.clone(),
},
)
.await
}
None => Err(err),
};
match result {
Ok(response) => return Ok(response),
Err(err) => {
if err.is_retryable() && attempt < MAX_RETRIES - 1 {
let delay = RETRY_INITIAL_INTERVAL * 2u32.saturating_pow(attempt);
let delay = delay.min(RETRY_MAX_INTERVAL);
debug!(
"execute_or_forward: retry {}/{} after {:?}: {}",
attempt + 1,
MAX_RETRIES,
delay,
err
);
sleep(delay).await;
continue;
}
return Err(err);
}
}
}
Err(Error::internal(
"No leader available to forward request after max retries",
))
}
pub async fn handle_forward_request(&self, request: ForwardRequest) -> Result<ForwardResponse> {
debug!("recv forward req: {:?}", request);
self.execute_or_forward(request.body).await
}
pub(crate) async fn send_forward_request(
&self,
addr: &String,
request: ForwardRequest,
) -> Result<RaftReply> {
let timeout = Some(Duration::from_millis(10_000));
let channel = JoinConnectionFactory::create_rpc_channel(addr, timeout, None)
.await
.map_err(|e| {
tracing::error!("Failed to connect to {}: {:?}", addr, e);
e
})?;
let max_message_size = self.config.raft.grpc_max_message_size();
let mut raft_client = RaftServiceClient::new(channel)
.max_decoding_message_size(max_message_size)
.max_encoding_message_size(max_message_size);
let response = raft_client
.forward(request)
.await
.map_err(|e| Error::internal(format!("Failed to forward request: {}", e)))?;
Ok(response.into_inner())
}
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| Error::internal(format!("Failed to get membership: {}", e)))?;
let leader_node = membership
.membership()
.get_node(&leader_id)
.ok_or_else(|| Error::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| Error::internal(format!("Failed to deserialize response: {}", e)))?;
Ok(forward_response)
} else {
let api_error: ApiError = decode(&reply.error)
.map_err(|e| Error::internal(format!("Failed to deserialize error response: {}", e)))?;
Err(Error::from(api_error))
}
}
pub(crate) async fn dispatch_leader_handler(
leader: LeaderHandler<'_>,
body: RequestPayload,
) -> Result<ForwardResponse> {
match body {
RequestPayload::Write(entry) => {
let result = leader.write(entry).await?;
Ok(ForwardResponse::Write(result))
}
RequestPayload::BatchWrite(req) => {
let result = leader.batch_write(req).await?;
Ok(ForwardResponse::BatchWrite(result))
}
RequestPayload::Txn(req) => {
let result = leader.txn(req).await?;
Ok(ForwardResponse::Txn(result))
}
RequestPayload::GetKV(req) => {
let result = leader.read(req).await?;
Ok(ForwardResponse::GetKV(result))
}
RequestPayload::ScanPrefix(req) => {
let result = leader.scan_prefix(req).await?;
Ok(ForwardResponse::ScanPrefix(result))
}
RequestPayload::Join(req) => {
leader.add_node(req).await?;
Ok(ForwardResponse::Join(()))
}
RequestPayload::Leave(req) => {
leader.remove_node(req).await?;
Ok(ForwardResponse::Leave(()))
}
RequestPayload::GetMembers(req) => {
let result = leader.get_members(req).await?;
Ok(ForwardResponse::GetMembers(result))
}
}
}
}