use std::fmt;
use std::io;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, RwLock, Weak};
use std::time::Duration;
use dashmap::DashMap;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf};
use tokio::net::{TcpListener, TcpStream};
use tokio::runtime::Handle;
use tokio::sync::{Mutex, Notify};
use tokio::task::JoinHandle;
use crate::atom::{Atom, AtomTable};
use crate::distribution::handshake::{
HandshakeNode, initiate_handshake_async, respond_handshake_async,
};
use crate::distribution::resolver::NodeResolver;
const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum ConnectError {
ResolveFailure,
ConnectionRefused,
Timeout,
Io(String),
}
impl fmt::Display for ConnectError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::ResolveFailure => formatter.write_str("distribution node resolution failed"),
Self::ConnectionRefused => formatter.write_str("distribution TCP connection refused"),
Self::Timeout => formatter.write_str("distribution TCP connection timed out"),
Self::Io(error) => write!(formatter, "distribution TCP connection failed: {error}"),
}
}
}
impl std::error::Error for ConnectError {}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub enum ConnectionDownReason {
PeerClosed,
ReadError,
WriteError,
WriteTimeout,
ManualDisconnect,
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub struct ConnectionDownEvent {
pub node: Atom,
pub reason: ConnectionDownReason,
}
type ConnectionDownCallback = dyn Fn(ConnectionDownEvent) + Send + Sync + 'static;
type ControlFrameHandler = dyn Fn(&[u8], &[u8]) + Send + Sync + 'static;
#[derive(Clone, Default)]
pub struct ConnectionDownHook {
callback: Arc<RwLock<Option<Arc<ConnectionDownCallback>>>>,
}
impl ConnectionDownHook {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn register<F>(&self, callback: F)
where
F: Fn(ConnectionDownEvent) + Send + Sync + 'static,
{
let mut slot = self
.callback
.write()
.unwrap_or_else(|error| error.into_inner());
*slot = Some(Arc::new(callback));
}
pub fn unregister(&self) {
let mut slot = self
.callback
.write()
.unwrap_or_else(|error| error.into_inner());
*slot = None;
}
#[must_use]
pub fn is_registered(&self) -> bool {
self.callback
.read()
.unwrap_or_else(|error| error.into_inner())
.is_some()
}
fn invoke(&self, event: ConnectionDownEvent) {
let callback = self
.callback
.read()
.unwrap_or_else(|error| error.into_inner())
.clone();
if let Some(callback) = callback {
callback(event);
}
}
}
pub struct DistConnection {
node: Atom,
peer_addr: SocketAddr,
writer: Mutex<OwnedWriteHalf>,
down: AtomicBool,
manager: Weak<ConnectionManagerInner>,
}
impl DistConnection {
fn new(
node: Atom,
peer_addr: SocketAddr,
writer: OwnedWriteHalf,
manager: Weak<ConnectionManagerInner>,
) -> Self {
Self {
node,
peer_addr,
writer: Mutex::new(writer),
down: AtomicBool::new(false),
manager,
}
}
#[must_use]
pub fn node(&self) -> Atom {
self.node
}
#[must_use]
pub fn peer_addr(&self) -> SocketAddr {
self.peer_addr
}
#[must_use]
pub fn is_down(&self) -> bool {
self.down.load(Ordering::Acquire)
}
pub async fn write_raw(self: &Arc<Self>, bytes: &[u8]) -> io::Result<()> {
let result = {
let mut writer = self.writer.lock().await;
writer.write_all(bytes).await
};
if result.is_err() {
self.mark_down(ConnectionDownReason::WriteError);
}
result
}
pub fn mark_down_write_timeout(self: &Arc<Self>) {
self.mark_down(ConnectionDownReason::WriteTimeout);
}
fn mark_down(self: &Arc<Self>, reason: ConnectionDownReason) {
if self.down.swap(true, Ordering::AcqRel) {
return;
}
if let Some(manager) = self.manager.upgrade() {
manager.connection_down(self.node, self, reason);
}
}
}
pub struct AcceptHandle {
local_addr: SocketAddr,
shutdown: Arc<Notify>,
task: JoinHandle<()>,
}
impl AcceptHandle {
#[must_use]
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
pub fn shutdown(&self) {
self.shutdown.notify_waiters();
}
#[must_use]
pub fn is_finished(&self) -> bool {
self.task.is_finished()
}
}
impl Drop for AcceptHandle {
fn drop(&mut self) {
self.shutdown.notify_waiters();
self.task.abort();
}
}
struct ConnectionManagerInner {
connections: DashMap<Atom, Arc<DistConnection>>,
atom_table: Arc<AtomTable>,
resolver: Arc<dyn NodeResolver + Send + Sync>,
connect_timeout: Duration,
connection_down_hook: ConnectionDownHook,
control_frame_handler: RwLock<Option<Arc<ControlFrameHandler>>>,
cookie: String,
local_node_name: String,
local_creation: u32,
runtime_handle: RwLock<Option<Handle>>,
}
impl ConnectionManagerInner {
fn spawn_lifecycle<F>(&self, future: F) -> JoinHandle<()>
where
F: std::future::Future<Output = ()> + Send + 'static,
{
let handle = self
.runtime_handle
.read()
.unwrap_or_else(|error| error.into_inner())
.clone();
match handle {
Some(handle) => handle.spawn(future),
None => tokio::spawn(future),
}
}
fn handshake_node(&self) -> Result<HandshakeNode, ConnectError> {
HandshakeNode::with_default_flags(self.local_node_name.clone(), self.local_creation)
.map_err(|error| ConnectError::Io(error.to_string()))
}
fn gen_challenge(&self) -> u32 {
rand::random::<u32>()
}
}
impl ConnectionManagerInner {
fn connection_down(
&self,
node: Atom,
connection: &Arc<DistConnection>,
reason: ConnectionDownReason,
) {
let removed = self
.connections
.remove_if(&node, |_, current| Arc::ptr_eq(current, connection))
.is_some();
if removed {
self.connection_down_hook
.invoke(ConnectionDownEvent { node, reason });
}
}
}
#[derive(Clone)]
pub struct ConnectionManager {
inner: Arc<ConnectionManagerInner>,
}
impl ConnectionManager {
#[must_use]
pub fn new(
atom_table: Arc<AtomTable>,
resolver: Arc<dyn NodeResolver + Send + Sync>,
cookie: impl Into<String>,
local_node_name: impl Into<String>,
local_creation: u32,
) -> Self {
Self::with_connect_timeout(
atom_table,
resolver,
cookie,
local_node_name,
local_creation,
DEFAULT_CONNECT_TIMEOUT,
)
}
#[must_use]
pub fn with_connect_timeout(
atom_table: Arc<AtomTable>,
resolver: Arc<dyn NodeResolver + Send + Sync>,
cookie: impl Into<String>,
local_node_name: impl Into<String>,
local_creation: u32,
connect_timeout: Duration,
) -> Self {
Self {
inner: Arc::new(ConnectionManagerInner {
connections: DashMap::new(),
atom_table,
resolver,
connect_timeout,
connection_down_hook: ConnectionDownHook::new(),
control_frame_handler: RwLock::new(None),
cookie: cookie.into(),
local_node_name: local_node_name.into(),
local_creation,
runtime_handle: RwLock::new(None),
}),
}
}
pub fn set_runtime_handle(&self, handle: Handle) {
*self
.inner
.runtime_handle
.write()
.unwrap_or_else(|error| error.into_inner()) = Some(handle);
}
#[must_use]
pub fn connect_timeout(&self) -> Duration {
self.inner.connect_timeout
}
#[must_use]
pub fn connection_down_hook(&self) -> ConnectionDownHook {
self.inner.connection_down_hook.clone()
}
pub fn register_connection_down<F>(&self, callback: F)
where
F: Fn(ConnectionDownEvent) + Send + Sync + 'static,
{
self.inner.connection_down_hook.register(callback);
}
pub fn register_control_frame_handler<F>(&self, handler: F)
where
F: Fn(&[u8], &[u8]) + Send + Sync + 'static,
{
let mut slot = self
.inner
.control_frame_handler
.write()
.unwrap_or_else(|error| error.into_inner());
*slot = Some(Arc::new(handler));
}
#[must_use]
pub fn connection_count(&self) -> usize {
self.inner.connections.len()
}
#[must_use]
pub fn get_connection(&self, node: Atom) -> Option<Arc<DistConnection>> {
self.inner
.connections
.get(&node)
.map(|entry| Arc::clone(entry.value()))
}
#[must_use]
pub fn connected_nodes(&self) -> Vec<Atom> {
let mut nodes: Vec<_> = self
.inner
.connections
.iter()
.map(|entry| *entry.key())
.collect();
nodes.sort_unstable_by_key(|node| node.index());
nodes
}
pub async fn connect_node(&self, node: Atom) -> bool {
if self.get_connection(node).is_some() {
return true;
}
let Some(node_name) = self.inner.atom_table.resolve(node).map(str::to_owned) else {
return false;
};
self.connect(&node_name).await.is_ok()
}
pub fn disconnect_node(&self, node: Atom) -> bool {
let Some(connection) = self.get_connection(node) else {
return true;
};
connection.mark_down(ConnectionDownReason::ManualDisconnect);
true
}
pub async fn start(
listen_addr: SocketAddr,
resolver: Arc<dyn NodeResolver + Send + Sync>,
cookie: impl Into<String>,
local_node_name: impl Into<String>,
local_creation: u32,
) -> io::Result<(Self, AcceptHandle)> {
let manager = Self::new(
Arc::new(AtomTable::with_common_atoms()),
resolver,
cookie,
local_node_name,
local_creation,
);
let handle = manager.listen(listen_addr).await?;
Ok((manager, handle))
}
pub async fn listen(&self, listen_addr: SocketAddr) -> io::Result<AcceptHandle> {
let listener = TcpListener::bind(listen_addr).await?;
let local_addr = listener.local_addr()?;
let shutdown = Arc::new(Notify::new());
let task_shutdown = Arc::clone(&shutdown);
let manager = self.clone();
let task = self.inner.spawn_lifecycle(async move {
manager.accept_loop(listener, task_shutdown).await;
});
Ok(AcceptHandle {
local_addr,
shutdown,
task,
})
}
pub async fn connect(&self, node_name: &str) -> Result<Arc<DistConnection>, ConnectError> {
let addr = self
.inner
.resolver
.resolve(node_name)
.await
.map_err(|_| ConnectError::ResolveFailure)?;
let mut stream = match tokio::time::timeout(
self.inner.connect_timeout,
TcpStream::connect(addr),
)
.await
{
Ok(Ok(stream)) => stream,
Ok(Err(error)) if error.kind() == io::ErrorKind::ConnectionRefused => {
return Err(ConnectError::ConnectionRefused);
}
Ok(Err(error)) => return Err(ConnectError::Io(error.to_string())),
Err(_) => return Err(ConnectError::Timeout),
};
let peer_addr = stream.peer_addr().unwrap_or(addr);
let local = self.inner.handshake_node()?;
let result = initiate_handshake_async(
&mut stream,
&local,
&self.inner.cookie,
self.inner.gen_challenge(),
)
.await
.map_err(|error| ConnectError::Io(error.to_string()))?;
let node = self.inner.atom_table.intern(result.remote_name());
Ok(self.register_connection(node, peer_addr, stream))
}
fn register_connection(
&self,
node: Atom,
peer_addr: SocketAddr,
stream: TcpStream,
) -> Arc<DistConnection> {
let (read_half, write_half) = stream.into_split();
let connection = Arc::new(DistConnection::new(
node,
peer_addr,
write_half,
Arc::downgrade(&self.inner),
));
self.inner.connections.insert(node, Arc::clone(&connection));
self.spawn_read_lifecycle(Arc::clone(&connection), read_half);
connection
}
#[cfg(test)]
pub(crate) fn register_test_connection(
&self,
node: Atom,
peer_addr: SocketAddr,
stream: std::net::TcpStream,
) -> io::Result<Arc<DistConnection>> {
stream.set_nonblocking(true)?;
let stream = TcpStream::from_std(stream)?;
Ok(self.register_connection(node, peer_addr, stream))
}
fn spawn_read_lifecycle(&self, connection: Arc<DistConnection>, mut read_half: OwnedReadHalf) {
let manager = Arc::clone(&self.inner);
self.inner.spawn_lifecycle(async move {
loop {
let mut header = [0_u8; 8];
match read_half.read_exact(&mut header).await {
Ok(0) => {
connection.mark_down(ConnectionDownReason::PeerClosed);
break;
}
Ok(_) => {
let control_len =
u32::from_be_bytes([header[0], header[1], header[2], header[3]])
as usize;
let payload_len =
u32::from_be_bytes([header[4], header[5], header[6], header[7]])
as usize;
let Some(total_len) = control_len.checked_add(payload_len) else {
connection.mark_down(ConnectionDownReason::ReadError);
break;
};
let mut frame = vec![0_u8; total_len];
if read_half.read_exact(&mut frame).await.is_err() {
connection.mark_down(ConnectionDownReason::ReadError);
break;
}
let handler = manager
.control_frame_handler
.read()
.unwrap_or_else(|error| error.into_inner())
.clone();
if let Some(handler) = handler {
let (control, payload) = frame.split_at(control_len);
handler(control, payload);
}
}
Err(_) => {
connection.mark_down(ConnectionDownReason::ReadError);
break;
}
}
}
});
}
async fn accept_loop(&self, listener: TcpListener, shutdown: Arc<Notify>) {
loop {
tokio::select! {
_ = shutdown.notified() => {
break;
}
accepted = listener.accept() => {
let Ok((stream, peer_addr)) = accepted else {
continue;
};
self.handle_accepted(stream, peer_addr);
}
}
}
}
fn handle_accepted(&self, mut stream: TcpStream, peer_addr: SocketAddr) {
let manager = self.clone();
self.inner.spawn_lifecycle(async move {
let local = match manager.inner.handshake_node() {
Ok(local) => local,
Err(_) => return,
};
match respond_handshake_async(
&mut stream,
&local,
&manager.inner.cookie,
manager.inner.gen_challenge(),
)
.await
{
Ok(result) => {
let node = manager.inner.atom_table.intern(result.remote_name());
manager.register_connection(node, peer_addr, stream);
}
Err(_) => {
drop(stream);
}
}
});
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::net::TcpListener;
use tokio::task::JoinHandle;
use super::*;
use crate::distribution::handshake::HandshakeNode;
use crate::distribution::resolver::StaticResolver;
const TEST_COOKIE: &str = "test-cookie";
fn manager_with_resolver(resolver: Arc<StaticResolver>) -> ConnectionManager {
ConnectionManager::new(
Arc::new(AtomTable::with_common_atoms()),
resolver,
TEST_COOKIE,
"local@127.0.0.1",
1,
)
}
fn spawn_responder(
listener: TcpListener,
name: &'static str,
cookie: &'static str,
) -> JoinHandle<()> {
tokio::spawn(async move {
let Ok((mut stream, _peer)) = listener.accept().await else {
return;
};
let local = HandshakeNode::with_default_flags(name, 7)
.expect("responder node name should be valid");
let _ = crate::distribution::handshake::respond_handshake_async(
&mut stream,
&local,
cookie,
99,
)
.await;
tokio::time::sleep(Duration::from_millis(200)).await;
})
}
fn spawn_responder_handoff(
listener: TcpListener,
name: &'static str,
) -> tokio::sync::oneshot::Receiver<TcpStream> {
let (sender, receiver) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let Ok((mut stream, _peer)) = listener.accept().await else {
return;
};
let local = HandshakeNode::with_default_flags(name, 7)
.expect("responder node name should be valid");
if crate::distribution::handshake::respond_handshake_async(
&mut stream,
&local,
TEST_COOKIE,
99,
)
.await
.is_ok()
{
let _ = sender.send(stream);
}
});
receiver
}
#[tokio::test]
async fn empty_manager_has_no_connections() {
let manager = manager_with_resolver(Arc::new(StaticResolver::new(
std::collections::HashMap::new(),
)));
let node = manager.inner.atom_table.intern("missing@127.0.0.1");
assert_eq!(manager.connection_count(), 0);
assert!(manager.get_connection(node).is_none());
}
#[tokio::test]
async fn outbound_connect_inserts_table_entry() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.unwrap_or_else(|error| {
panic!("failed to bind local listener: {error}");
});
let addr = listener.local_addr().unwrap_or_else(|error| {
panic!("failed to inspect local listener: {error}");
});
let _responder = spawn_responder(listener, "remote@127.0.0.1", TEST_COOKIE);
let resolver = Arc::new(StaticResolver::new(std::collections::HashMap::from([(
"remote@127.0.0.1".to_string(),
addr,
)])));
let manager = manager_with_resolver(resolver);
let connection = manager
.connect("remote@127.0.0.1")
.await
.unwrap_or_else(|error| panic!("connect failed: {error}"));
let node = manager.inner.atom_table.intern("remote@127.0.0.1");
assert!(Arc::ptr_eq(
&connection,
&manager
.get_connection(node)
.expect("connection should be present"),
));
}
#[tokio::test]
async fn connect_keys_table_by_remote_handshake_name_not_resolver_key() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.unwrap_or_else(|error| panic!("failed to bind local listener: {error}"));
let addr = listener
.local_addr()
.unwrap_or_else(|error| panic!("failed to inspect local listener: {error}"));
let _responder = spawn_responder(listener, "advertised@127.0.0.1", TEST_COOKIE);
let resolver = Arc::new(StaticResolver::new(std::collections::HashMap::from([(
"dialed@127.0.0.1".to_string(),
addr,
)])));
let manager = manager_with_resolver(resolver);
let connection = manager
.connect("dialed@127.0.0.1")
.await
.unwrap_or_else(|error| panic!("connect failed: {error}"));
let advertised = manager.inner.atom_table.intern("advertised@127.0.0.1");
let dialed = manager.inner.atom_table.intern("dialed@127.0.0.1");
assert_eq!(connection.node(), advertised);
assert!(manager.get_connection(advertised).is_some());
assert!(
manager.get_connection(dialed).is_none(),
"connection must not be keyed by the resolver key"
);
}
#[tokio::test]
async fn connect_rejects_wrong_cookie_and_records_no_entry() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.unwrap_or_else(|error| panic!("failed to bind local listener: {error}"));
let addr = listener
.local_addr()
.unwrap_or_else(|error| panic!("failed to inspect local listener: {error}"));
let _responder = spawn_responder(listener, "remote@127.0.0.1", "other-cookie");
let resolver = Arc::new(StaticResolver::new(std::collections::HashMap::from([(
"remote@127.0.0.1".to_string(),
addr,
)])));
let manager = manager_with_resolver(resolver);
let result = manager.connect("remote@127.0.0.1").await;
assert!(
matches!(result, Err(ConnectError::Io(_))),
"connect must fail with Io on cookie mismatch"
);
assert_eq!(manager.connection_count(), 0);
let remote = manager.inner.atom_table.intern("remote@127.0.0.1");
assert!(manager.get_connection(remote).is_none());
}
#[tokio::test]
async fn inbound_wrong_cookie_registers_no_entry() {
let resolver = Arc::new(StaticResolver::new(std::collections::HashMap::new()));
let manager = manager_with_resolver(resolver);
let accept = manager
.listen("127.0.0.1:0".parse().unwrap_or_else(|error| {
panic!("failed to parse listen address: {error}");
}))
.await
.unwrap_or_else(|error| panic!("failed to start accept loop: {error}"));
let mut client = TcpStream::connect(accept.local_addr())
.await
.unwrap_or_else(|error| panic!("failed to open inbound stream: {error}"));
let client_node = HandshakeNode::with_default_flags("client@127.0.0.1", 5)
.expect("client node name should be valid");
let result = crate::distribution::handshake::initiate_handshake_async(
&mut client,
&client_node,
"wrong-cookie",
42,
)
.await;
assert!(
result.is_err(),
"inbound handshake with wrong cookie must fail"
);
let node = manager.inner.atom_table.intern("client@127.0.0.1");
for _ in 0..40 {
assert_eq!(
manager.connection_count(),
0,
"wrong-cookie peer must never register a connection"
);
assert!(
manager.get_connection(node).is_none(),
"wrong-cookie peer must not appear in the connection table"
);
tokio::time::sleep(Duration::from_millis(10)).await;
}
drop(client);
}
#[tokio::test]
async fn connect_node_is_idempotent_and_lists_node() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.unwrap_or_else(|error| panic!("failed to bind local listener: {error}"));
let addr = listener
.local_addr()
.unwrap_or_else(|error| panic!("failed to inspect local listener: {error}"));
let _responder = spawn_responder(listener, "remote@127.0.0.1", TEST_COOKIE);
let resolver = Arc::new(StaticResolver::new(std::collections::HashMap::from([(
"remote@127.0.0.1".to_string(),
addr,
)])));
let manager = manager_with_resolver(resolver);
let node = manager.inner.atom_table.intern("remote@127.0.0.1");
assert!(manager.connect_node(node).await);
assert!(manager.connect_node(node).await);
assert_eq!(manager.connected_nodes(), vec![node]);
assert_eq!(manager.connection_count(), 1);
}
#[tokio::test]
async fn connect_node_returns_false_for_unresolved_node() {
let manager = manager_with_resolver(Arc::new(StaticResolver::new(
std::collections::HashMap::new(),
)));
let node = manager.inner.atom_table.intern("missing@127.0.0.1");
assert!(!manager.connect_node(node).await);
assert!(manager.connected_nodes().is_empty());
}
#[tokio::test]
async fn inbound_peer_registers_under_its_handshake_name() {
let resolver = Arc::new(StaticResolver::new(std::collections::HashMap::new()));
let manager = manager_with_resolver(resolver);
let accept = manager
.listen("127.0.0.1:0".parse().unwrap_or_else(|error| {
panic!("failed to parse listen address: {error}");
}))
.await
.unwrap_or_else(|error| panic!("failed to start accept loop: {error}"));
let mut client = TcpStream::connect(accept.local_addr())
.await
.unwrap_or_else(|error| panic!("failed to open inbound stream: {error}"));
let client_node = HandshakeNode::with_default_flags("client@127.0.0.1", 5)
.expect("client node name should be valid");
crate::distribution::handshake::initiate_handshake_async(
&mut client,
&client_node,
TEST_COOKIE,
42,
)
.await
.expect("inbound peer handshake should succeed");
let node = manager.inner.atom_table.intern("client@127.0.0.1");
let mut connected = false;
for _ in 0..40 {
if manager.get_connection(node).is_some() {
connected = true;
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
assert!(
connected,
"inbound peer should register under its handshake name"
);
assert_eq!(manager.connected_nodes(), vec![node]);
drop(client);
}
#[tokio::test]
async fn dropping_peer_removes_connection_and_notifies_once() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.unwrap_or_else(|error| {
panic!("failed to bind local listener: {error}");
});
let addr = listener.local_addr().unwrap_or_else(|error| {
panic!("failed to inspect local listener: {error}");
});
let remote_stream = spawn_responder_handoff(listener, "remote@127.0.0.1");
let resolver = Arc::new(StaticResolver::new(std::collections::HashMap::from([(
"remote@127.0.0.1".to_string(),
addr,
)])));
let manager = manager_with_resolver(resolver);
let callback_count = Arc::new(AtomicUsize::new(0));
let callback_count_for_hook = Arc::clone(&callback_count);
manager.register_connection_down(move |_| {
callback_count_for_hook.fetch_add(1, Ordering::SeqCst);
});
let node = manager.inner.atom_table.intern("remote@127.0.0.1");
let _connection = manager
.connect("remote@127.0.0.1")
.await
.unwrap_or_else(|error| panic!("connect failed: {error}"));
let remote_stream = remote_stream
.await
.expect("responder did not complete handshake");
drop(remote_stream);
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(manager.get_connection(node).is_none());
assert_eq!(callback_count.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn manual_disconnect_removes_connection_and_notifies_once() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.unwrap_or_else(|error| panic!("failed to bind local listener: {error}"));
let addr = listener
.local_addr()
.unwrap_or_else(|error| panic!("failed to inspect local listener: {error}"));
let _responder = spawn_responder(listener, "remote@127.0.0.1", TEST_COOKIE);
let resolver = Arc::new(StaticResolver::new(std::collections::HashMap::from([(
"remote@127.0.0.1".to_string(),
addr,
)])));
let manager = manager_with_resolver(resolver);
let callback_count = Arc::new(AtomicUsize::new(0));
let callback_count_for_hook = Arc::clone(&callback_count);
manager.register_connection_down(move |event| {
assert_eq!(event.reason, ConnectionDownReason::ManualDisconnect);
callback_count_for_hook.fetch_add(1, Ordering::SeqCst);
});
let node = manager.inner.atom_table.intern("remote@127.0.0.1");
assert!(manager.connect_node(node).await);
assert!(manager.disconnect_node(node));
assert!(manager.disconnect_node(node));
assert!(manager.get_connection(node).is_none());
assert!(manager.connected_nodes().is_empty());
assert_eq!(callback_count.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn write_error_removes_connection_and_notifies_once() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.unwrap_or_else(|error| {
panic!("failed to bind local listener: {error}");
});
let addr = listener.local_addr().unwrap_or_else(|error| {
panic!("failed to inspect local listener: {error}");
});
let remote_stream = spawn_responder_handoff(listener, "remote@127.0.0.1");
let resolver = Arc::new(StaticResolver::new(std::collections::HashMap::from([(
"remote@127.0.0.1".to_string(),
addr,
)])));
let manager = manager_with_resolver(resolver);
let callback_count = Arc::new(AtomicUsize::new(0));
let callback_count_for_hook = Arc::clone(&callback_count);
manager.register_connection_down(move |_| {
callback_count_for_hook.fetch_add(1, Ordering::SeqCst);
});
let node = manager.inner.atom_table.intern("remote@127.0.0.1");
let connection = manager
.connect("remote@127.0.0.1")
.await
.unwrap_or_else(|error| panic!("connect failed: {error}"));
let remote_stream = remote_stream
.await
.expect("responder did not complete handshake");
drop(remote_stream);
for _ in 0..8 {
if connection.write_raw(b"probe").await.is_err() {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
tokio::time::sleep(Duration::from_millis(25)).await;
assert!(manager.get_connection(node).is_none());
assert_eq!(callback_count.load(Ordering::SeqCst), 1);
}
}