use crate::Node;
use crate::NodeId;
use crate::app_state::RaftType;
use crate::helpers::{deserialize, serialize};
use crate::network::raft_server::{
RaftStreamRequest, RaftStreamResponse, RaftStreamResponsePayload,
};
use crate::network::web_socket_connect;
use fastwebsockets::{FragmentCollectorRead, Frame, OpCode, Payload, WebSocketWrite};
use hyper::upgrade::Upgraded;
use hyper_util::rt::TokioIo;
use openraft::error::RPCError;
use openraft::error::Unreachable;
use std::collections::HashMap;
use std::ops::Deref;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use tokio::io::{ReadHalf, WriteHalf};
use tokio::sync::oneshot;
use tokio::task::JoinHandle;
use tokio::{select, task, time};
use tracing::{debug, error, info};
#[cfg(feature = "cache")]
use crate::store::state_machine::memory::TypeConfigKV;
#[cfg(feature = "sqlite")]
use crate::store::state_machine::sqlite::TypeConfigSqlite;
#[cfg(any(feature = "cache", feature = "sqlite"))]
use crate::Error;
#[cfg(any(feature = "cache", feature = "sqlite"))]
use openraft::{
error::{InstallSnapshotError, RaftError},
network::{RPCOption, RaftNetwork, RaftNetworkFactory},
raft::{
AppendEntriesRequest, AppendEntriesResponse, InstallSnapshotRequest,
InstallSnapshotResponse, VoteRequest, VoteResponse,
},
};
pub struct NetworkStreaming {
pub node_id: NodeId,
pub tls_config: Option<Arc<rustls::ClientConfig>>,
pub secret_raft: Vec<u8>,
pub raft_type: RaftType,
pub heartbeat_interval: u64,
pub is_raft_stopped: Arc<AtomicBool>,
pub is_startup_finished: Arc<AtomicBool>,
}
#[cfg(feature = "cache")]
impl RaftNetworkFactory<TypeConfigKV> for NetworkStreaming {
type Network = NetworkConnectionStreaming;
#[tracing::instrument(level = "debug", skip_all)]
async fn new_client(&mut self, _target: NodeId, node: &Node) -> Self::Network {
debug!("Building new Raft Cache client with target {}", node);
let (sender, rx) = flume::bounded(1);
let task = tokio::task::spawn(Self::ws_handler(
self.node_id,
self.raft_type.clone(),
node.clone(),
self.tls_config.clone(),
self.secret_raft.clone(),
rx,
self.heartbeat_interval,
self.is_raft_stopped.clone(),
self.is_startup_finished.clone(),
));
NetworkConnectionStreaming {
node: node.clone(),
sender,
task: Some(task),
}
}
}
#[cfg(feature = "sqlite")]
impl RaftNetworkFactory<TypeConfigSqlite> for NetworkStreaming {
type Network = NetworkConnectionStreaming;
#[tracing::instrument(level = "debug", skip_all)]
async fn new_client(&mut self, _target: NodeId, node: &Node) -> Self::Network {
debug!("Building new Raft DB client with target {}", node);
let (sender, rx) = flume::bounded(1);
let task = tokio::task::spawn(Self::ws_handler(
self.node_id,
self.raft_type.clone(),
node.clone(),
self.tls_config.clone(),
self.secret_raft.clone(),
rx,
self.heartbeat_interval,
self.is_raft_stopped.clone(),
self.is_startup_finished.clone(),
));
NetworkConnectionStreaming {
node: node.clone(),
sender,
task: Some(task),
}
}
}
#[derive(Debug)]
enum RaftRequest {
#[cfg(feature = "sqlite")]
AppendDB(
(
oneshot::Sender<Result<RaftStreamResponsePayload, Error>>,
AppendEntriesRequest<TypeConfigSqlite>,
),
),
#[cfg(feature = "sqlite")]
VoteDB(
(
oneshot::Sender<Result<RaftStreamResponsePayload, Error>>,
VoteRequest<u64>,
),
),
#[cfg(feature = "sqlite")]
SnapshotDB(
(
oneshot::Sender<Result<RaftStreamResponsePayload, Error>>,
InstallSnapshotRequest<TypeConfigSqlite>,
),
),
#[cfg(feature = "cache")]
AppendCache(
(
oneshot::Sender<Result<RaftStreamResponsePayload, Error>>,
AppendEntriesRequest<TypeConfigKV>,
),
),
#[cfg(feature = "cache")]
VoteCache(
(
oneshot::Sender<Result<RaftStreamResponsePayload, Error>>,
VoteRequest<u64>,
),
),
#[cfg(feature = "cache")]
SnapshotCache(
(
oneshot::Sender<Result<RaftStreamResponsePayload, Error>>,
InstallSnapshotRequest<TypeConfigKV>,
),
),
StreamResponse(RaftStreamResponse),
ReaderExit,
Shutdown,
}
#[derive(Debug)]
enum WritePayload {
Payload(Vec<u8>),
Close,
}
#[allow(clippy::type_complexity)]
impl NetworkStreaming {
#[allow(clippy::too_many_arguments)]
async fn ws_handler(
this_node: NodeId,
raft_type: RaftType,
node: Node,
tls_config: Option<Arc<rustls::ClientConfig>>,
secret: Vec<u8>,
rx: flume::Receiver<RaftRequest>,
heartbeat_interval: u64,
is_raft_stopped: Arc<AtomicBool>,
is_startup_finished: Arc<AtomicBool>,
) {
let mut request_id = 0usize;
let mut in_flight: HashMap<
usize,
oneshot::Sender<Result<RaftStreamResponsePayload, Error>>,
> = HashMap::with_capacity(4);
let mut shutdown = false;
'outer: loop {
if is_raft_stopped.load(Ordering::Relaxed) {
if !is_startup_finished.load(Ordering::Relaxed) {
debug!("Raft is still starting up - skipping initial connection");
time::sleep(Duration::from_secs(1)).await;
continue;
}
debug!("Raft is stopped - exiting NetworkStreaming::ws_handler()");
break;
}
info!("Trying to open WebSocket stream");
let socket = {
match web_socket_connect::try_connect(
this_node,
&node.addr_raft,
&raft_type,
tls_config.clone(),
&secret,
)
.await
{
Ok(socket) => {
info!("WebSocket connected successfully");
socket
}
Err(err) => {
error!("Socket connect error to node {}: {:?}", node.id, err);
for _ in 0..3 {
time::sleep(Duration::from_millis(heartbeat_interval)).await;
if let Ok(req) = rx.try_recv() {
let ack = match req {
#[cfg(feature = "sqlite")]
RaftRequest::AppendDB((ack, _)) => Some(ack),
#[cfg(feature = "sqlite")]
RaftRequest::VoteDB((ack, _)) => Some(ack),
#[cfg(feature = "sqlite")]
RaftRequest::SnapshotDB((ack, _)) => Some(ack),
#[cfg(feature = "cache")]
RaftRequest::AppendCache((ack, _)) => Some(ack),
#[cfg(feature = "cache")]
RaftRequest::VoteCache((ack, _)) => Some(ack),
#[cfg(feature = "cache")]
RaftRequest::SnapshotCache((ack, _)) => Some(ack),
RaftRequest::StreamResponse(_) => None,
RaftRequest::ReaderExit => {
continue;
}
RaftRequest::Shutdown => {
break 'outer;
}
};
if let Some(ack) = ack {
let _ = ack.send(Err(Error::Connect(err.to_string())));
}
}
}
continue;
}
}
};
assert!(
in_flight.is_empty(),
"raft in flight buffer should always be empty when restoring a connection"
);
let (tx_write, rx_write) = flume::bounded(1);
let (tx_read, rx_read) = flume::bounded(1);
let (read, write) = socket.split(tokio::io::split);
let read = FragmentCollectorRead::new(read);
let handle_read = task::spawn(Self::stream_reader(read, tx_read.clone()));
let handle_write = task::spawn(Self::stream_writer(write, rx_write));
loop {
let res = select! {
res = rx_read.recv_async() => res,
res = rx.recv_async() => res,
};
let req = match res {
Ok(r) => r,
Err(err) => {
error!("Client stream reader error: {}", err,);
if rx.is_disconnected() {
debug!("Raft tx dropped - exiting Stream Reader");
shutdown = true;
}
if rx_read.is_disconnected() {
debug!("Client Stream reader exited - initiating shutdown + reconnect");
}
break;
}
};
let stream_req = match req {
#[cfg(feature = "sqlite")]
RaftRequest::AppendDB((ack, req)) => {
Some((ack, RaftStreamRequest::AppendDB((request_id, req))))
}
#[cfg(feature = "sqlite")]
RaftRequest::VoteDB((ack, req)) => {
Some((ack, RaftStreamRequest::VoteDB((request_id, req))))
}
#[cfg(feature = "sqlite")]
RaftRequest::SnapshotDB((ack, req)) => {
Some((ack, RaftStreamRequest::SnapshotDB((request_id, req))))
}
#[cfg(feature = "cache")]
RaftRequest::AppendCache((ack, req)) => {
Some((ack, RaftStreamRequest::AppendCache((request_id, req))))
}
#[cfg(feature = "cache")]
RaftRequest::VoteCache((ack, req)) => {
Some((ack, RaftStreamRequest::VoteCache((request_id, req))))
}
#[cfg(feature = "cache")]
RaftRequest::SnapshotCache((ack, req)) => {
Some((ack, RaftStreamRequest::SnapshotCache((request_id, req))))
}
RaftRequest::StreamResponse(resp) => {
match in_flight.remove(&resp.request_id) {
None => {
error!("client ack for RaftStreamResponse missing");
}
Some(ack) => {
if ack.send(Ok(resp.payload)).is_err() {
error!("sending back stream response from raft server");
}
}
}
None
}
RaftRequest::ReaderExit => {
debug!(
"ReaderExit - Client Stream reader exited - initiating shutdown + reconnect"
);
break;
}
RaftRequest::Shutdown => {
debug!("RaftRequest::Shutdown");
shutdown = true;
break;
}
};
if let Some((ack, payload)) = stream_req {
let bytes = serialize(&payload).unwrap();
if let Err(err) = tx_write.send_async(WritePayload::Payload(bytes)).await {
let _ = ack.send(Err(Error::Connect(format!(
"Error sending Write Request to WebSocket writer: {err}"
))));
break;
}
in_flight.insert(request_id, ack);
request_id += 1;
}
}
let _ = tx_write.send_async(WritePayload::Close).await;
for (_, ack) in in_flight.drain() {
let _ = ack.send(Err(Error::Connect("Raft WebSocket stream ended".into())));
}
in_flight = HashMap::with_capacity(4);
time::sleep(Duration::from_millis(250)).await;
handle_write.abort();
handle_read.abort();
if shutdown {
break;
}
}
debug!("Raft Client shut down, tx closed, exiting WsHandler");
}
async fn stream_reader(
mut read: FragmentCollectorRead<ReadHalf<TokioIo<Upgraded>>>,
tx: flume::Sender<RaftRequest>,
) {
while let Ok(frame) = read
.read_frame(&mut |frame| async move {
debug!(
"Received obligated send in stream client: OpCode: {:?}: {:?}",
frame.opcode.clone(),
frame.payload
);
Ok::<(), Error>(())
})
.await
{
match frame.opcode {
OpCode::Continuation => {}
OpCode::Text => {}
OpCode::Binary => {
let bytes = frame.payload.deref();
let payload = deserialize::<RaftStreamResponse>(bytes).unwrap();
if let Err(err) = tx.send_async(RaftRequest::StreamResponse(payload)).await {
error!(
"Error sending Response to Raft client stream manager: {:?}",
err
);
}
}
OpCode::Close => break,
OpCode::Ping => {}
OpCode::Pong => {}
}
}
let _ = tx.send_async(RaftRequest::ReaderExit).await;
debug!("Exiting Client Stream Reader");
}
async fn stream_writer(
mut write: WebSocketWrite<WriteHalf<TokioIo<Upgraded>>>,
rx: flume::Receiver<WritePayload>,
) {
while let Ok(payload) = rx.recv_async().await {
match payload {
WritePayload::Payload(bytes) => {
let frame = Frame::binary(Payload::from(bytes));
if let Err(err) = write.write_frame(frame).await {
error!("Client Stream error: {:?}", err);
break;
}
}
WritePayload::Close => {
debug!("Received Close request in Client Stream Writer");
let _ = write.write_frame(Frame::close(1000, b"go away")).await;
break;
}
}
}
debug!("Exiting Client Stream Writer");
}
}
#[allow(clippy::type_complexity)]
pub struct NetworkConnectionStreaming {
node: Node,
sender: flume::Sender<RaftRequest>,
task: Option<JoinHandle<()>>,
}
impl Drop for NetworkConnectionStreaming {
fn drop(&mut self) {
let _ = self.sender.try_send(RaftRequest::Shutdown);
if let Some(task) = self.task.take() {
task.abort();
}
}
}
impl NetworkConnectionStreaming {
#[inline(always)]
async fn send<Err>(
&mut self,
req: RaftRequest,
rx: oneshot::Receiver<Result<RaftStreamResponsePayload, Error>>,
) -> Result<RaftStreamResponsePayload, RPCError<NodeId, Node, Err>>
where
Err: std::error::Error + 'static + Clone,
{
tracing::debug!(
req = debug(&req),
"sending rpc request to {}",
self.node.addr_raft
);
self.sender.send_async(req).await.map_err(|err| {
error!(
"NetworkConnectionStreaming::send to node {}: {}",
self.node.id,
err.to_string()
);
RPCError::Unreachable(Unreachable::new(&err))
})?;
rx.await
.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))?
.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))
}
}
#[cfg(feature = "sqlite")]
impl RaftNetwork<TypeConfigSqlite> for NetworkConnectionStreaming {
#[tracing::instrument(level = "debug", skip_all, err(Debug))]
async fn append_entries(
&mut self,
req: AppendEntriesRequest<TypeConfigSqlite>,
_option: RPCOption,
) -> Result<AppendEntriesResponse<NodeId>, RPCError<NodeId, Node, RaftError<NodeId>>> {
let (ack, rx) = oneshot::channel();
match self.send(RaftRequest::AppendDB((ack, req)), rx).await? {
RaftStreamResponsePayload::AppendDB(resp) => {
resp.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))
}
_ => unreachable!(),
}
}
#[tracing::instrument(level = "debug", skip_all, err(Debug))]
async fn install_snapshot(
&mut self,
req: InstallSnapshotRequest<TypeConfigSqlite>,
_option: RPCOption,
) -> Result<
InstallSnapshotResponse<NodeId>,
RPCError<NodeId, Node, RaftError<NodeId, InstallSnapshotError>>,
> {
let (ack, rx) = oneshot::channel();
match self.send(RaftRequest::SnapshotDB((ack, req)), rx).await? {
RaftStreamResponsePayload::SnapshotDB(resp) => {
resp.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))
}
_ => unreachable!(),
}
}
#[tracing::instrument(level = "debug", skip_all, err(Debug))]
async fn vote(
&mut self,
req: VoteRequest<NodeId>,
_option: RPCOption,
) -> Result<VoteResponse<NodeId>, RPCError<NodeId, Node, RaftError<NodeId>>> {
let (ack, rx) = oneshot::channel();
match self.send(RaftRequest::VoteDB((ack, req)), rx).await? {
RaftStreamResponsePayload::VoteDB(resp) => {
resp.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))
}
_ => unreachable!(),
}
}
}
#[cfg(feature = "cache")]
impl RaftNetwork<TypeConfigKV> for NetworkConnectionStreaming {
#[tracing::instrument(level = "debug", skip_all, err(Debug))]
async fn append_entries(
&mut self,
req: AppendEntriesRequest<TypeConfigKV>,
_option: RPCOption,
) -> Result<AppendEntriesResponse<NodeId>, RPCError<NodeId, Node, RaftError<NodeId>>> {
let (ack, rx) = oneshot::channel();
match self.send(RaftRequest::AppendCache((ack, req)), rx).await? {
RaftStreamResponsePayload::AppendCache(resp) => {
resp.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))
}
_ => unreachable!(),
}
}
#[tracing::instrument(level = "debug", skip_all, err(Debug))]
async fn install_snapshot(
&mut self,
req: InstallSnapshotRequest<TypeConfigKV>,
_option: RPCOption,
) -> Result<
InstallSnapshotResponse<NodeId>,
RPCError<NodeId, Node, RaftError<NodeId, InstallSnapshotError>>,
> {
let (ack, rx) = oneshot::channel();
match self
.send(RaftRequest::SnapshotCache((ack, req)), rx)
.await?
{
RaftStreamResponsePayload::SnapshotCache(resp) => {
resp.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))
}
_ => unreachable!(),
}
}
#[tracing::instrument(level = "debug", skip_all, err(Debug))]
async fn vote(
&mut self,
req: VoteRequest<NodeId>,
_option: RPCOption,
) -> Result<VoteResponse<NodeId>, RPCError<NodeId, Node, RaftError<NodeId>>> {
let (ack, rx) = oneshot::channel();
match self.send(RaftRequest::VoteCache((ack, req)), rx).await? {
RaftStreamResponsePayload::VoteCache(resp) => {
resp.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))
}
_ => unreachable!(),
}
}
}