use std::fmt::Display;
use std::future::Future;
use std::io;
use openraft::AnyError;
use openraft::BasicNode;
use openraft::OptionalSend;
use openraft::error::Infallible;
use openraft::error::NetworkError;
use openraft::error::RPCError;
use openraft::error::ReplicationClosed;
use openraft::error::StreamingError;
use openraft::error::Unreachable;
use openraft::network::RPCOption;
use openraft::network::RaftNetworkFactory;
use openraft::network::v2::RaftNetworkV2;
use openraft::raft::AppendEntriesRequest;
use openraft::raft::AppendEntriesResponse;
use openraft::raft::SnapshotResponse;
use openraft::raft::VoteRequest;
use openraft::raft::VoteResponse;
use openraft::type_config::alias::SnapshotOf;
use openraft::type_config::alias::VoteOf;
use reqwest::Client;
use serde::Serialize;
use serde::de::DeserializeOwned;
use crate::app::EzApp;
use crate::snapshot::EzSnapshotData;
use crate::snapshot::EzSnapshotMeta;
use crate::type_config::EzVote;
use crate::type_config::OpenRaftTypes;
type C<T> = OpenRaftTypes<T>;
pub struct EzNetworkFactory {
client: Client,
}
impl EzNetworkFactory {
pub fn new() -> Result<Self, io::Error> {
let client = Client::builder().no_proxy().build().map_err(io::Error::other)?;
Ok(Self { client })
}
}
impl<T> RaftNetworkFactory<C<T>> for EzNetworkFactory
where T: EzApp
{
type Network = Network;
async fn new_client(&mut self, _target: u64, node: &BasicNode) -> Self::Network {
let addr = node.addr.clone();
let client = self.client.clone();
Network { addr, client }
}
}
pub struct Network {
addr: String,
client: Client,
}
impl Network {
async fn request<Req, Resp, Err, Cfg>(
&mut self,
uri: impl Display,
req: Req,
option: &RPCOption,
) -> Result<Result<Resp, Err>, RPCError<Cfg>>
where
Cfg: openraft::RaftTypeConfig,
Req: Serialize + 'static,
Resp: Serialize + DeserializeOwned,
Err: std::error::Error + Serialize + DeserializeOwned,
{
let url = format!("http://{}/{}", self.addr, uri);
let resp = self.client.post(url.clone()).timeout(option.soft_ttl()).json(&req).send().await.map_err(|e| {
if e.is_connect() || e.is_timeout() {
RPCError::Unreachable(Unreachable::new(&e))
} else {
RPCError::Network(NetworkError::new(&e))
}
})?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
let err = AnyError::error(format!("{} responded {}: {}", url, status, body));
return Err(RPCError::Network(NetworkError::new(&err)));
}
let res: Result<Resp, Err> = resp.json().await.map_err(|e| NetworkError::new(&e))?;
Ok(res)
}
}
impl<T> RaftNetworkV2<C<T>> for Network
where T: EzApp
{
type SnapshotData = EzSnapshotData;
async fn append_entries(
&mut self,
req: AppendEntriesRequest<C<T>>,
option: RPCOption,
) -> Result<AppendEntriesResponse<C<T>>, RPCError<C<T>>> {
let res = self.request::<_, _, Infallible, C<T>>("raft/append", req, &option).await?;
Ok(res.unwrap())
}
async fn full_snapshot(
&mut self,
vote: VoteOf<C<T>>,
snapshot: SnapshotOf<C<T>, Self::SnapshotData>,
cancel: impl Future<Output = ReplicationClosed> + OptionalSend + 'static,
option: RPCOption,
) -> Result<SnapshotResponse<C<T>>, StreamingError<C<T>>> {
let req = SnapshotTransfer {
vote,
meta: snapshot.meta,
data: snapshot.snapshot.into_inner(),
};
tokio::pin!(cancel);
tokio::select! {
closed = &mut cancel => Err(StreamingError::Closed(closed)),
res = self.request::<_, _, Infallible, C<T>>("raft/snapshot", req, &option) => Ok(res?.unwrap()),
}
}
async fn vote(&mut self, req: VoteRequest<C<T>>, option: RPCOption) -> Result<VoteResponse<C<T>>, RPCError<C<T>>> {
let res = self.request::<_, _, Infallible, C<T>>("raft/vote", req, &option).await?;
Ok(res.unwrap())
}
}
#[derive(serde::Deserialize, serde::Serialize)]
pub(crate) struct SnapshotTransfer {
pub vote: EzVote,
pub meta: EzSnapshotMeta,
pub data: Vec<u8>,
}