use std::{
any::Any,
collections::{HashMap, hash_map::Entry},
future::{Future, poll_fn},
io,
net::SocketAddr,
panic::{AssertUnwindSafe, resume_unwind},
sync::Arc,
time::Duration,
};
#[cfg(doc)]
use bytes::Bytes;
use futures_util::{FutureExt, SinkExt};
use parking_lot::RwLock;
use tokio::{
io::AsyncWrite,
sync::{mpsc, oneshot},
task::JoinHandle,
time::timeout,
};
use tokio_util::codec::{Encoder, FramedWrite};
use tracing::*;
#[cfg(doc)]
use crate::{Config, Node, protocols::Handshake};
use crate::{
Connection, ConnectionSide, Pea2Pea,
connections::{DisconnectOrigin, create_connection_span},
node::NodeTask,
protocols::{
DisconnectOnDrop, Protocol, ProtocolHandler, ReturnableConnection, catch_setup_panic,
install_protocol_handler, panic_message, run_setup_handler_loop,
},
};
type WritingSenders = Arc<RwLock<HashMap<SocketAddr, (u64, Arc<dyn Any + Send + Sync>)>>>;
pub trait Writing: Pea2Pea
where
Self: Clone + Send + Sync + 'static,
{
const MESSAGE_QUEUE_DEPTH: usize = 64;
const INITIAL_BUFFER_SIZE: usize = 64 * 1024;
const TIMEOUT_MS: u64 = 10_000;
type Message: Send;
type Codec: Encoder<Self::Message, Error = io::Error> + Send;
fn enable_writing(&self) -> impl Future<Output = ()> {
async {
assert!(
Self::MESSAGE_QUEUE_DEPTH != 0,
"Writing::MESSAGE_QUEUE_DEPTH must not be 0"
);
let (conn_sender, conn_receiver) =
mpsc::channel(self.node().config().max_connecting as usize);
let conn_senders: WritingSenders = Default::default();
let senders = conn_senders.clone();
let self_clone = self.clone();
let handler_loop = async move {
let node = self_clone.node().clone();
run_setup_handler_loop(node, "Writing", conn_receiver, |conn, setup_tasks| {
let self_clone = self_clone.clone();
let senders = conn_senders.clone();
setup_tasks.spawn(async move {
self_clone.handle_new_connection(conn, &senders).await;
});
})
.await;
};
install_protocol_handler(
self.node(),
NodeTask::Writing,
"Writing",
|protocols| &protocols.writing,
WritingHandler {
handler: ProtocolHandler(conn_sender),
senders,
},
handler_loop,
)
.await;
}
}
fn codec(&self, addr: SocketAddr, side: ConnectionSide) -> Self::Codec;
fn unicast(
&self,
addr: SocketAddr,
message: Self::Message,
) -> io::Result<oneshot::Receiver<io::Result<()>>> {
self.queue_message(addr, message, true)
.map(|delivery| delivery.unwrap()) }
fn unicast_fast(&self, addr: SocketAddr, message: Self::Message) -> io::Result<()> {
self.queue_message(addr, message, false).map(|_| ())
}
#[deprecated(
since = "0.57.0",
note = "use Node::connected_addrs + unicast(_fast) depending on your specific use case"
)]
fn broadcast(&self, message: Self::Message) -> io::Result<()>
where
Self::Message: Clone,
{
let Some(handler) = self.node().protocols.writing.get() else {
return Err(io::ErrorKind::Unsupported.into());
};
let addrs: Vec<_> = handler.senders.read().keys().copied().collect();
for addr in addrs {
if let Err(e) = self.queue_message(addr, message.clone(), false) {
if e.kind() == io::ErrorKind::Unsupported {
return Err(e);
}
}
}
Ok(())
}
}
trait WritingInternal: Writing {
fn queue_message(
&self,
addr: SocketAddr,
message: Self::Message,
confirm_delivery: bool,
) -> io::Result<Option<oneshot::Receiver<io::Result<()>>>>;
async fn write_batch<W: AsyncWrite + Unpin + Send>(
&self,
messages: &mut Vec<WrappedMessage<Self::Message>>,
writer: &mut FramedWrite<W, Self::Codec>,
) -> io::Result<(usize, usize)>;
async fn handle_new_connection(
&self,
conn_with_returner: ReturnableConnection,
conn_senders: &WritingSenders,
);
}
impl<W: Writing> WritingInternal for W {
fn queue_message(
&self,
addr: SocketAddr,
message: Self::Message,
confirm_delivery: bool,
) -> io::Result<Option<oneshot::Receiver<io::Result<()>>>> {
let Some(handler) = self.node().protocols.writing.get() else {
return Err(io::ErrorKind::Unsupported.into());
};
let senders = handler.senders.read();
let Some((_, erased)) = senders.get(&addr) else {
return Err(io::ErrorKind::NotConnected.into());
};
let sender = erased
.downcast_ref::<mpsc::Sender<WrappedMessage<Self::Message>>>()
.ok_or(io::ErrorKind::Unsupported)?;
let (msg, delivery) = WrappedMessage::new(message, confirm_delivery);
let result = sender.try_send(msg);
drop(senders);
result.map(|_| delivery).map_err(|e| {
let conn_span = create_connection_span(addr, self.node().span());
error!(parent: conn_span, "can't send a message: {e}");
match e {
mpsc::error::TrySendError::Full(_) => io::ErrorKind::QuotaExceeded.into(),
mpsc::error::TrySendError::Closed(_) => io::ErrorKind::BrokenPipe.into(),
}
})
}
async fn write_batch<A: AsyncWrite + Unpin + Send>(
&self,
messages: &mut Vec<WrappedMessage<Self::Message>>,
writer: &mut FramedWrite<A, Self::Codec>,
) -> io::Result<(usize, usize)> {
let write = async move {
let msgs = messages.len();
let mut bytes = 0;
for wrapped in messages {
let msg = wrapped.msg.take().unwrap(); poll_fn(|cx| writer.poll_ready_unpin(cx)).await?;
let prev = writer.write_buffer().len();
writer.start_send_unpin(msg)?;
bytes += writer.write_buffer().len().saturating_sub(prev);
}
writer.flush().await?;
Ok::<_, io::Error>((msgs, bytes))
};
match timeout(Duration::from_millis(Self::TIMEOUT_MS), write).await {
Ok(Ok(stats)) => Ok(stats),
Ok(Err(e)) => Err(e),
Err(_) => {
self.node().heuristics().register_write_timeout();
Err(io::Error::new(io::ErrorKind::TimedOut, "write timed out"))
}
}
}
async fn handle_new_connection(
&self,
(mut conn, conn_returner): ReturnableConnection,
conn_senders: &WritingSenders,
) {
let addr = conn.addr();
let conn_id = conn.id;
let codec = match catch_setup_panic(conn.span(), "Writing::codec", || {
self.codec(addr, !conn.side())
}) {
Ok(codec) => codec,
Err(err) => {
let _ = conn_returner.send(Err(err));
return;
}
};
let Some(writer) = conn.writer.take() else {
let err = io::Error::other("the stream was not returned during the handshake");
error!(parent: conn.span(), "{err}");
let _ = conn_returner.send(Err(err));
return;
};
let mut framed = FramedWrite::new(writer, codec);
if Self::INITIAL_BUFFER_SIZE != 0 {
framed.write_buffer_mut().reserve(Self::INITIAL_BUFFER_SIZE);
let boundary = framed
.backpressure_boundary()
.max(Self::INITIAL_BUFFER_SIZE);
framed.set_backpressure_boundary(boundary);
}
let (outbound_message_sender, mut outbound_message_receiver) =
mpsc::channel(Self::MESSAGE_QUEUE_DEPTH);
let sender_cleanup = SenderCleanup {
addr,
conn_id,
senders: Arc::clone(conn_senders),
};
let (tx_writer, rx_writer) = oneshot::channel();
let self_clone = self.clone();
let conn_stats = conn.stats().clone();
let conn_span = conn.span().clone();
let writer_task = tokio::spawn(Box::pin(async move {
let node = self_clone.node();
trace!(parent: &conn_span, "spawned a task for writing messages");
if tx_writer.send(()).is_err() {
error!(parent: &conn_span, "Writing was interrupted; shutting down its task");
return;
}
let _sender_cleanup = sender_cleanup;
let _conn_cleanup =
DisconnectOnDrop::new(node.clone(), addr, conn_id, DisconnectOrigin::Writing);
let writing = AssertUnwindSafe(async {
let mut batch: Vec<WrappedMessage<Self::Message>> = Vec::new();
while outbound_message_receiver
.recv_many(&mut batch, Self::MESSAGE_QUEUE_DEPTH)
.await
> 0
{
match self_clone.write_batch(&mut batch, &mut framed).await {
Ok((msgs, bytes)) => {
for tx in batch.drain(..).filter_map(|w| w.delivery_notification) {
let _ = tx.send(Ok(()));
}
conn_stats.register_sent_messages(msgs, bytes);
node.stats().register_sent_messages(msgs, bytes);
trace!(parent: &conn_span, "wrote {bytes}B ({msgs} messages)");
}
Err(e) => {
error!(parent: &conn_span, "couldn't write a batch of {} message(s): {e}", batch.len());
for tx in batch.drain(..).filter_map(|w| w.delivery_notification) {
let _ = tx.send(Err(io::Error::new(e.kind(), e.to_string())));
}
break;
}
}
}
});
if let Err(payload) = writing.catch_unwind().await {
error!(parent: &conn_span, "Writing::Codec panicked while encoding: {}", panic_message(&*payload));
resume_unwind(payload);
}
}));
let _ = rx_writer.await;
adopt_writer(
&mut conn,
conn_senders,
conn_id,
Arc::new(outbound_message_sender),
writer_task,
);
let conn_span = conn.span().clone();
if conn_returner.send(Ok(conn)).is_err() {
error!(parent: &conn_span, "couldn't return a Connection from the Writing handler");
}
}
}
fn adopt_writer(
conn: &mut Connection,
conn_senders: &WritingSenders,
conn_id: u64,
sender: Arc<dyn Any + Send + Sync>,
writer_task: JoinHandle<()>,
) {
conn_senders.write().insert(conn.addr(), (conn_id, sender));
conn.tasks.push(writer_task);
}
pub(crate) struct WrappedMessage<T> {
msg: Option<T>,
delivery_notification: Option<oneshot::Sender<io::Result<()>>>,
}
impl<T> WrappedMessage<T> {
fn new(msg: T, confirmation: bool) -> (Self, Option<oneshot::Receiver<io::Result<()>>>) {
let (tx, rx) = if confirmation {
let (tx, rx) = oneshot::channel();
(Some(tx), Some(rx))
} else {
(None, None)
};
let wrapped_msg = Self {
msg: Some(msg),
delivery_notification: tx,
};
(wrapped_msg, rx)
}
}
pub(crate) struct WritingHandler {
handler: ProtocolHandler<Connection, io::Result<Connection>>,
pub(crate) senders: WritingSenders,
}
impl WritingHandler {
pub(crate) async fn closed(&self) {
self.handler.closed().await;
}
}
impl Protocol<Connection, io::Result<Connection>> for WritingHandler {
async fn trigger(&self, item: ReturnableConnection) {
self.handler.trigger(item).await;
}
}
struct SenderCleanup {
addr: SocketAddr,
senders: WritingSenders,
conn_id: u64,
}
impl Drop for SenderCleanup {
fn drop(&mut self) {
let mut map = self.senders.write();
if let Entry::Occupied(e) = map.entry(self.addr) {
if e.get().0 == self.conn_id {
e.remove();
}
}
}
}