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;
pub const DIST_SEND_QUEUE_CAP: usize = 1024;
const WRITE_TIMEOUT: Duration = Duration::from_secs(5);
#[derive(Clone, Debug)]
pub enum DistOutbound {
ToNode {
node: Atom,
frame: Arc<[u8]>,
},
}
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>,
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 drain = runtime.spawn(async move {
while let Some(DistOutbound::ToNode { node, frame }) = rx.recv().await {
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();
}
}
}
});
let handle = runtime.handle().clone();
Some(Self {
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 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::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);
}
}