use std::{
collections::{hash_map, HashMap},
sync::Arc,
time::Duration,
};
use message_encoding::MessageEncoding;
use tokio::sync::{
mpsc::{error::TrySendError, Receiver, Sender},
oneshot, Mutex,
};
use tokio_util::sync::CancellationToken;
use crate::{
cluster::node_state::NodeState,
protocol::messages::{LeaderInfo, LeaderState, SharePeerDetails, SyncRequest, SyncResponse, PROTOCOL_VERSION},
state::deterministic_state::DeterministicState,
transport::{
channels::NetIoSettings,
traits::{SyncIO, SyncIOAddress},
},
};
const RPC_QUEUE_CAPACITY: usize = 512;
const CONNECT_RETRY_LIMIT: u64 = 3;
const CONNECT_RETRY_DELAY: Duration = Duration::from_millis(100);
const CONNECT_FAIL_HOLD_DELAY: Duration = Duration::from_secs(5);
type ConnectionMap<I, D> = HashMap<<I as SyncIO>::Address, Connection<<I as SyncIO>::Address, D>>;
pub struct PeerConnections<I: SyncIO, D: DeterministicState> {
io: Arc<I>,
conn_settings: NetIoSettings,
state: Arc<NodeState<I::Address, D>>,
connections: Mutex<ConnectionMap<I, D>>,
}
struct Connection<A: SyncIOAddress, D: DeterministicState> {
tx: Sender<RpcMessage<A, D>>,
cancel: CancellationToken,
}
struct RpcMessage<A: SyncIOAddress, D: DeterministicState> {
request: SyncRequest<A, D>,
response: RpcResponseHandler<A, D>,
}
enum RpcResponseHandler<A: SyncIOAddress, D: DeterministicState> {
Sender(oneshot::Sender<Result<SyncResponse<A, D>, PeerRpcError>>),
ForwardedAction { target: A, source: A },
}
impl<A: SyncIOAddress, D: DeterministicState> RpcResponseHandler<A, D> {
fn handle(self, response: Result<SyncResponse<A, D>, PeerRpcError>) {
match self {
Self::Sender(tx) => {
let _ = tx.send(response);
}
Self::ForwardedAction { target, source } => match response {
Ok(SyncResponse::Ok) => {}
Ok(response) => {
tracing::warn!(?target, ?source, response = response.name(), "sync target rejected forwarded action");
}
Err(error) => {
tracing::warn!(?target, ?source, ?error, "failed to forward action to sync target");
}
},
}
}
}
impl<I, D> PeerConnections<I, D>
where
I: SyncIO,
D: DeterministicState + MessageEncoding,
D::Action: MessageEncoding,
D::AuthorityAction: MessageEncoding,
{
pub fn new(io: Arc<I>, conn_settings: NetIoSettings, state: Arc<NodeState<I::Address, D>>) -> Self {
Self {
io,
conn_settings,
state,
connections: Mutex::new(HashMap::new()),
}
}
pub async fn send_rpc(
&self,
peer: I::Address,
request: SyncRequest<I::Address, D>,
) -> Result<SyncResponse<I::Address, D>, PeerRpcError> {
let rx = self.enqueue_rpc(peer, request).await?;
await_rpc_response(rx).await
}
pub async fn enqueue_rpc(
&self,
peer: I::Address,
request: SyncRequest<I::Address, D>,
) -> Result<oneshot::Receiver<Result<SyncResponse<I::Address, D>, PeerRpcError>>, PeerRpcError> {
if matches!(&request, SyncRequest::SubscribeFresh | SyncRequest::SubscribeRecovery(_)) {
return Err(PeerRpcError::RequestNotAllowedOverRpc);
}
let (tx, rx) = oneshot::channel();
self.enqueue_message(
peer,
RpcMessage {
request,
response: RpcResponseHandler::Sender(tx),
},
)
.await?;
Ok(rx)
}
pub async fn enqueue_forwarded_action(
&self,
peer: I::Address,
source: I::Address,
action: D::Action,
) -> Result<(), PeerRpcError> {
self.enqueue_message(
peer,
RpcMessage {
request: SyncRequest::Action { source, action },
response: RpcResponseHandler::ForwardedAction { target: peer, source },
},
)
.await
}
async fn enqueue_message(
&self,
peer: I::Address,
msg: RpcMessage<I::Address, D>,
) -> Result<(), PeerRpcError> {
let mut pending = Some(msg);
loop {
let msg = pending.take().expect("rpc message missing while send is still pending");
let mut connections = self.connections.lock().await;
match connections.entry(peer) {
hash_map::Entry::Occupied(entry) => {
if entry.get().cancel.is_cancelled() {
let connection = entry.remove();
drop(connections);
connection.kill().await;
pending = Some(msg);
} else {
match entry.get().tx.try_send(msg) {
Ok(()) => {
drop(connections);
return Ok(());
}
Err(TrySendError::Closed(msg)) => {
entry.remove();
pending = Some(msg);
}
Err(TrySendError::Full(msg)) => {
let sender = entry.get().tx.clone();
drop(connections);
match sender.send(msg).await {
Ok(()) => return Ok(()),
Err(error) => {
pending = Some(error.0);
let mut connections = self.connections.lock().await;
if connections
.get(&peer)
.is_some_and(|conn| conn.tx.is_closed() || conn.cancel.is_cancelled())
{
connections.remove(&peer);
}
}
}
}
}
}
}
hash_map::Entry::Vacant(entry) => {
let conn = Connection::create(
peer,
self.state.my_address,
self.state.clone(),
self.io.clone(),
self.conn_settings.clone(),
);
entry.insert(conn);
pending = Some(msg);
}
}
}
}
pub async fn kill_connection(&self, peer: I::Address) {
let connection = self.connections.lock().await.remove(&peer);
if let Some(connection) = connection {
connection.kill().await;
}
}
pub async fn send_peers_info(
&self,
peer: I::Address,
peers: Vec<SharePeerDetails<I::Address>>,
) -> Result<Vec<SharePeerDetails<I::Address>>, PeerRpcError> {
let response = self.send_rpc(peer, SyncRequest::SharePeers(peers)).await?;
match response {
SyncResponse::Peers(peers) => Ok(peers),
response => self.unexpected_response(peer, "Peers", response).await,
}
}
pub async fn send_leader_info(&self, peer: I::Address, info: LeaderInfo<I::Address>) -> Result<(), PeerRpcError> {
let response = self.send_rpc(peer, SyncRequest::LeaderInformation(info)).await?;
match response {
SyncResponse::Ok => Ok(()),
response => self.unexpected_response(peer, "Ok", response).await,
}
}
pub async fn query_leader(&self, peer: I::Address) -> Result<LeaderState<I::Address>, PeerRpcError> {
let response = self.send_rpc(peer, SyncRequest::LeaderQuery).await?;
match response {
SyncResponse::LeaderState(state) => Ok(state),
response => self.unexpected_response(peer, "LeaderState", response).await,
}
}
async fn unexpected_response<T>(
&self,
peer: I::Address,
expected: &'static str,
response: SyncResponse<I::Address, D>,
) -> Result<T, PeerRpcError> {
let actual = response.name();
tracing::debug!(?peer, expected, actual, "peer returned unexpected rpc response");
self.kill_connection(peer).await;
Err(PeerRpcError::UnexpectedResponse { expected, actual })
}
}
async fn await_rpc_response<A: SyncIOAddress, D: DeterministicState>(
rx: oneshot::Receiver<Result<SyncResponse<A, D>, PeerRpcError>>,
) -> Result<SyncResponse<A, D>, PeerRpcError> {
match rx.await {
Ok(res) => res,
Err(_) => Err(PeerRpcError::ResponseDropped),
}
}
impl<A, D> Connection<A, D>
where
A: SyncIOAddress,
D: DeterministicState + MessageEncoding,
D::Action: MessageEncoding,
D::AuthorityAction: MessageEncoding,
{
pub fn create<I>(
remote_addr: A,
local_addr: A,
state: Arc<NodeState<A, D>>,
io: Arc<I>,
settings: NetIoSettings,
) -> Self
where
I: SyncIO<Address = A>,
{
let (tx, rx) = tokio::sync::mpsc::channel(RPC_QUEUE_CAPACITY);
let cancel = CancellationToken::new();
tokio::spawn(
ConnectionWorker {
rx,
cancel: cancel.clone(),
remote_addr,
local_addr,
state,
}
.run(io, settings),
);
Self { tx, cancel }
}
async fn kill(self) {
self.cancel.cancel();
self.tx.closed().await;
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PeerRpcError {
RequestNotAllowedOverRpc,
FailedToConnectToPeer,
FailedToSendRequest,
FailedToReceiveResponse,
HandshakeRejected,
ResponseTimedOut,
ResponseDropped,
UnexpectedResponse {
expected: &'static str,
actual: &'static str,
},
}
struct ConnectionWorker<A: SyncIOAddress, D: DeterministicState> {
rx: Receiver<RpcMessage<A, D>>,
cancel: CancellationToken,
remote_addr: A,
local_addr: A,
state: Arc<NodeState<A, D>>,
}
impl<A: SyncIOAddress, D: DeterministicState> ConnectionWorker<A, D>
where
D: MessageEncoding,
D::Action: MessageEncoding,
D::AuthorityAction: MessageEncoding,
{
async fn run<I: SyncIO<Address = A>>(mut self, io: Arc<I>, settings: NetIoSettings) {
let mut repeat_failures = 0;
let mut retry_wait = CONNECT_RETRY_DELAY;
let connection = loop {
let connect_result = match tokio::time::timeout(settings.message_timeout, io.connect(&self.remote_addr)).await
{
Ok(result) => result,
Err(_) => Err(std::io::Error::new(
std::io::ErrorKind::TimedOut,
"timed out connecting to peer",
)),
};
match connect_result {
Ok(connection) => break connection,
Err(error) => {
tracing::error!(?error, repeat_failures, "failed to connect to peer {:?}", self.remote_addr);
repeat_failures += 1;
if repeat_failures > CONNECT_RETRY_LIMIT {
self.state.mark_peer_failed_to_connect(self.remote_addr).await;
self.drain_with_error(PeerRpcError::FailedToConnectToPeer).await;
tokio::time::sleep(CONNECT_FAIL_HOLD_DELAY).await;
return;
}
tokio::time::sleep(retry_wait).await;
retry_wait += CONNECT_RETRY_DELAY;
}
}
};
let (_remote, write, mut read) = connection.client_channels::<D>(settings.clone());
if let Err(error) = self.handshake(&write, &mut read, settings.message_timeout).await {
self.state.mark_peer_not_connected(self.remote_addr).await;
self.drain_with_error(error).await;
return;
}
self.state.mark_peer_connected(self.remote_addr).await;
let (response_tx, response_rx) = tokio::sync::mpsc::channel(RPC_QUEUE_CAPACITY);
tokio::spawn(
ResponseReader::<A, D> {
responses: read,
response_handlers: response_rx,
timeout: settings.message_timeout,
cancel: self.cancel.clone(),
}
.run(),
);
loop {
tokio::select! {
biased;
_ = self.cancel.cancelled() => {
self.drain_with_error(PeerRpcError::FailedToReceiveResponse).await;
break;
}
msg = self.rx.recv() => {
let Some(msg) = msg else {
break;
};
if write.send(msg.request).await.is_err() {
msg.response.handle(Err(PeerRpcError::FailedToSendRequest));
self.drain_with_error(PeerRpcError::FailedToSendRequest).await;
break;
}
if let Err(error) = response_tx.send(msg.response).await {
error.0.handle(Err(PeerRpcError::FailedToReceiveResponse));
self.drain_with_error(PeerRpcError::FailedToReceiveResponse).await;
break;
}
}
}
}
self.state.mark_peer_not_connected(self.remote_addr).await;
self.cancel.cancel();
}
async fn handshake(
&self,
write: &Sender<SyncRequest<A, D>>,
read: &mut Receiver<SyncResponse<A, D>>,
timeout: Duration,
) -> Result<(), PeerRpcError> {
write
.send(SyncRequest::ProtocolVersion(PROTOCOL_VERSION))
.await
.map_err(|_| PeerRpcError::FailedToSendRequest)?;
require_ok(read_response(read, timeout).await?)?;
write
.send(SyncRequest::MyAddress(self.local_addr))
.await
.map_err(|_| PeerRpcError::FailedToSendRequest)?;
require_ok(read_response(read, timeout).await?)?;
Ok(())
}
async fn drain_with_error(&mut self, error: PeerRpcError) {
self.rx.close();
while let Some(msg) = self.rx.recv().await {
msg.response.handle(Err(error.clone()));
}
}
}
struct ResponseReader<A: SyncIOAddress, D: DeterministicState> {
responses: Receiver<SyncResponse<A, D>>,
response_handlers: Receiver<RpcResponseHandler<A, D>>,
timeout: Duration,
cancel: CancellationToken,
}
impl<A: SyncIOAddress, D: DeterministicState> ResponseReader<A, D> {
async fn run(mut self) {
while let Some(response_handler) = self.response_handlers.recv().await {
let response = read_response(&mut self.responses, self.timeout).await;
let fatal_error = response.as_ref().err().cloned();
response_handler.handle(response);
if let Some(error) = fatal_error {
self.cancel.cancel();
self.drain_with_error(error).await;
return;
}
}
}
async fn drain_with_error(&mut self, error: PeerRpcError) {
self.response_handlers.close();
while let Some(response_handler) = self.response_handlers.recv().await {
response_handler.handle(Err(error.clone()));
}
}
}
async fn read_response<A: SyncIOAddress, D: DeterministicState>(
read: &mut Receiver<SyncResponse<A, D>>,
timeout: Duration,
) -> Result<SyncResponse<A, D>, PeerRpcError> {
match tokio::time::timeout(timeout, read.recv()).await {
Ok(Some(response)) => Ok(response),
Ok(None) => Err(PeerRpcError::FailedToReceiveResponse),
Err(_) => Err(PeerRpcError::ResponseTimedOut),
}
}
fn require_ok<A: SyncIOAddress, D: DeterministicState>(response: SyncResponse<A, D>) -> Result<(), PeerRpcError> {
match response {
SyncResponse::Ok => Ok(()),
_ => Err(PeerRpcError::HandshakeRejected),
}
}
#[cfg(test)]
mod tests {
use std::{
collections::HashMap,
sync::{
atomic::{AtomicUsize, Ordering},
Arc,
},
time::Duration,
};
use message_encoding::MessageEncoding;
use sequenced_broadcast::SequencedBroadcastSettings;
use tokio::{
io::{duplex, split, DuplexStream, ReadHalf, WriteHalf},
sync::{Mutex, Notify, Semaphore},
};
use super::*;
use crate::{
protocol::messages::{ElectionTerm, LeaderMode},
state::{recoverable_state::RecoverableState, subscribable_state::SubscribableState},
transport::traits::SyncConnection,
};
#[derive(Clone, Debug)]
struct TestState;
impl DeterministicState for TestState {
type Action = u64;
type AuthorityAction = u64;
fn accept_seq(&self) -> u64 {
0
}
fn authority(&self, action: Self::Action) -> Self::AuthorityAction {
action
}
fn update(&mut self, _action: &Self::AuthorityAction) {}
}
impl MessageEncoding for TestState {
fn write_to<T: std::io::Write>(&self, _out: &mut T) -> std::io::Result<usize> {
Ok(0)
}
fn read_from<T: std::io::Read>(_read: &mut T) -> std::io::Result<Self> {
Ok(Self)
}
}
struct DelayedResponseIo {
actions_seen: Arc<AtomicUsize>,
second_action_seen: Arc<Notify>,
release_responses: Arc<Semaphore>,
}
impl DelayedResponseIo {
fn new() -> Self {
Self {
actions_seen: Arc::new(AtomicUsize::new(0)),
second_action_seen: Arc::new(Notify::new()),
release_responses: Arc::new(Semaphore::new(0)),
}
}
}
impl SyncIO for DelayedResponseIo {
type Address = u64;
type Read = ReadHalf<DuplexStream>;
type Write = WriteHalf<DuplexStream>;
async fn connect(&self, remote: &u64) -> std::io::Result<SyncConnection<Self>> {
let (client, server) = duplex(64 * 1024);
let (client_read, client_write) = split(client);
let (server_read, server_write) = split(server);
let (_, write, read) = SyncConnection::<Self> {
remote: *remote,
read: server_read,
write: server_write,
}
.server_channels::<TestState>(test_settings());
tokio::spawn(serve_delayed_responses(
write,
read,
self.actions_seen.clone(),
self.second_action_seen.clone(),
self.release_responses.clone(),
));
Ok(SyncConnection {
remote: *remote,
read: client_read,
write: client_write,
})
}
}
async fn serve_delayed_responses(
write: Sender<SyncResponse<u64, TestState>>,
mut read: Receiver<SyncRequest<u64, TestState>>,
actions_seen: Arc<AtomicUsize>,
second_action_seen: Arc<Notify>,
release_responses: Arc<Semaphore>,
) {
while let Some(request) = read.recv().await {
match request {
SyncRequest::ProtocolVersion(_) | SyncRequest::MyAddress(_) => {
if write.send(SyncResponse::Ok).await.is_err() {
return;
}
}
SyncRequest::Action { .. } => {
if actions_seen.fetch_add(1, Ordering::SeqCst) == 1 {
second_action_seen.notify_waiters();
}
let write = write.clone();
let release_responses = release_responses.clone();
tokio::spawn(async move {
let _permit = release_responses.acquire().await;
let _ = write.send(SyncResponse::Ok).await;
});
}
_ => {
let _ = write.send(SyncResponse::UnexpectedRequest).await;
return;
}
}
}
}
fn test_settings() -> NetIoSettings {
NetIoSettings {
process_timeout: Duration::from_millis(100),
message_timeout: Duration::from_secs(5),
}
}
fn node_state() -> Arc<NodeState<u64, TestState>> {
Arc::new(NodeState {
my_address: 1,
can_lead: true,
peers: Mutex::new(HashMap::new()),
state: SubscribableState::new(RecoverableState::new(1, TestState), SequencedBroadcastSettings::default())
.unwrap(),
leader_state: Mutex::new(LeaderState {
term: ElectionTerm::from_term(0),
mode: LeaderMode::Following { leader: 2 },
}),
})
}
#[tokio::test]
async fn enqueue_rpc_does_not_wait_for_previous_response() {
let io = Arc::new(DelayedResponseIo::new());
let connections = PeerConnections::new(io.clone(), test_settings(), node_state());
let first = connections
.enqueue_rpc(2, SyncRequest::Action { source: 1, action: 10 })
.await
.unwrap();
let second = connections
.enqueue_rpc(2, SyncRequest::Action { source: 1, action: 20 })
.await
.unwrap();
tokio::time::timeout(Duration::from_secs(1), io.second_action_seen.notified())
.await
.expect("second action should be sent before the first response arrives");
io.release_responses.add_permits(2);
assert!(matches!(first.await.unwrap().unwrap(), SyncResponse::Ok));
assert!(matches!(second.await.unwrap().unwrap(), SyncResponse::Ok));
}
}