use std::sync::Arc;
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};
pub const DIST_SEND_QUEUE_CAP: usize = 1024;
pub const DIST_CONTROL_QUEUE_CAP: usize = 256;
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: Option<Runtime>,
handle: Handle,
drain: JoinHandle<()>,
}
impl Drop for DistSenderInner {
fn drop(&mut self) {
self.drain.abort();
if let Some(runtime) = self.runtime.take() {
std::thread::spawn(move || drop(runtime));
}
}
}
#[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 runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(1)
.thread_name("beamr-dist-send")
.enable_all()
.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: Some(runtime),
handle,
drain,
}),
})
}
#[must_use]
pub fn handle(&self) -> Handle {
self.inner.handle.clone()
}
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();
}
}
#[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);
}
}