use std::collections::BTreeMap;
use std::collections::BTreeSet;
use std::io;
use std::sync::Arc;
use std::time::Duration;
use openraft::BasicNode;
use openraft::ChangeMembers;
use openraft::Raft;
use openraft::ReadPolicy;
use openraft::async_runtime::WatchReceiver;
use openraft::errors::ChangeMembershipError;
use openraft::errors::ClientWriteError;
use openraft::errors::InitializeError;
use openraft::errors::RaftError;
use serde::Serialize;
use tokio::time::sleep;
use crate::app::EzApp;
use crate::config::EzConfig;
use crate::network::EzNetworkFactory;
use crate::storage::EzStorage;
use crate::storage::adapter::StorageAdapter;
use crate::type_config::OpenRaftTypes;
type ORTypes<T> = OpenRaftTypes<T>;
pub type ORRaft<T> = Raft<ORTypes<T>, Arc<StorageAdapter<T>>>;
pub struct EzRaft<T>
where T: EzApp
{
node_id: u64,
addr: String,
storage: Arc<StorageAdapter<T>>,
raft: ORRaft<T>,
}
impl<T> Clone for EzRaft<T>
where T: EzApp
{
fn clone(&self) -> Self {
Self {
node_id: self.node_id,
addr: self.addr.clone(),
storage: self.storage.clone(),
raft: self.raft.clone(),
}
}
}
impl<T> EzRaft<T>
where T: EzApp
{
pub async fn create(
http_addr: impl ToString,
app: T,
storage: impl EzStorage<T>,
config: EzConfig,
) -> Result<Self, io::Error> {
Self::new(http_addr, app, storage, config, None).await
}
pub async fn join(
http_addr: impl ToString,
seed_addr: impl ToString,
app: T,
storage: impl EzStorage<T>,
config: EzConfig,
) -> Result<Self, io::Error> {
Self::new(http_addr, app, storage, config, Some(seed_addr.to_string())).await
}
async fn new(
http_addr: impl ToString,
app: T,
storage: impl EzStorage<T>,
config: EzConfig,
seed_addr: Option<String>,
) -> Result<Self, io::Error> {
let http_addr = http_addr.to_string();
let adapter = StorageAdapter::new(storage, app).await?;
let adapter = Arc::new(adapter);
let node_id = if let Some(id) = adapter.node_id().await {
id
} else if let Some(seed) = &seed_addr {
let id = request_join(seed, &http_addr).await?;
adapter.save_meta(|m| m.node_id = Some(id)).await?;
id
} else {
let id = 0;
adapter.save_meta(|m| m.node_id = Some(id)).await?;
id
};
let (log_store, sm_store) = (adapter.clone(), adapter.clone());
let raft_config = config.to_raft_config()?;
let raft_config = Arc::new(raft_config);
let network = EzNetworkFactory::new()?;
let raft = Raft::new(node_id, raft_config, network, log_store, sm_store)
.await
.map_err(|e| io::Error::other(e.to_string()))?;
if node_id == 0 {
let nodes = BTreeMap::from_iter([(node_id, BasicNode::new(http_addr.clone()))]);
match raft.initialize(nodes).await {
Ok(()) | Err(RaftError::APIError(InitializeError::NotAllowed(_))) => {}
Err(e) => return Err(io::Error::other(e.to_string())),
}
}
let this = Self {
node_id,
addr: http_addr,
storage: adapter,
raft,
};
tokio::spawn(this.clone().reconcile_learners());
Ok(this)
}
async fn reconcile_learners(self) {
let promotable = |m: &openraft::RaftMetrics<ORTypes<T>>| -> Option<u64> {
let replication = m.replication.as_ref()?;
let membership = m.membership_config.membership();
membership.learner_ids().find(|id| {
let matched = replication.get(id).and_then(|log_id| log_id.as_ref()).map(|log_id| log_id.index);
matched >= m.last_log_index
})
};
loop {
let res =
self.raft.wait(None).metrics(|m| promotable(m).is_some(), "a learner is ready for promotion").await;
let Ok(metrics) = res else {
return;
};
let Some(node_id) = promotable(&metrics) else {
continue;
};
if let Err(e) = self.promote_to_voter(node_id).await {
tracing::error!("failed to promote node {} to voter: {}", node_id, e);
sleep(PROMOTE_RETRY_INTERVAL).await;
}
}
}
pub async fn write(&self, req: T::Request) -> Result<T::Response, io::Error> {
let mut last_err = String::new();
for _ in 0..WRITE_ATTEMPTS {
let err = match self.raft.client_write(req.clone()).await {
Ok(resp) => return resp.data.ok_or_else(|| io::Error::other("write produced no response")),
Err(e) => e,
};
let RaftError::APIError(ClientWriteError::ForwardToLeader(forward)) = &err else {
return Err(io::Error::other(err.to_string()));
};
match forward.leader_node.as_ref().map(|n| n.addr.as_str()) {
Some(leader) if leader != self.addr => return forward_write::<T>(leader, &req).await,
_ => {
last_err = err.to_string();
sleep(WRITE_RETRY_INTERVAL).await;
}
}
}
Err(io::Error::other(format!(
"write gave up after {} attempts: {}",
WRITE_ATTEMPTS, last_err
)))
}
pub async fn read<F, R>(&self, read: F) -> R
where F: FnOnce(&T) -> R {
let sm = self.storage.sm_state.lock().await;
read(&sm.app)
}
pub async fn linearizable(&self) -> Result<(), io::Error> {
self.raft
.ensure_linearizable(ReadPolicy::ReadIndex)
.await
.map_err(|e| io::Error::other(e.to_string()))?;
Ok(())
}
pub(crate) async fn add_learner(&self, node_id: u64, addr: String) -> Result<(), io::Error> {
let node = BasicNode::new(addr);
self.raft.add_learner(node_id, node, false).await.map_err(|e| io::Error::other(e.to_string()))?;
Ok(())
}
async fn promote_to_voter(&self, node_id: u64) -> Result<(), io::Error> {
let caught_up = |m: &openraft::RaftMetrics<ORTypes<T>>| {
let Some(replication) = m.replication.as_ref() else {
return true;
};
let matched = replication.get(&node_id).and_then(|log_id| log_id.as_ref()).map(|log_id| log_id.index);
matched >= m.last_log_index
};
let mut last_err = String::new();
for _ in 0..PROMOTE_ATTEMPTS {
let metrics = self
.raft
.wait(None)
.metrics(caught_up, "learner catches up before promotion")
.await
.map_err(|e| io::Error::other(e.to_string()))?;
if metrics.current_leader != Some(self.node_id) {
return Ok(());
}
if metrics.membership_config.membership().voter_ids().any(|id| id == node_id) {
return Ok(());
}
let change = ChangeMembers::AddVoterIds(BTreeSet::from([node_id]));
let err = match self.raft.change_membership(change, false).await {
Ok(_) => return Ok(()),
Err(e) => e,
};
match &err {
RaftError::APIError(ClientWriteError::ChangeMembershipError(ChangeMembershipError::InProgress(_))) => {
last_err = err.to_string();
sleep(PROMOTE_RETRY_INTERVAL).await;
}
RaftError::APIError(ClientWriteError::ForwardToLeader(_)) => return Ok(()),
_ => return Err(io::Error::other(err.to_string())),
}
}
Err(io::Error::other(format!(
"promotion of node {} gave up after {} attempts: {}",
node_id, PROMOTE_ATTEMPTS, last_err
)))
}
pub async fn change_membership(&self, change: ChangeMembers<u64, BasicNode>) -> Result<(), io::Error> {
self.raft.change_membership(change, false).await.map_err(|e| io::Error::other(e.to_string()))?;
Ok(())
}
pub fn is_leader(&self) -> bool {
self.raft.is_leader()
}
pub async fn metrics(&self) -> openraft::RaftMetrics<ORTypes<T>> {
self.raft.metrics().borrow_watched().clone()
}
pub async fn serve(self) -> Result<(), io::Error> {
crate::server::run(self).await
}
pub fn node_id(&self) -> u64 {
self.node_id
}
pub fn addr(&self) -> &str {
&self.addr
}
pub fn inner(&self) -> &ORRaft<T> {
&self.raft
}
}
const WRITE_ATTEMPTS: usize = 20;
const WRITE_RETRY_INTERVAL: Duration = Duration::from_millis(500);
const PROMOTE_ATTEMPTS: usize = 20;
const PROMOTE_RETRY_INTERVAL: Duration = Duration::from_millis(500);
const FORWARD_WRITE_TIMEOUT: Duration = Duration::from_secs(10);
async fn forward_write<T>(leader_addr: &str, req: &T::Request) -> Result<T::Response, io::Error>
where T: EzApp {
let client = reqwest::Client::builder()
.no_proxy()
.timeout(FORWARD_WRITE_TIMEOUT)
.build()
.map_err(|e| io::Error::other(e.to_string()))?;
let url = format!("http://{}/api/write", leader_addr);
let resp = client
.post(&url)
.json(req)
.send()
.await
.map_err(|e| io::Error::other(format!("forwarding write to {} failed: {}", url, e)))?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
return Err(io::Error::other(format!("{} responded {}: {}", url, status, body)));
}
resp.json().await.map_err(|e| io::Error::other(format!("failed to parse write response: {}", e)))
}
#[derive(Debug, Serialize)]
struct JoinRequest {
addr: String,
}
type JoinResponse = Result<u64, Option<String>>;
const JOIN_ATTEMPTS: usize = 20;
const JOIN_RETRY_INTERVAL: Duration = Duration::from_millis(500);
const JOIN_TIMEOUT: Duration = Duration::from_secs(5);
async fn request_join(seed_addr: &str, my_addr: &str) -> Result<u64, io::Error> {
let client = reqwest::Client::builder()
.no_proxy()
.timeout(JOIN_TIMEOUT)
.build()
.map_err(|e| io::Error::other(e.to_string()))?;
let mut target_addr = seed_addr.to_string();
let mut last_err = "cluster did not accept the join".to_string();
for _ in 0..JOIN_ATTEMPTS {
let url = format!("http://{}/api/join", target_addr);
let req = JoinRequest {
addr: my_addr.to_string(),
};
let resp = match client.post(&url).json(&req).send().await {
Ok(resp) => resp,
Err(e) => {
last_err = format!("join request to {} failed: {}", url, e);
sleep(JOIN_RETRY_INTERVAL).await;
continue;
}
};
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
last_err = format!("{} responded {}: {}", url, status, body);
sleep(JOIN_RETRY_INTERVAL).await;
continue;
}
let join_resp: JoinResponse =
resp.json().await.map_err(|e| io::Error::other(format!("failed to parse join response: {}", e)))?;
match join_resp {
Ok(node_id) => return Ok(node_id),
Err(Some(leader)) => {
last_err = format!("{} redirected to {}", url, leader);
target_addr = leader;
}
Err(None) => {
last_err = format!("{} knows of no leader", url);
sleep(JOIN_RETRY_INTERVAL).await;
}
}
}
Err(io::Error::other(format!(
"join gave up after {} attempts: {}",
JOIN_ATTEMPTS, last_err
)))
}