use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use openraft::error::{
InstallSnapshotError, NetworkError as OpenraftNetworkError, RPCError, RaftError, RemoteError,
Timeout, Unreachable,
};
use openraft::network::{RPCOption, RPCTypes, RaftNetwork, RaftNetworkFactory};
use openraft::raft::{
AppendEntriesRequest, AppendEntriesResponse, InstallSnapshotRequest, InstallSnapshotResponse,
VoteRequest, VoteResponse,
};
use openraft::{BasicNode, Raft};
use serde::{Deserialize, Serialize};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::RwLock;
use crate::raft::{OxirsNodeId, OxirsTypeConfig};
const MAX_MESSAGE_SIZE: usize = 16 * 1024 * 1024;
#[derive(Debug, Serialize, Deserialize)]
enum RaftWireRequest {
AppendEntries(AppendEntriesRequest<OxirsTypeConfig>),
Vote(VoteRequest<OxirsNodeId>),
InstallSnapshot(InstallSnapshotRequest<OxirsTypeConfig>),
}
#[derive(Debug, Serialize, Deserialize)]
enum RaftWireResponse {
AppendEntries(Result<AppendEntriesResponse<OxirsNodeId>, RaftError<OxirsNodeId>>),
Vote(Result<VoteResponse<OxirsNodeId>, RaftError<OxirsNodeId>>),
InstallSnapshot(
Result<InstallSnapshotResponse<OxirsNodeId>, RaftError<OxirsNodeId, InstallSnapshotError>>,
),
}
#[derive(Debug, thiserror::Error)]
enum RaftTransportError {
#[error("connect to {addr}: {source}")]
Connect {
addr: SocketAddr,
#[source]
source: std::io::Error,
},
#[error("io error talking to {addr}: {source}")]
Io {
addr: SocketAddr,
#[source]
source: std::io::Error,
},
#[error("frame to/from {addr} ({len} bytes) exceeds max message size ({max} bytes)")]
FrameTooLarge {
addr: SocketAddr,
len: usize,
max: usize,
},
#[error("failed to encode raft RPC for {addr}: {message}")]
Encode { addr: SocketAddr, message: String },
#[error("failed to decode raft RPC frame from {addr}: {message}")]
Decode { addr: SocketAddr, message: String },
#[error("no known network address for peer node {peer}")]
UnknownPeer { peer: OxirsNodeId },
#[error("peer {addr} replied with a response of the wrong RPC kind")]
MismatchedResponse { addr: SocketAddr },
#[error("raft RPC to {addr} timed out after {timeout:?}")]
Timeout { addr: SocketAddr, timeout: Duration },
}
async fn write_frame<W, T>(
stream: &mut W,
message: &T,
peer_addr: SocketAddr,
max_size: usize,
) -> Result<(), RaftTransportError>
where
W: AsyncWrite + Unpin,
T: Serialize,
{
let body =
oxicode::serde::encode_to_vec(message, oxicode::config::standard()).map_err(|e| {
RaftTransportError::Encode {
addr: peer_addr,
message: e.to_string(),
}
})?;
if body.len() > max_size {
return Err(RaftTransportError::FrameTooLarge {
addr: peer_addr,
len: body.len(),
max: max_size,
});
}
let io_err = |source: std::io::Error| RaftTransportError::Io {
addr: peer_addr,
source,
};
let len = body.len() as u32;
stream.write_all(&len.to_be_bytes()).await.map_err(io_err)?;
stream.write_all(&body).await.map_err(io_err)?;
stream.flush().await.map_err(io_err)?;
Ok(())
}
async fn read_frame<R, T>(
stream: &mut R,
peer_addr: SocketAddr,
max_size: usize,
) -> Result<T, RaftTransportError>
where
R: AsyncRead + Unpin,
T: for<'de> Deserialize<'de>,
{
try_read_frame(stream, peer_addr, max_size)
.await?
.ok_or(RaftTransportError::Io {
addr: peer_addr,
source: std::io::Error::from(std::io::ErrorKind::UnexpectedEof),
})
}
async fn try_read_frame<R, T>(
stream: &mut R,
peer_addr: SocketAddr,
max_size: usize,
) -> Result<Option<T>, RaftTransportError>
where
R: AsyncRead + Unpin,
T: for<'de> Deserialize<'de>,
{
let io_err = |source: std::io::Error| RaftTransportError::Io {
addr: peer_addr,
source,
};
let mut len_buf = [0u8; 4];
let first_byte = stream.read(&mut len_buf[0..1]).await.map_err(io_err)?;
if first_byte == 0 {
return Ok(None);
}
stream
.read_exact(&mut len_buf[1..4])
.await
.map_err(io_err)?;
let len = u32::from_be_bytes(len_buf) as usize;
if len > max_size {
return Err(RaftTransportError::FrameTooLarge {
addr: peer_addr,
len,
max: max_size,
});
}
let mut body = vec![0u8; len];
stream.read_exact(&mut body).await.map_err(io_err)?;
let (message, _) = oxicode::serde::decode_from_slice(&body, oxicode::config::standard())
.map_err(|e| RaftTransportError::Decode {
addr: peer_addr,
message: e.to_string(),
})?;
Ok(Some(message))
}
pub(crate) struct OxirsRaftNetworkFactory {
self_id: OxirsNodeId,
peer_addresses: Arc<RwLock<HashMap<OxirsNodeId, SocketAddr>>>,
}
impl OxirsRaftNetworkFactory {
pub(crate) fn new(
self_id: OxirsNodeId,
peer_addresses: Arc<RwLock<HashMap<OxirsNodeId, SocketAddr>>>,
) -> Self {
Self {
self_id,
peer_addresses,
}
}
}
impl RaftNetworkFactory<OxirsTypeConfig> for OxirsRaftNetworkFactory {
type Network = OxirsRaftNetworkClient;
async fn new_client(&mut self, target: OxirsNodeId, node: &BasicNode) -> Self::Network {
let addr = node.addr.parse::<SocketAddr>().ok();
OxirsRaftNetworkClient {
self_id: self.self_id,
target,
addr,
peer_addresses: Arc::clone(&self.peer_addresses),
}
}
}
pub(crate) struct OxirsRaftNetworkClient {
self_id: OxirsNodeId,
target: OxirsNodeId,
addr: Option<SocketAddr>,
peer_addresses: Arc<RwLock<HashMap<OxirsNodeId, SocketAddr>>>,
}
impl OxirsRaftNetworkClient {
async fn resolve_address(&self) -> Result<SocketAddr, RaftTransportError> {
if let Some(addr) = self.addr {
return Ok(addr);
}
self.peer_addresses
.read()
.await
.get(&self.target)
.copied()
.ok_or(RaftTransportError::UnknownPeer { peer: self.target })
}
async fn call(
&self,
addr: SocketAddr,
request: RaftWireRequest,
hard_ttl: Duration,
) -> Result<RaftWireResponse, RaftTransportError> {
let attempt = async {
let mut stream = TcpStream::connect(addr)
.await
.map_err(|e| RaftTransportError::Connect { addr, source: e })?;
if let Err(e) = stream.set_nodelay(true) {
tracing::debug!("failed to set TCP_NODELAY on raft RPC connection to {addr}: {e}");
}
if let Err(e) = stream.set_zero_linger() {
tracing::debug!(
"failed to set zero SO_LINGER on raft RPC connection to {addr}: {e}"
);
}
write_frame(&mut stream, &request, addr, MAX_MESSAGE_SIZE).await?;
read_frame(&mut stream, addr, MAX_MESSAGE_SIZE).await
};
let result = match tokio::time::timeout(hard_ttl, attempt).await {
Ok(result) => result,
Err(_elapsed) => Err(RaftTransportError::Timeout {
addr,
timeout: hard_ttl,
}),
};
if let Err(ref e) = result {
tracing::debug!("raft RPC call to {addr} (hard_ttl={hard_ttl:?}) failed: {e}");
}
result
}
fn transport_err_to_rpc<E>(
&self,
err: RaftTransportError,
action: RPCTypes,
hard_ttl: Duration,
) -> RPCError<OxirsNodeId, BasicNode, RaftError<OxirsNodeId, E>>
where
E: std::error::Error,
{
match err {
RaftTransportError::Timeout { .. } => RPCError::Timeout(Timeout {
action,
id: self.self_id,
target: self.target,
timeout: hard_ttl,
}),
RaftTransportError::Connect { .. } | RaftTransportError::UnknownPeer { .. } => {
RPCError::Unreachable(Unreachable::new(&err))
}
RaftTransportError::Io { .. }
| RaftTransportError::FrameTooLarge { .. }
| RaftTransportError::Encode { .. }
| RaftTransportError::Decode { .. }
| RaftTransportError::MismatchedResponse { .. } => {
RPCError::Network(OpenraftNetworkError::new(&err))
}
}
}
}
impl RaftNetwork<OxirsTypeConfig> for OxirsRaftNetworkClient {
async fn append_entries(
&mut self,
rpc: AppendEntriesRequest<OxirsTypeConfig>,
option: RPCOption,
) -> Result<
AppendEntriesResponse<OxirsNodeId>,
RPCError<OxirsNodeId, BasicNode, RaftError<OxirsNodeId>>,
> {
let hard_ttl = option.hard_ttl();
let addr = self
.resolve_address()
.await
.map_err(|e| self.transport_err_to_rpc(e, RPCTypes::AppendEntries, hard_ttl))?;
let wire_resp = self
.call(addr, RaftWireRequest::AppendEntries(rpc), hard_ttl)
.await
.map_err(|e| self.transport_err_to_rpc(e, RPCTypes::AppendEntries, hard_ttl))?;
match wire_resp {
RaftWireResponse::AppendEntries(result) => {
result.map_err(|e| RPCError::RemoteError(RemoteError::new(self.target, e)))
}
_ => Err(RPCError::Network(OpenraftNetworkError::new(
&RaftTransportError::MismatchedResponse { addr },
))),
}
}
async fn vote(
&mut self,
rpc: VoteRequest<OxirsNodeId>,
option: RPCOption,
) -> Result<VoteResponse<OxirsNodeId>, RPCError<OxirsNodeId, BasicNode, RaftError<OxirsNodeId>>>
{
let hard_ttl = option.hard_ttl();
let addr = self
.resolve_address()
.await
.map_err(|e| self.transport_err_to_rpc(e, RPCTypes::Vote, hard_ttl))?;
let wire_resp = self
.call(addr, RaftWireRequest::Vote(rpc), hard_ttl)
.await
.map_err(|e| self.transport_err_to_rpc(e, RPCTypes::Vote, hard_ttl))?;
match wire_resp {
RaftWireResponse::Vote(result) => {
result.map_err(|e| RPCError::RemoteError(RemoteError::new(self.target, e)))
}
_ => Err(RPCError::Network(OpenraftNetworkError::new(
&RaftTransportError::MismatchedResponse { addr },
))),
}
}
async fn install_snapshot(
&mut self,
rpc: InstallSnapshotRequest<OxirsTypeConfig>,
option: RPCOption,
) -> Result<
InstallSnapshotResponse<OxirsNodeId>,
RPCError<OxirsNodeId, BasicNode, RaftError<OxirsNodeId, InstallSnapshotError>>,
> {
let hard_ttl = option.hard_ttl();
let addr = self
.resolve_address()
.await
.map_err(|e| self.transport_err_to_rpc(e, RPCTypes::InstallSnapshot, hard_ttl))?;
let wire_resp = self
.call(addr, RaftWireRequest::InstallSnapshot(rpc), hard_ttl)
.await
.map_err(|e| self.transport_err_to_rpc(e, RPCTypes::InstallSnapshot, hard_ttl))?;
match wire_resp {
RaftWireResponse::InstallSnapshot(result) => {
result.map_err(|e| RPCError::RemoteError(RemoteError::new(self.target, e)))
}
_ => Err(RPCError::Network(OpenraftNetworkError::new(
&RaftTransportError::MismatchedResponse { addr },
))),
}
}
}
pub(crate) async fn serve_raft_rpc(listener: TcpListener, raft: Raft<OxirsTypeConfig>) {
loop {
let (stream, peer_addr) = match listener.accept().await {
Ok(pair) => pair,
Err(e) => {
tracing::warn!("raft RPC listener accept() failed: {e}; continuing to listen");
continue;
}
};
if let Err(e) = stream.set_nodelay(true) {
tracing::debug!(
"failed to set TCP_NODELAY on raft RPC connection from {peer_addr}: {e}"
);
}
let raft = raft.clone();
tokio::spawn(async move {
if let Err(e) = serve_raft_connection(stream, peer_addr, raft).await {
tracing::debug!("raft RPC connection from {peer_addr} ended: {e}");
}
});
}
}
async fn serve_raft_connection(
mut stream: TcpStream,
peer_addr: SocketAddr,
raft: Raft<OxirsTypeConfig>,
) -> Result<(), RaftTransportError> {
loop {
let request: RaftWireRequest =
match try_read_frame(&mut stream, peer_addr, MAX_MESSAGE_SIZE).await? {
Some(request) => request,
None => return Ok(()), };
let response = match request {
RaftWireRequest::AppendEntries(rpc) => {
RaftWireResponse::AppendEntries(raft.append_entries(rpc).await)
}
RaftWireRequest::Vote(rpc) => RaftWireResponse::Vote(raft.vote(rpc).await),
RaftWireRequest::InstallSnapshot(rpc) => {
RaftWireResponse::InstallSnapshot(raft.install_snapshot(rpc).await)
}
};
write_frame(&mut stream, &response, peer_addr, MAX_MESSAGE_SIZE).await?;
}
}
#[cfg(test)]
mod tests {
use super::*;
use openraft::Vote;
#[tokio::test]
async fn request_frame_round_trips() {
let original = RaftWireRequest::Vote(VoteRequest::new(Vote::new(3, 7u64), None));
let (mut a, mut b) = tokio::io::duplex(64 * 1024);
let addr: SocketAddr = "127.0.0.1:0".parse().expect("valid addr");
let writer = tokio::spawn(async move {
write_frame(&mut a, &original, addr, MAX_MESSAGE_SIZE)
.await
.expect("write_frame failed");
});
let decoded: RaftWireRequest = read_frame(&mut b, addr, MAX_MESSAGE_SIZE)
.await
.expect("read_frame failed");
writer.await.expect("writer task panicked");
match decoded {
RaftWireRequest::Vote(v) => {
assert_eq!(v.vote.leader_id().term, 3);
}
other => panic!("unexpected variant: {other:?}"),
}
}
#[tokio::test]
async fn try_read_frame_reports_clean_eof_as_none() {
let (a, b) = tokio::io::duplex(1024);
drop(a); let mut b = b;
let addr: SocketAddr = "127.0.0.1:0".parse().expect("valid addr");
let result: Option<RaftWireRequest> = try_read_frame(&mut b, addr, MAX_MESSAGE_SIZE)
.await
.expect("try_read_frame should not error on clean EOF");
assert!(result.is_none());
}
#[tokio::test]
async fn write_frame_rejects_oversize_message() {
let msg = RaftWireRequest::Vote(VoteRequest::new(Vote::new(1, 1u64), None));
let (mut a, _b) = tokio::io::duplex(1024);
let addr: SocketAddr = "127.0.0.1:0".parse().expect("valid addr");
let err = write_frame(&mut a, &msg, addr, 1)
.await
.expect_err("oversize frame must be rejected");
assert!(
matches!(err, RaftTransportError::FrameTooLarge { .. }),
"unexpected error: {err}"
);
}
#[tokio::test]
async fn resolve_address_fails_loudly_for_unknown_peer() {
let client = OxirsRaftNetworkClient {
self_id: 1,
target: 99,
addr: None,
peer_addresses: Arc::new(RwLock::new(HashMap::new())),
};
let err = client
.resolve_address()
.await
.expect_err("must fail for an unregistered peer");
assert!(matches!(err, RaftTransportError::UnknownPeer { peer: 99 }));
}
}