use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::runtime::{Handle, Runtime};
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use crate::atom::Atom;
use crate::distribution::connection::{ConnectionManager, DistConnection};
use crate::distribution::join_runtime_drop;
pub const DIST_SEND_THREAD_NAME: &str = "beamr-dist-send";
pub const DIST_SEND_QUEUE_CAP: usize = 1024;
pub const DIST_CONTROL_QUEUE_CAP: usize = 256;
pub(crate) const WRITE_TIMEOUT: Duration = Duration::from_secs(5);
#[derive(Clone, Debug)]
pub enum DistOutbound {
ToNode {
node: Atom,
frame: Arc<[u8]>,
},
}
#[derive(Clone)]
pub struct ControlOutbound {
pub connection: Arc<DistConnection>,
pub frame: Arc<[u8]>,
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub enum ControlEnqueueError {
Overflow,
Closed,
}
struct DistSenderInner {
runtime: Mutex<Option<Runtime>>,
handle: Handle,
drain: JoinHandle<()>,
mark: u64,
}
impl Drop for DistSenderInner {
fn drop(&mut self) {
self.drain.abort();
let runtime = self
.runtime
.get_mut()
.unwrap_or_else(|error| error.into_inner())
.take();
join_runtime_drop(runtime, self.mark);
}
}
#[derive(Clone)]
pub struct DistSender {
tx: mpsc::Sender<DistOutbound>,
control_tx: mpsc::Sender<ControlOutbound>,
inner: Arc<DistSenderInner>,
}
impl DistSender {
#[must_use]
pub fn new(connections: ConnectionManager) -> Option<Self> {
let mark = crate::distribution::mint_runtime_mark();
let mut builder = tokio::runtime::Builder::new_multi_thread();
builder
.worker_threads(1)
.thread_name(DIST_SEND_THREAD_NAME)
.enable_all();
crate::distribution::stamp_runtime_threads(&mut builder, mark);
let runtime = builder.build().ok()?;
let (tx, mut rx) = mpsc::channel::<DistOutbound>(DIST_SEND_QUEUE_CAP);
let (control_tx, mut control_rx) = mpsc::channel::<ControlOutbound>(DIST_CONTROL_QUEUE_CAP);
let drain = runtime.spawn(async move {
let mut control_open = true;
let mut data_open = true;
while control_open || data_open {
tokio::select! {
biased;
item = control_rx.recv(), if control_open => match item {
Some(item) => {
if item.connection.is_down() {
continue;
}
if tokio::time::timeout(
WRITE_TIMEOUT,
item.connection.write_raw(&item.frame),
)
.await
.is_err()
{
item.connection.mark_down_write_timeout();
}
}
None => control_open = false,
},
item = rx.recv(), if data_open => match item {
Some(DistOutbound::ToNode { node, frame }) => {
if let Some(connection) = connections.get_connection(node) {
if tokio::time::timeout(WRITE_TIMEOUT, connection.write_raw(&frame))
.await
.is_err()
{
connection.mark_down_write_timeout();
}
}
}
None => data_open = false,
},
}
}
});
let handle = runtime.handle().clone();
Some(Self {
tx,
control_tx,
inner: Arc::new(DistSenderInner {
runtime: Mutex::new(Some(runtime)),
mark,
handle,
drain,
}),
})
}
#[must_use]
pub fn handle(&self) -> Handle {
self.inner.handle.clone()
}
#[must_use]
pub fn worker_thread_names(&self) -> Vec<String> {
if self
.inner
.runtime
.lock()
.unwrap_or_else(|error| error.into_inner())
.is_some()
{
vec![DIST_SEND_THREAD_NAME.to_owned()]
} else {
Vec::new()
}
}
pub fn enqueue(&self, item: DistOutbound) {
let _ = self.tx.try_send(item);
}
pub fn enqueue_control(&self, item: ControlOutbound) -> Result<(), ControlEnqueueError> {
match self.control_tx.try_send(item) {
Ok(()) => Ok(()),
Err(mpsc::error::TrySendError::Full(item)) => {
item.connection.mark_down_control_overflow();
Err(ControlEnqueueError::Overflow)
}
Err(mpsc::error::TrySendError::Closed(_)) => Err(ControlEnqueueError::Closed),
}
}
pub fn shutdown(&self) {
self.inner.drain.abort();
let runtime = self
.inner
.runtime
.lock()
.unwrap_or_else(|error| error.into_inner())
.take();
join_runtime_drop(runtime, self.inner.mark);
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::io::AsyncReadExt;
use tokio::net::TcpListener;
use super::*;
use crate::atom::AtomTable;
use crate::distribution::connection::{ConnectionDownReason, ConnectionManager};
use crate::distribution::resolver::StaticResolver;
fn manager() -> (ConnectionManager, Arc<AtomTable>) {
let atom_table = Arc::new(AtomTable::with_common_atoms());
let resolver = Arc::new(StaticResolver::new(HashMap::new()));
(
ConnectionManager::new(
Arc::clone(&atom_table),
resolver,
"test-cookie",
"local@test",
0,
),
atom_table,
)
}
fn framed(control: &[u8]) -> Arc<[u8]> {
let control_len = u32::try_from(control.len()).expect("control fits u32");
let mut frame = Vec::with_capacity(8 + control.len());
frame.extend_from_slice(&control_len.to_be_bytes());
frame.extend_from_slice(&0u32.to_be_bytes());
frame.extend_from_slice(control);
Arc::from(frame.into_boxed_slice())
}
#[test]
fn enqueue_is_non_blocking_and_drops_when_full() {
let (connections, atom_table) = manager();
let sender = DistSender::new(connections).expect("sender builds");
let node = atom_table.intern("absent@127.0.0.1");
for index in 0..(DIST_SEND_QUEUE_CAP * 4) {
sender.enqueue(DistOutbound::ToNode {
node,
frame: framed(&index.to_be_bytes()),
});
}
sender.shutdown();
}
#[tokio::test]
async fn per_node_fifo_ordering() {
let (connections, atom_table) = manager();
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
let received = Arc::new(Mutex::new(Vec::new()));
let received_for_task = Arc::clone(&received);
let count = 16usize;
let reader = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("accept");
for _ in 0..count {
let mut header = [0u8; 8];
if stream.read_exact(&mut header).await.is_err() {
break;
}
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 mut body = vec![0u8; control_len + payload_len];
if stream.read_exact(&mut body).await.is_err() {
break;
}
received_for_task
.lock()
.unwrap_or_else(|error| error.into_inner())
.push(body[0]);
}
});
let std_stream = std::net::TcpStream::connect(addr).expect("client connects");
let node = atom_table.intern("peer@127.0.0.1");
let peer_addr: SocketAddr = std_stream.peer_addr().expect("peer addr");
connections
.register_test_connection(node, peer_addr, std_stream)
.expect("register test connection");
let sender = DistSender::new(connections).expect("sender builds");
for index in 0..count {
let seq = u8::try_from(index).expect("seq fits u8");
sender.enqueue(DistOutbound::ToNode {
node,
frame: framed(&[seq]),
});
}
reader.await.expect("reader task joins");
let order = received
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone();
let expected: Vec<u8> = (0..count).map(|i| i as u8).collect();
assert_eq!(order, expected, "frames must arrive in enqueue order");
sender.shutdown();
drop(sender);
}
#[tokio::test]
async fn dead_peer_does_not_stall_drain() {
let (connections, atom_table) = manager();
let down_count = Arc::new(AtomicUsize::new(0));
let down_for_hook = Arc::clone(&down_count);
connections.register_connection_down(move |_| {
down_for_hook.fetch_add(1, Ordering::SeqCst);
});
let dead_listener = TcpListener::bind("127.0.0.1:0").await.expect("bind dead");
let dead_addr = dead_listener.local_addr().expect("dead addr");
let dead_node = atom_table.intern("dead@127.0.0.1");
let dead_stream = std::net::TcpStream::connect(dead_addr).expect("dead connects");
let dead_peer_addr = dead_stream.peer_addr().expect("dead peer addr");
let dead_accept = tokio::spawn(async move { dead_listener.accept().await });
connections
.register_test_connection(dead_node, dead_peer_addr, dead_stream)
.expect("register dead connection");
let accepted = dead_accept
.await
.expect("dead accept join")
.expect("accepted");
drop(accepted);
let live_listener = TcpListener::bind("127.0.0.1:0").await.expect("bind live");
let live_addr = live_listener.local_addr().expect("live addr");
let live_received = Arc::new(Mutex::new(Vec::new()));
let live_for_task = Arc::clone(&live_received);
let live_reader = tokio::spawn(async move {
let (mut stream, _) = live_listener.accept().await.expect("live accept");
let mut header = [0u8; 8];
if stream.read_exact(&mut header).await.is_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 mut body = vec![0u8; control_len + payload_len];
if stream.read_exact(&mut body).await.is_ok() {
live_for_task
.lock()
.unwrap_or_else(|error| error.into_inner())
.push(body[0]);
}
}
});
let live_stream = std::net::TcpStream::connect(live_addr).expect("live connects");
let live_node = atom_table.intern("live@127.0.0.1");
let live_peer_addr = live_stream.peer_addr().expect("live peer addr");
connections
.register_test_connection(live_node, live_peer_addr, live_stream)
.expect("register live connection");
let sender = DistSender::new(connections.clone()).expect("sender builds");
for index in 0..32u8 {
sender.enqueue(DistOutbound::ToNode {
node: dead_node,
frame: framed(&[index]),
});
}
sender.enqueue(DistOutbound::ToNode {
node: live_node,
frame: framed(&[0xAB]),
});
live_reader.await.expect("live reader joins");
let got = live_received
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone();
assert_eq!(got, vec![0xAB], "live node must still receive its frame");
let deadline = std::time::Instant::now() + Duration::from_secs(5);
while down_count.load(Ordering::SeqCst) == 0 {
assert!(
std::time::Instant::now() < deadline,
"dead peer down-hook never fired"
);
tokio::time::sleep(Duration::from_millis(10)).await;
}
assert!(connections.get_connection(dead_node).is_none());
sender.shutdown();
drop(sender);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn wedged_peer_does_not_stall_drain_indefinitely() {
let (connections, atom_table) = manager();
let down_count = Arc::new(AtomicUsize::new(0));
let down_for_hook = Arc::clone(&down_count);
connections.register_connection_down(move |_| {
down_for_hook.fetch_add(1, Ordering::SeqCst);
});
let wedged_listener = TcpListener::bind("127.0.0.1:0").await.expect("bind wedged");
let wedged_addr = wedged_listener.local_addr().expect("wedged addr");
let wedged_node = atom_table.intern("wedged@127.0.0.1");
let wedged_stream = std::net::TcpStream::connect(wedged_addr).expect("wedged connects");
let wedged_peer_addr = wedged_stream.peer_addr().expect("wedged peer addr");
let wedged_accept = tokio::spawn(async move { wedged_listener.accept().await });
connections
.register_test_connection(wedged_node, wedged_peer_addr, wedged_stream)
.expect("register wedged connection");
let wedged_accepted = wedged_accept
.await
.expect("wedged accept join")
.expect("wedged accepted");
let _wedged_held = wedged_accepted;
let live_listener = TcpListener::bind("127.0.0.1:0").await.expect("bind live");
let live_addr = live_listener.local_addr().expect("live addr");
let live_received = Arc::new(Mutex::new(Vec::new()));
let live_for_task = Arc::clone(&live_received);
let live_reader = tokio::spawn(async move {
let (mut stream, _) = live_listener.accept().await.expect("live accept");
let mut header = [0u8; 8];
if stream.read_exact(&mut header).await.is_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 mut body = vec![0u8; control_len + payload_len];
if stream.read_exact(&mut body).await.is_ok() {
live_for_task
.lock()
.unwrap_or_else(|error| error.into_inner())
.push(body[0]);
}
}
});
let live_stream = std::net::TcpStream::connect(live_addr).expect("live connects");
let live_node = atom_table.intern("live@127.0.0.1");
let live_peer_addr = live_stream.peer_addr().expect("live peer addr");
connections
.register_test_connection(live_node, live_peer_addr, live_stream)
.expect("register live connection");
let sender = DistSender::new(connections.clone()).expect("sender builds");
let mut big = vec![0u8; 16 * 1024 * 1024];
big[0] = 0x01;
let big_control_len = u32::try_from(big.len()).expect("control fits u32");
let mut wedged_frame = Vec::with_capacity(8 + big.len());
wedged_frame.extend_from_slice(&big_control_len.to_be_bytes());
wedged_frame.extend_from_slice(&0u32.to_be_bytes());
wedged_frame.extend_from_slice(&big);
sender.enqueue(DistOutbound::ToNode {
node: wedged_node,
frame: Arc::from(wedged_frame.into_boxed_slice()),
});
sender.enqueue(DistOutbound::ToNode {
node: live_node,
frame: framed(&[0xAB]),
});
let live_join = tokio::time::timeout(Duration::from_secs(30), live_reader)
.await
.expect("healthy peer received within the bounded window (not a ~2h stall)");
live_join.expect("live reader task joins");
let got = live_received
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone();
assert_eq!(got, vec![0xAB], "healthy node must still receive its frame");
let deadline = std::time::Instant::now() + Duration::from_secs(10);
while down_count.load(Ordering::SeqCst) == 0 {
assert!(
std::time::Instant::now() < deadline,
"wedged peer down-hook never fired after the write timeout"
);
tokio::time::sleep(Duration::from_millis(10)).await;
}
assert!(
connections.get_connection(wedged_node).is_none(),
"wedged connection must be purged after the write timeout"
);
sender.shutdown();
drop(sender);
}
#[tokio::test]
async fn control_lane_per_node_fifo_ordering() {
let (connections, atom_table) = manager();
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
let received = Arc::new(Mutex::new(Vec::new()));
let received_for_task = Arc::clone(&received);
let count = 16usize;
let reader = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("accept");
for _ in 0..count {
let mut header = [0u8; 8];
if stream.read_exact(&mut header).await.is_err() {
break;
}
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 mut body = vec![0u8; control_len + payload_len];
if stream.read_exact(&mut body).await.is_err() {
break;
}
received_for_task
.lock()
.unwrap_or_else(|error| error.into_inner())
.push(body[0]);
}
});
let std_stream = std::net::TcpStream::connect(addr).expect("client connects");
let node = atom_table.intern("peer@127.0.0.1");
let peer_addr: SocketAddr = std_stream.peer_addr().expect("peer addr");
let connection = connections
.register_test_connection(node, peer_addr, std_stream)
.expect("register test connection");
let sender = DistSender::new(connections).expect("sender builds");
for index in 0..count {
let seq = u8::try_from(index).expect("seq fits u8");
sender
.enqueue_control(ControlOutbound {
connection: Arc::clone(&connection),
frame: framed(&[seq]),
})
.expect("control lane accepts below capacity");
}
reader.await.expect("reader task joins");
let order = received
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone();
let expected: Vec<u8> = (0..count).map(|i| i as u8).collect();
assert_eq!(order, expected, "controls must arrive in enqueue order");
sender.shutdown();
drop(sender);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn control_lane_overflow_flood_marks_wedged_peer_down_exactly_once() {
let (connections, atom_table) = manager();
let down_reasons = Arc::new(Mutex::new(Vec::new()));
let down_for_hook = Arc::clone(&down_reasons);
connections.register_connection_down(move |event| {
down_for_hook
.lock()
.unwrap_or_else(|error| error.into_inner())
.push(event.reason);
});
let wedged_listener = TcpListener::bind("127.0.0.1:0").await.expect("bind wedged");
let wedged_addr = wedged_listener.local_addr().expect("wedged addr");
let wedged_node = atom_table.intern("wedged@127.0.0.1");
let wedged_stream = std::net::TcpStream::connect(wedged_addr).expect("wedged connects");
let wedged_peer_addr = wedged_stream.peer_addr().expect("wedged peer addr");
let wedged_accept = tokio::spawn(async move { wedged_listener.accept().await });
let wedged_connection = connections
.register_test_connection(wedged_node, wedged_peer_addr, wedged_stream)
.expect("register wedged connection");
let _wedged_held = wedged_accept
.await
.expect("wedged accept join")
.expect("wedged accepted");
let sender = DistSender::new(connections.clone()).expect("sender builds");
let mut big = vec![0u8; 16 * 1024 * 1024];
big[0] = 0x01;
let big_control_len = u32::try_from(big.len()).expect("control fits u32");
let mut big_frame = Vec::with_capacity(8 + big.len());
big_frame.extend_from_slice(&big_control_len.to_be_bytes());
big_frame.extend_from_slice(&0u32.to_be_bytes());
big_frame.extend_from_slice(&big);
sender
.enqueue_control(ControlOutbound {
connection: Arc::clone(&wedged_connection),
frame: Arc::from(big_frame.into_boxed_slice()),
})
.expect("first control accepted into an empty lane");
let start = std::time::Instant::now();
let mut overflowed = 0usize;
for index in 0..(DIST_CONTROL_QUEUE_CAP * 4) {
match sender.enqueue_control(ControlOutbound {
connection: Arc::clone(&wedged_connection),
frame: framed(&index.to_be_bytes()),
}) {
Ok(()) => {}
Err(ControlEnqueueError::Overflow) => {
assert!(
wedged_connection.is_down(),
"Overflow must mark the pinned connection down before returning"
);
overflowed += 1;
}
Err(ControlEnqueueError::Closed) => {
panic!("control lane must not close while the sender is live")
}
}
}
let elapsed = start.elapsed();
assert!(
overflowed > 0,
"flooding 4x capacity behind a wedged write must overflow the lane"
);
assert!(
elapsed < WRITE_TIMEOUT,
"enqueue_control must be non-blocking; flood took {elapsed:?}"
);
let deadline = std::time::Instant::now() + Duration::from_secs(10);
while down_reasons
.lock()
.unwrap_or_else(|error| error.into_inner())
.is_empty()
{
assert!(
std::time::Instant::now() < deadline,
"overflow down-hook never fired"
);
tokio::time::sleep(Duration::from_millis(10)).await;
}
let reasons = down_reasons
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone();
assert_eq!(
reasons.len(),
1,
"down-hook must fire exactly once, got {reasons:?}"
);
assert!(
matches!(
reasons[0],
ConnectionDownReason::ControlOverflow | ConnectionDownReason::WriteTimeout
),
"down reason must be a DC-1(b) arm, got {:?}",
reasons[0]
);
assert!(wedged_connection.is_down());
assert!(
connections.get_connection(wedged_node).is_none(),
"wedged connection must be purged from the table"
);
sender.shutdown();
drop(sender);
}
#[tokio::test]
async fn control_frame_pinned_to_down_generation_never_reaches_redialed_socket() {
let (connections, atom_table) = manager();
let node = atom_table.intern("peer@127.0.0.1");
let listener_g = TcpListener::bind("127.0.0.1:0").await.expect("bind G");
let addr_g = listener_g.local_addr().expect("G addr");
let stream_g = std::net::TcpStream::connect(addr_g).expect("G connects");
let peer_addr_g = stream_g.peer_addr().expect("G peer addr");
let accept_g = tokio::spawn(async move { listener_g.accept().await });
let pinned = connections
.register_test_connection(node, peer_addr_g, stream_g)
.expect("register generation G");
let _held_g = accept_g.await.expect("G accept join").expect("G accepted");
pinned.mark_down_write_timeout();
assert!(pinned.is_down());
assert!(
connections.get_connection(node).is_none(),
"downed generation must leave the table before the redial"
);
let listener_new = TcpListener::bind("127.0.0.1:0").await.expect("bind new");
let addr_new = listener_new.local_addr().expect("new addr");
let received = Arc::new(Mutex::new(Vec::new()));
let received_for_task = Arc::clone(&received);
let reader = tokio::spawn(async move {
let (mut stream, _) = listener_new.accept().await.expect("new accept");
let mut header = [0u8; 8];
if stream.read_exact(&mut header).await.is_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 mut body = vec![0u8; control_len + payload_len];
if stream.read_exact(&mut body).await.is_ok() {
received_for_task
.lock()
.unwrap_or_else(|error| error.into_inner())
.push(body[0]);
}
}
});
let stream_new = std::net::TcpStream::connect(addr_new).expect("new connects");
let peer_addr_new = stream_new.peer_addr().expect("new peer addr");
connections
.register_test_connection(node, peer_addr_new, stream_new)
.expect("register redialed connection");
let sender = DistSender::new(connections.clone()).expect("sender builds");
sender
.enqueue_control(ControlOutbound {
connection: Arc::clone(&pinned),
frame: framed(&[0xCC]),
})
.expect("enqueue accepts; the is_down skip happens at the drain");
sender.enqueue(DistOutbound::ToNode {
node,
frame: framed(&[0xDD]),
});
reader.await.expect("reader task joins");
let got = received
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone();
assert_eq!(
got,
vec![0xDD],
"first frame on the redialed socket must be the sentinel, never the pinned control"
);
sender.shutdown();
drop(sender);
}
#[test]
fn disconnect_all_closes_a_connection_whose_writer_is_wedged() {
use std::io::Read;
let (manager, atom_table) = manager();
let sender = DistSender::new(manager.clone()).expect("sender builds");
manager.set_runtime_handle(sender.handle());
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
let mut client = std::net::TcpStream::connect(addr).expect("connect");
let (server, _) = listener.accept().expect("accept");
let node = atom_table.intern("peer@wedged");
let connection = {
let handle = sender.handle();
let _context = handle.enter();
manager
.register_test_connection(node, addr, server)
.expect("register test connection")
};
let (write_tx, write_rx) = std::sync::mpsc::channel();
let write_connection = Arc::clone(&connection);
sender.handle().spawn(async move {
let payload = vec![0u8; 16 * 1024 * 1024];
let result = write_connection.write_raw(&payload).await;
let _ = write_tx.send(result.is_err());
});
std::thread::sleep(Duration::from_millis(300));
manager.disconnect_all();
assert!(
write_rx
.recv_timeout(Duration::from_secs(10))
.expect("the wedged write must finish once the socket is shut down"),
"the wedged write reports an error after socket shutdown"
);
assert!(manager.connected_nodes().is_empty());
client
.set_read_timeout(Some(Duration::from_secs(10)))
.expect("set read timeout");
let mut buffer = [0u8; 64];
loop {
match client.read(&mut buffer) {
Ok(0) => break,
Ok(_) => continue,
Err(error) if error.kind() == std::io::ErrorKind::ConnectionReset => break,
Err(error) => panic!("peer expected EOF, read failed: {error}"),
}
}
sender.shutdown();
}
#[test]
fn worker_side_inventory_during_shutdown_does_not_deadlock() {
let (manager, _atom_table) = manager();
let sender = DistSender::new(manager).expect("sender builds");
let probe = sender.clone();
let (started_tx, started_rx) = std::sync::mpsc::channel();
let (names_tx, names_rx) = std::sync::mpsc::channel();
sender.handle().spawn(async move {
let _ = started_tx.send(());
std::thread::sleep(Duration::from_millis(100));
let _ = names_tx.send(probe.worker_thread_names());
});
started_rx
.recv_timeout(Duration::from_secs(5))
.expect("worker task starts");
let (done_tx, done_rx) = std::sync::mpsc::channel();
let owner = sender.clone();
let shutdown_thread = std::thread::spawn(move || {
owner.shutdown();
let _ = done_tx.send(());
});
done_rx
.recv_timeout(Duration::from_secs(10))
.expect("shutdown must not deadlock against a worker-side inventory read");
let _ = shutdown_thread.join();
names_rx
.recv_timeout(Duration::from_secs(5))
.expect("the worker-side inventory read completes");
}
#[test]
fn shutdown_from_an_unrelated_runtime_still_joins_the_worker() {
let (manager, _atom_table) = manager();
let sender = DistSender::new(manager).expect("sender builds");
let (busy_tx, busy_rx) = std::sync::mpsc::channel();
sender.handle().spawn(async move {
let _ = busy_tx.send(());
std::thread::sleep(Duration::from_millis(1200));
});
busy_rx
.recv_timeout(Duration::from_secs(5))
.expect("busy task starts");
let unrelated = tokio::runtime::Builder::new_multi_thread()
.worker_threads(1)
.enable_all()
.build()
.expect("unrelated runtime builds");
let for_shutdown = sender.clone();
let started = std::time::Instant::now();
unrelated.block_on(async move {
for_shutdown.shutdown();
});
assert!(
started.elapsed() >= Duration::from_millis(800),
"shutdown from an unrelated runtime must JOIN the busy worker \
(returned in {:?} — the background fallback)",
started.elapsed()
);
drop(unrelated);
}
#[test]
fn teardown_dup_is_cloexec_and_released_at_mark_down() {
let (manager, atom_table) = manager();
let sender = DistSender::new(manager.clone()).expect("sender builds");
manager.set_runtime_handle(sender.handle());
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
let _client = std::net::TcpStream::connect(addr).expect("connect");
let (server, _) = listener.accept().expect("accept");
let node = atom_table.intern("peer@cloexec");
let connection = {
let handle = sender.handle();
let _context = handle.enter();
manager
.register_test_connection(node, addr, server)
.expect("register test connection")
};
assert_eq!(
connection.teardown_fd_cloexec(),
Some(true),
"the teardown dup is created atomically CLOEXEC"
);
manager.disconnect_node(node);
assert_eq!(
connection.teardown_fd_cloexec(),
None,
"mark_down releases the dup even while this Arc is retained"
);
sender.shutdown();
}
#[test]
fn teardown_dup_failure_refuses_the_connection() {
let (manager, atom_table) = manager();
let sender = DistSender::new(manager.clone()).expect("sender builds");
manager.set_runtime_handle(sender.handle());
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
let _client = std::net::TcpStream::connect(addr).expect("connect");
let (server, _) = listener.accept().expect("accept");
let node = atom_table.intern("peer@dupfail");
crate::distribution::connection::FAIL_TEARDOWN_DUP_FOR_TEST.with(|flag| flag.set(true));
let refused = {
let handle = sender.handle();
let _context = handle.enter();
manager.register_test_connection(node, addr, server)
};
crate::distribution::connection::FAIL_TEARDOWN_DUP_FOR_TEST.with(|flag| flag.set(false));
assert!(refused.is_err(), "dup failure refuses the install");
assert!(
manager.get_connection(node).is_none(),
"a refused connection is never tabled"
);
sender.shutdown();
}
#[test]
fn cross_scheduler_same_named_worker_still_joins_the_other_runtime() {
let (manager_a, _t1) = manager();
let sender_a = DistSender::new(manager_a).expect("sender A builds");
let (manager_b, _t2) = manager();
let sender_b = DistSender::new(manager_b).expect("sender B builds");
let (busy_tx, busy_rx) = std::sync::mpsc::channel();
sender_a.handle().spawn(async move {
let _ = busy_tx.send(());
std::thread::sleep(Duration::from_millis(1200));
});
busy_rx
.recv_timeout(Duration::from_secs(5))
.expect("A's busy task starts");
let (elapsed_tx, elapsed_rx) = std::sync::mpsc::channel();
let a_for_b = sender_a.clone();
sender_b.handle().spawn(async move {
let started = std::time::Instant::now();
a_for_b.shutdown();
let _ = elapsed_tx.send(started.elapsed());
});
let elapsed = elapsed_rx
.recv_timeout(Duration::from_secs(10))
.expect("A's shutdown from B's worker completes");
assert!(
elapsed >= Duration::from_millis(800),
"shutdown from another scheduler's same-named worker must JOIN \
(returned in {elapsed:?} — the self-runtime background fallback)"
);
sender_b.shutdown();
}
#[test]
fn shutdown_from_the_senders_own_blocking_pool_does_not_deadlock() {
let (manager, _atom_table) = manager();
let sender = DistSender::new(manager).expect("sender builds");
let (done_tx, done_rx) = std::sync::mpsc::channel();
let own = sender.clone();
sender.handle().spawn_blocking(move || {
own.shutdown();
let _ = done_tx.send(());
});
done_rx
.recv_timeout(Duration::from_secs(10))
.expect("shutdown on the sender's own blocking pool must complete, not deadlock");
}
#[test]
fn shutdown_from_the_senders_own_runtime_worker_does_not_deadlock() {
let (manager, _table) = manager();
let sender = DistSender::new(manager).unwrap_or_else(|| panic!("sender builds"));
let (done_tx, done_rx) = std::sync::mpsc::channel();
let on_runtime = sender.clone();
sender.handle().spawn(async move {
on_runtime.shutdown();
let _ = done_tx.send(());
});
done_rx
.recv_timeout(Duration::from_secs(10))
.expect("shutdown on the sender's own worker must complete, not deadlock");
sender.shutdown();
assert!(sender.worker_thread_names().is_empty());
}
#[test]
fn final_clone_drop_on_the_senders_own_runtime_worker_does_not_deadlock() {
let (manager, _table) = manager();
let sender = DistSender::new(manager).unwrap_or_else(|| panic!("sender builds"));
let handle = sender.handle().clone();
let (done_tx, done_rx) = std::sync::mpsc::channel();
handle.spawn(async move {
drop(sender);
let _ = done_tx.send(());
});
done_rx
.recv_timeout(Duration::from_secs(10))
.expect("final-clone drop on the sender's own worker must complete, not deadlock");
}
}