use std::{
future::Future,
net::SocketAddr,
panic::{AssertUnwindSafe, resume_unwind},
time::Duration,
};
use futures_util::FutureExt;
use tokio::{
sync::{mpsc, oneshot},
task::JoinHandle,
time::timeout,
};
use tracing::*;
#[cfg(doc)]
use crate::{
Connection, Node,
protocols::{OnConnect, Reading, Writing},
};
use crate::{
Pea2Pea,
connections::{DisconnectOrigin, create_connection_span},
node::NodeTask,
protocols::{ProtocolHandler, install_protocol_handler, panic_message, run_hook_handler_loop},
};
pub(crate) type OnDisconnectBundle = (JoinHandle<()>, oneshot::Receiver<()>);
pub trait OnDisconnect: Pea2Pea
where
Self: Clone + Send + Sync + 'static,
{
const TIMEOUT_MS: u64 = 3_000;
fn enable_on_disconnect(&self) -> impl Future<Output = ()> {
async {
let (from_node_sender, from_node_receiver) =
mpsc::channel::<(
(SocketAddr, DisconnectOrigin),
oneshot::Sender<OnDisconnectBundle>,
)>(self.node().config().max_connections as usize);
let self_clone = self.clone();
let handler_loop = async move {
let node = self_clone.node().clone();
run_hook_handler_loop(node, from_node_receiver, |((addr, origin), notifier)| {
let self_clone = self_clone.clone();
let (done_tx, done_rx) = oneshot::channel();
let handle = tokio::spawn(async move {
let hook = AssertUnwindSafe(self_clone.on_disconnect(addr, origin))
.catch_unwind();
match timeout(Duration::from_millis(Self::TIMEOUT_MS), hook).await {
Ok(Ok(())) => {}
Ok(Err(payload)) => {
let conn_span =
create_connection_span(addr, self_clone.node().span());
error!(parent: conn_span, "OnDisconnect::on_disconnect panicked: {}", panic_message(&*payload));
resume_unwind(payload);
}
Err(_) => {
let conn_span =
create_connection_span(addr, self_clone.node().span());
warn!(parent: conn_span, "OnDisconnect logic timed out");
}
}
let _ = done_tx.send(());
});
if let Err((handle, _)) = notifier.send((handle, done_rx)) {
handle.abort();
}
})
.await;
};
install_protocol_handler(
self.node(),
NodeTask::OnDisconnect,
"OnDisconnect",
|protocols| &protocols.on_disconnect,
ProtocolHandler(from_node_sender),
handler_loop,
)
.await;
}
}
fn on_disconnect(
&self,
addr: SocketAddr,
origin: DisconnectOrigin,
) -> impl Future<Output = ()> + Send;
}