use std::{
collections::{HashMap, hash_map::Entry},
fmt,
fmt::Display,
sync::Arc,
time::Duration,
};
use log::*;
use tari_shutdown::{
ShutdownSignal,
oneshot_trigger::{OneshotSignal, OneshotTrigger},
};
use thiserror::Error;
use tokio::{
io::{AsyncRead, AsyncWrite},
sync::{Semaphore, broadcast, mpsc, oneshot},
task::JoinHandle,
time,
};
use tokio_util::codec::{Framed, LengthDelimitedCodec};
use super::error::MessagingProtocolError;
use crate::{
PeerConnection,
RefKind,
connectivity::ConnectivityRequester,
framing,
message::{InboundMessage, MessageTag, OutboundMessage},
multiplexing::Substream,
peer_manager::NodeId,
protocol::{
ProtocolEvent,
ProtocolId,
ProtocolNotification,
messaging::{inbound::InboundMessaging, outbound::OutboundMessaging},
},
};
const LOG_TARGET: &str = "comms::protocol::messaging";
const INTERNAL_MESSAGING_EVENT_CHANNEL_SIZE: usize = 10;
const MAX_FRAME_LENGTH: usize = 8 * 1_024 * 1_024;
const CONNECTION_LOOKUP_RETRY_INTERVAL: Duration = Duration::from_millis(20);
const CONNECTION_LOOKUP_TIMEOUT: Duration = Duration::from_secs(2);
const MAX_PENDING_SUBSTREAM_RESOLUTIONS: usize = 128;
const REPLACEMENT_RATE_WINDOW: Duration = Duration::from_secs(1);
const MAX_REPLACEMENTS_PER_WINDOW: u32 = 10;
const SHED_LOG_INTERVAL: Duration = Duration::from_secs(2);
const MAX_STOPPING_SESSIONS_PER_PEER: usize = 4;
pub type MessagingEventSender = broadcast::Sender<MessagingEvent>;
pub type MessagingEventReceiver = broadcast::Receiver<MessagingEvent>;
#[derive(Debug, Error, Copy, Clone)]
pub enum SendFailReason {
#[error("Dial was attempted, but failed")]
PeerDialFailed,
#[error("Failed to open a messaging substream to peer")]
SubstreamOpenFailed,
#[error("Failed to send on substream channel")]
SubstreamSendFailed,
#[error("Message was dropped before sending")]
Dropped,
#[error("Message could not send after {0} attempt(s)")]
MaxRetriesReached(usize),
}
#[derive(Debug, Clone)]
pub enum MessagingEvent {
MessageReceived(NodeId, MessageTag),
OutboundProtocolExited(NodeId),
InboundProtocolExited(NodeId),
InboundSessionExited(NodeId, u64),
ProtocolViolation {
peer_node_id: NodeId,
details: String,
},
}
impl fmt::Display for MessagingEvent {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
use MessagingEvent::*;
match self {
MessageReceived(node_id, tag) => write!(f, "MessageReceived({node_id}, {tag})"),
OutboundProtocolExited(node_id) => write!(f, "OutboundProtocolExited({node_id})"),
InboundProtocolExited(node_id) => write!(f, "InboundProtocolExited({node_id})"),
InboundSessionExited(node_id, session_id) => {
write!(f, "InboundSessionExited({node_id}, {session_id})")
},
ProtocolViolation { peer_node_id, details } => {
write!(f, "ProtocolViolation({peer_node_id}, {details})")
},
}
}
}
struct ActiveInboundSession {
id: u64,
handle: JoinHandle<()>,
stop_tx: oneshot::Sender<()>,
}
#[derive(Default)]
struct PeerInboundSessions {
current: Option<ActiveInboundSession>,
stopping: Vec<(u64, JoinHandle<()>)>,
}
impl PeerInboundSessions {
fn reap_finished(&mut self) {
self.stopping.retain(|(_, handle)| !handle.is_finished());
}
}
#[derive(Default)]
struct ReplacementBudget {
window_start: Option<time::Instant>,
count: u32,
last_shed_log: Option<time::Instant>,
}
impl ReplacementBudget {
fn try_consume(&mut self) -> bool {
let now = time::Instant::now();
let within_window = self
.window_start
.is_some_and(|start| now.saturating_duration_since(start) < REPLACEMENT_RATE_WINDOW);
if within_window {
if self.count >= MAX_REPLACEMENTS_PER_WINDOW {
return false;
}
self.count = self.count.saturating_add(1);
} else {
self.window_start = Some(now);
self.count = 1;
}
true
}
fn should_warn_on_shed(&mut self) -> bool {
let now = time::Instant::now();
let should = match self.last_shed_log {
Some(last) => now.saturating_duration_since(last) >= SHED_LOG_INTERVAL,
None => true,
};
if should {
self.last_shed_log = Some(now);
}
should
}
}
pub struct MessagingProtocol {
protocol_id: ProtocolId,
connectivity: ConnectivityRequester,
proto_notification: mpsc::Receiver<ProtocolNotification<Substream>>,
active_queues: HashMap<NodeId, mpsc::UnboundedSender<OutboundMessage>>,
active_inbound: HashMap<NodeId, PeerInboundSessions>,
replacement_budgets: HashMap<NodeId, ReplacementBudget>,
next_inbound_session_id: u64,
outbound_message_rx: mpsc::UnboundedReceiver<OutboundMessage>,
messaging_events_tx: MessagingEventSender,
enable_message_received_event: bool,
ban_duration: Option<Duration>,
inbound_message_tx: mpsc::Sender<InboundMessage>,
internal_messaging_event_tx: mpsc::Sender<MessagingEvent>,
internal_messaging_event_rx: mpsc::Receiver<MessagingEvent>,
retry_queue_tx: mpsc::UnboundedSender<OutboundMessage>,
retry_queue_rx: mpsc::UnboundedReceiver<OutboundMessage>,
resolved_substream_tx: mpsc::Sender<(PeerConnection, Substream)>,
resolved_substream_rx: mpsc::Receiver<(PeerConnection, Substream)>,
pending_resolution_permits: Arc<Semaphore>,
last_pending_resolution_shed_log: Option<time::Instant>,
shutdown_signal: ShutdownSignal,
complete_trigger: OneshotTrigger<()>,
}
impl MessagingProtocol {
pub(super) fn new(
protocol_id: ProtocolId,
connectivity: ConnectivityRequester,
proto_notification: mpsc::Receiver<ProtocolNotification<Substream>>,
outbound_message_rx: mpsc::UnboundedReceiver<OutboundMessage>,
messaging_events_tx: MessagingEventSender,
inbound_message_tx: mpsc::Sender<InboundMessage>,
shutdown_signal: ShutdownSignal,
) -> Self {
let (internal_messaging_event_tx, internal_messaging_event_rx) =
mpsc::channel(INTERNAL_MESSAGING_EVENT_CHANNEL_SIZE);
let (retry_queue_tx, retry_queue_rx) = mpsc::unbounded_channel();
let (resolved_substream_tx, resolved_substream_rx) = mpsc::channel(MAX_PENDING_SUBSTREAM_RESOLUTIONS);
Self {
protocol_id,
connectivity,
proto_notification,
outbound_message_rx,
active_inbound: Default::default(),
active_queues: Default::default(),
replacement_budgets: Default::default(),
next_inbound_session_id: 0,
messaging_events_tx,
enable_message_received_event: false,
internal_messaging_event_rx,
internal_messaging_event_tx,
ban_duration: None,
retry_queue_tx,
retry_queue_rx,
inbound_message_tx,
resolved_substream_tx,
resolved_substream_rx,
pending_resolution_permits: Arc::new(Semaphore::new(MAX_PENDING_SUBSTREAM_RESOLUTIONS)),
last_pending_resolution_shed_log: None,
shutdown_signal,
complete_trigger: OneshotTrigger::new(),
}
}
pub fn set_message_received_event_enabled(mut self, enabled: bool) -> Self {
self.enable_message_received_event = enabled;
self
}
pub fn with_ban_duration(mut self, ban_duration: Duration) -> Self {
self.ban_duration = Some(ban_duration);
self
}
pub fn complete_signal(&self) -> OneshotSignal<()> {
self.complete_trigger.to_signal()
}
pub async fn run(mut self) {
let mut shutdown_signal = self.shutdown_signal.clone();
loop {
tokio::select! {
Some(event) = self.internal_messaging_event_rx.recv() => {
self.handle_internal_messaging_event(event).await;
},
Some(msg) = self.retry_queue_rx.recv() => {
if let Err(err) = self.handle_retry_queue_messages(msg) {
error!(
target: LOG_TARGET,
"Failed to retry outbound message because '{err}'"
);
}
},
Some(msg) = self.outbound_message_rx.recv() => {
if let Err(err) = self.send_message(msg) {
error!(
target: LOG_TARGET,
"Failed to handle request because '{err}'"
);
}
},
Some(notification) = self.proto_notification.recv() => {
if let Err(err) = self.handle_protocol_notification(notification).await {
error!(target: LOG_TARGET, "handle_protocol_notification failed: {err}");
}
},
Some((conn, substream)) = self.resolved_substream_rx.recv() => {
self.spawn_inbound_handler(conn, substream);
},
_ = &mut shutdown_signal => {
info!(target: LOG_TARGET, "MessagingProtocol is shutting down because the shutdown signal was triggered");
break;
}
}
}
}
#[inline]
pub(super) fn framed<TSubstream>(socket: TSubstream) -> Framed<TSubstream, LengthDelimitedCodec>
where TSubstream: AsyncRead + AsyncWrite + Unpin {
framing::canonical(socket, MAX_FRAME_LENGTH)
}
async fn handle_internal_messaging_event(&mut self, event: MessagingEvent) {
use MessagingEvent::*;
trace!(target: LOG_TARGET, "Internal messaging event '{event}'" );
match &event {
OutboundProtocolExited(node_id) => {
debug!(
target: LOG_TARGET,
"Outbound protocol handler exited for peer `{}`",
node_id.short_str()
);
if self.active_queues.remove(node_id).is_none() {
debug!(
target: LOG_TARGET,
"OutboundProtocolExited event, but MessagingProtocol has no record of the outbound protocol \
for peer `{}`",
node_id.short_str()
);
}
},
InboundProtocolExited(_) => {},
InboundSessionExited(node_id, session_id) => {
self.prune_inbound_session(node_id, *session_id);
return;
},
ProtocolViolation { peer_node_id, details } => {
self.ban_peer(peer_node_id.clone(), details.to_string()).await;
},
_ => {},
}
let _result = self.messaging_events_tx.send(event);
}
fn prune_inbound_session(&mut self, node_id: &NodeId, session_id: u64) {
let Entry::Occupied(mut entry) = self.active_inbound.entry(node_id.clone()) else {
return;
};
let sessions = entry.get_mut();
if sessions
.current
.as_ref()
.is_some_and(|current| current.id == session_id)
{
sessions.current = None;
} else {
sessions.stopping.retain(|(id, _)| *id != session_id);
}
if sessions.current.is_none() && sessions.stopping.is_empty() {
entry.remove();
self.replacement_budgets.remove(node_id);
}
}
fn handle_retry_queue_messages(&mut self, msg: OutboundMessage) -> Result<(), MessagingProtocolError> {
debug!(target: LOG_TARGET, "Retrying outbound message ({msg})");
self.send_message(msg)?;
Ok(())
}
fn send_message(&mut self, out_msg: OutboundMessage) -> Result<(), MessagingProtocolError> {
trace!(target: LOG_TARGET, "Received request to send message ({out_msg})");
let peer_node_id = out_msg.peer_node_id.clone();
let sender = loop {
match self.active_queues.entry(peer_node_id.clone()) {
Entry::Occupied(entry) => {
if entry.get().is_closed() {
entry.remove();
continue;
}
break entry.into_mut();
},
Entry::Vacant(entry) => {
let sender = Self::spawn_outbound_handler(
self.connectivity.clone(),
self.internal_messaging_event_tx.clone(),
peer_node_id,
self.retry_queue_tx.clone(),
self.protocol_id.clone(),
);
break entry.insert(sender);
},
}
};
trace!(target: LOG_TARGET, "Sending message {out_msg}");
let tag = out_msg.tag;
match sender.send(out_msg) {
Ok(_) => {
trace!(target: LOG_TARGET, "Message ({tag}) dispatched to outbound handler");
Ok(())
},
Err(err) => {
debug!(
target: LOG_TARGET,
"Failed to send message on channel because '{err:?}'"
);
Err(MessagingProtocolError::MessageSendFailed)
},
}
}
fn spawn_outbound_handler(
connectivity: ConnectivityRequester,
events_tx: mpsc::Sender<MessagingEvent>,
peer_node_id: NodeId,
retry_queue_tx: mpsc::UnboundedSender<OutboundMessage>,
protocol_id: ProtocolId,
) -> mpsc::UnboundedSender<OutboundMessage> {
let (msg_tx, msg_rx) = mpsc::unbounded_channel();
let outbound_messaging = OutboundMessaging::new(
connectivity,
events_tx,
msg_rx,
retry_queue_tx,
peer_node_id,
protocol_id,
);
tokio::spawn(outbound_messaging.run());
msg_tx
}
fn spawn_inbound_handler(&mut self, conn: PeerConnection, substream: Substream) {
let peer = conn.peer_node_id().clone();
if let Some(sessions) = self.active_inbound.get_mut(&peer) {
sessions.reap_finished();
}
if let Some(sessions) = self.active_inbound.get(&peer) {
match &sessions.current {
Some(current) if !current.handle.is_finished() => {
let budget = self.replacement_budgets.entry(peer.clone()).or_default();
if !budget.try_consume() {
let msg = format!(
"Peer '{}' exceeded the inbound session replacement rate ({} within {:.0?}); dropping \
this substream and keeping the existing session.",
peer.short_str(),
MAX_REPLACEMENTS_PER_WINDOW,
REPLACEMENT_RATE_WINDOW
);
if budget.should_warn_on_shed() {
warn!(target: LOG_TARGET, "{msg}");
} else {
debug!(target: LOG_TARGET, "{msg}");
}
return;
}
if sessions.stopping.len() >= MAX_STOPPING_SESSIONS_PER_PEER {
let msg = format!(
"Peer '{}' already has {} inbound session(s) still stopping; dropping this substream \
rather than letting them accumulate further.",
peer.short_str(),
MAX_STOPPING_SESSIONS_PER_PEER
);
if budget.should_warn_on_shed() {
warn!(target: LOG_TARGET, "{msg}");
} else {
debug!(target: LOG_TARGET, "{msg}");
}
return;
}
debug!(
target: LOG_TARGET,
"Replacing InboundMessaging session for peer '{}' with a session for its newest \
substream",
peer.short_str()
);
},
_ => {},
}
}
let messaging_events_tx = self.messaging_events_tx.clone();
let inbound_message_tx = self.inbound_message_tx.clone();
let (stop_tx, stop_rx) = oneshot::channel();
let session_id = self.next_inbound_session_id;
self.next_inbound_session_id = self.next_inbound_session_id.wrapping_add(1);
let inbound_messaging = InboundMessaging::new(
conn,
inbound_message_tx,
messaging_events_tx,
self.internal_messaging_event_tx.clone(),
session_id,
self.enable_message_received_event,
self.shutdown_signal.clone(),
stop_rx,
);
let handle = tokio::spawn(inbound_messaging.run(substream));
let new_session = ActiveInboundSession {
id: session_id,
handle,
stop_tx,
};
let sessions = self.active_inbound.entry(peer).or_default();
if let Some(outgoing) = sessions.current.replace(new_session) {
let _ignore = outgoing.stop_tx.send(());
sessions.stopping.push((outgoing.id, outgoing.handle));
}
}
async fn wait_for_connection(
connectivity: &mut ConnectivityRequester,
node_id: &NodeId,
) -> Result<Option<PeerConnection>, MessagingProtocolError> {
let deadline = time::Instant::now()
.checked_add(CONNECTION_LOOKUP_TIMEOUT)
.unwrap_or_else(time::Instant::now);
loop {
if let Some(conn) = connectivity.get_connection(node_id.clone(), RefKind::Weak).await? {
return Ok(Some(conn));
}
if time::Instant::now() >= deadline {
return Ok(None);
}
time::sleep(CONNECTION_LOOKUP_RETRY_INTERVAL).await;
}
}
async fn handle_protocol_notification(
&mut self,
notification: ProtocolNotification<Substream>,
) -> Result<(), MessagingProtocolError> {
match notification.event {
ProtocolEvent::NewInboundSubstream(node_id, substream) => {
trace!(
target: LOG_TARGET,
"NewInboundSubstream for peer '{}'",
node_id.short_str()
);
match Arc::clone(&self.pending_resolution_permits).try_acquire_owned() {
Ok(permit) => {
let mut connectivity = self.connectivity.clone();
let resolved_tx = self.resolved_substream_tx.clone();
let mut shutdown_signal = self.shutdown_signal.clone();
tokio::spawn(async move {
let _permit = permit;
let wait = Self::wait_for_connection(&mut connectivity, &node_id);
tokio::pin!(wait);
tokio::select! {
biased;
_ = &mut shutdown_signal => {},
result = &mut wait => match result {
Ok(Some(conn)) => {
if resolved_tx.send((conn, substream)).await.is_err() {
debug!(
target: LOG_TARGET,
"MessagingProtocol shut down before a resolved substream for \
peer '{}' could be handed off",
node_id.short_str()
);
}
},
Ok(None) => {
info!(
target: LOG_TARGET,
"No active connection for new inbound substream for node {node_id}"
);
},
Err(err) => {
error!(
target: LOG_TARGET,
"Failed to resolve connection for new inbound substream for node \
{node_id}: {err}"
);
},
},
}
});
},
Err(_) => {
let msg = format!(
"Already resolving {MAX_PENDING_SUBSTREAM_RESOLUTIONS} inbound substream(s); dropping the \
new substream for peer '{}' rather than growing that further unbounded.",
node_id.short_str()
);
let now = time::Instant::now();
let should_warn = match self.last_pending_resolution_shed_log {
Some(last) => now.saturating_duration_since(last) >= SHED_LOG_INTERVAL,
None => true,
};
if should_warn {
self.last_pending_resolution_shed_log = Some(now);
warn!(target: LOG_TARGET, "{msg}");
} else {
debug!(target: LOG_TARGET, "{msg}");
}
},
}
},
}
Ok(())
}
async fn ban_peer<T: Display>(&mut self, peer_node_id: NodeId, reason: T) {
warn!(
target: LOG_TARGET,
"Banning peer '{}' because it violated the messaging protocol: {}", peer_node_id.short_str(), reason
);
if let Some(sessions) = self.active_inbound.remove(&peer_node_id) {
if let Some(current) = sessions.current {
current.handle.abort();
}
for (_, handle) in sessions.stopping {
handle.abort();
}
}
self.replacement_budgets.remove(&peer_node_id);
drop(self.active_queues.remove(&peer_node_id));
match self.ban_duration {
Some(ban_duration) => {
if let Err(err) = self
.connectivity
.ban_peer_until(peer_node_id.clone(), ban_duration, reason.to_string())
.await
{
error!(
target: LOG_TARGET,
"Failed to ban peer '{}' because '{:?}'", peer_node_id.short_str(), err
);
}
},
None => {
warn!(
target: LOG_TARGET,
"Banning disabled in MessagingProtocol, so peer '{peer_node_id}' will not be banned (reason: {reason})",
);
},
}
}
}