use std::{io, task::Poll};
use futures::StreamExt;
use log::*;
use tari_shutdown::ShutdownSignal;
use tokio::{
io::{AsyncRead, AsyncWrite},
sync::{broadcast, mpsc, oneshot},
};
#[cfg(feature = "metrics")]
use super::metrics;
use super::{MessagingEvent, MessagingProtocol};
use crate::{PeerConnection, message::InboundMessage, peer_manager::NodeId};
const LOG_TARGET: &str = "comms::protocol::messaging::inbound";
pub struct InboundMessaging {
connection: PeerConnection,
inbound_message_tx: mpsc::Sender<InboundMessage>,
messaging_events_tx: broadcast::Sender<MessagingEvent>,
internal_events_tx: mpsc::Sender<MessagingEvent>,
session_id: u64,
enable_message_received_event: bool,
shutdown_signal: ShutdownSignal,
replaced: oneshot::Receiver<()>,
}
impl InboundMessaging {
#[allow(clippy::too_many_arguments)]
pub fn new(
connection: PeerConnection,
inbound_message_tx: mpsc::Sender<InboundMessage>,
messaging_events_tx: broadcast::Sender<MessagingEvent>,
internal_events_tx: mpsc::Sender<MessagingEvent>,
session_id: u64,
enable_message_received_event: bool,
shutdown_signal: ShutdownSignal,
replaced: oneshot::Receiver<()>,
) -> Self {
Self {
connection,
inbound_message_tx,
messaging_events_tx,
internal_events_tx,
session_id,
enable_message_received_event,
shutdown_signal,
replaced,
}
}
pub async fn run<S>(mut self, socket: S)
where S: AsyncRead + AsyncWrite + Unpin {
let peer = self.connection.peer_node_id().clone();
#[cfg(feature = "metrics")]
metrics::num_sessions().inc();
debug!(
target: LOG_TARGET,
"Starting inbound messaging protocol for peer '{}'",
peer.short_str()
);
let mut stream = MessagingProtocol::framed(socket);
let on_disconnect = self.connection.on_disconnect();
tokio::pin!(on_disconnect);
let mut disconnected = false;
loop {
let maybe_result = if disconnected {
match futures::poll!(stream.next()) {
Poll::Ready(item) => item,
Poll::Pending => None,
}
} else {
tokio::select! {
biased;
item = stream.next() => item,
_ = &mut on_disconnect => {
disconnected = true;
match futures::poll!(stream.next()) {
Poll::Ready(item) => item,
Poll::Pending => None,
}
},
_ = &mut self.replaced => {
debug!(
target: LOG_TARGET,
"Inbound messaging session for peer '{}' replaced by a newer substream", peer.short_str()
);
disconnected = true;
match futures::poll!(stream.next()) {
Poll::Ready(item) => item,
Poll::Pending => None,
}
},
_ = self.shutdown_signal.wait() => None,
}
};
let Some(result) = maybe_result else { break };
if !self.handle_frame_result(result, &peer).await {
break;
}
}
drop(stream);
let _ignore = self
.messaging_events_tx
.send(MessagingEvent::InboundProtocolExited(peer.clone()));
let _ignore = self
.internal_events_tx
.send(MessagingEvent::InboundSessionExited(peer.clone(), self.session_id))
.await;
#[cfg(feature = "metrics")]
metrics::num_sessions().dec();
debug!(
target: LOG_TARGET,
"Inbound messaging handler exited for peer `{}`",
peer.short_str()
);
}
async fn handle_frame_result(&mut self, result: io::Result<bytes::BytesMut>, peer: &NodeId) -> bool {
match result {
Ok(raw_msg) => {
#[cfg(feature = "metrics")]
metrics::inbound_message_count().inc();
let msg_len = raw_msg.len();
let inbound_msg = InboundMessage::new(peer.clone(), raw_msg.freeze());
debug!(
target: LOG_TARGET,
"Received message {} from peer '{}' ({} bytes)",
inbound_msg.tag,
peer.short_str(),
msg_len
);
let message_tag = inbound_msg.tag;
if self.inbound_message_tx.send(inbound_msg).await.is_err() {
warn!(
target: LOG_TARGET,
"Failed to send InboundMessage {} for peer '{}' because inbound message channel closed",
message_tag,
peer.short_str(),
);
return false;
}
if self.enable_message_received_event {
let _result = self
.messaging_events_tx
.send(MessagingEvent::MessageReceived(peer.clone(), message_tag));
}
true
},
Err(err) if err.kind() == io::ErrorKind::InvalidData => {
#[cfg(feature = "metrics")]
metrics::error_count().inc();
debug!(
target: LOG_TARGET,
"Failed to receive from peer '{}' because '{}'",
peer.short_str(),
err
);
let _result = self.messaging_events_tx.send(MessagingEvent::ProtocolViolation {
peer_node_id: peer.clone(),
details: err.to_string(),
});
false
},
Err(err) => {
#[cfg(feature = "metrics")]
metrics::error_count().inc();
error!(
target: LOG_TARGET,
"Failed to receive from peer '{}' because '{}'",
peer.short_str(),
err
);
false
},
}
}
}