use std::{
future::Future,
net::SocketAddr,
panic::{AssertUnwindSafe, resume_unwind},
};
use futures_util::FutureExt;
use tokio::{
sync::{mpsc, oneshot},
task::JoinHandle,
};
use tracing::*;
#[cfg(doc)]
use crate::{
Connection, Node,
protocols::{Handshake, OnDisconnect, Reading, Writing},
};
use crate::{
Pea2Pea,
connections::{DisconnectOrigin, create_connection_span},
node::NodeTask,
protocols::{
DisconnectOnDrop, ProtocolHandler, install_protocol_handler, panic_message,
run_hook_handler_loop,
},
};
pub(crate) type OnConnectBundle = (JoinHandle<()>, bool);
pub trait OnConnect: Pea2Pea
where
Self: Clone + Send + Sync + 'static,
{
const ABORTABLE: bool = true;
fn enable_on_connect(&self) -> impl Future<Output = ()> {
async {
let (from_node_sender, from_node_receiver) =
mpsc::channel::<((SocketAddr, u64), oneshot::Sender<OnConnectBundle>)>(
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, conn_id), notifier)| {
let self_clone = self_clone.clone();
let handle = tokio::spawn(async move {
let mut conn_cleanup = DisconnectOnDrop::new(
self_clone.node().clone(),
addr,
conn_id,
DisconnectOrigin::OnConnectAbort,
);
match AssertUnwindSafe(self_clone.on_connect(addr))
.catch_unwind()
.await
{
Ok(()) => {
conn_cleanup.node.take();
}
Err(payload) => {
let conn_span =
create_connection_span(addr, self_clone.node().span());
error!(parent: conn_span, "OnConnect::on_connect panicked: {}", panic_message(&*payload));
resume_unwind(payload); }
}
});
let _ = notifier.send((handle, Self::ABORTABLE)); })
.await;
};
install_protocol_handler(
self.node(),
NodeTask::OnConnect,
"OnConnect",
|protocols| &protocols.on_connect,
ProtocolHandler(from_node_sender),
handler_loop,
)
.await;
}
}
fn on_connect(&self, addr: SocketAddr) -> impl Future<Output = ()> + Send;
}