Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
81 changes: 75 additions & 6 deletions crates/mqdb-cluster/src/cluster/quic_transport.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@ use std::collections::{HashMap, VecDeque};
use std::io::BufReader;
use std::net::SocketAddr;
use std::path::Path;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use tokio::sync::{Notify, RwLock};
use tracing::{debug, error, info, trace, warn};
Expand All @@ -22,6 +22,7 @@ struct PeerConnection {
_connection: Connection,
control_tx: flume::Sender<Vec<u8>>,
bulk_tx: flume::Sender<Vec<u8>>,
generation: u64,
}

const MAX_MESSAGE_SIZE: usize = 10 * 1024 * 1024;
Expand All @@ -38,6 +39,8 @@ pub struct QuicDirectTransport {
node_id: NodeId,
endpoint: Arc<RwLock<Option<Endpoint>>>,
peers: Arc<RwLock<HashMap<NodeId, PeerConnection>>>,
peer_addrs: Arc<RwLock<HashMap<NodeId, SocketAddr>>>,
generation: Arc<AtomicU64>,
inbox_tx: flume::Sender<InboundMessage>,
inbox_rx: flume::Receiver<InboundMessage>,
requeue_buffer: Arc<Mutex<VecDeque<InboundMessage>>>,
Expand Down Expand Up @@ -67,6 +70,8 @@ impl Clone for QuicDirectTransport {
node_id: self.node_id,
endpoint: self.endpoint.clone(),
peers: self.peers.clone(),
peer_addrs: self.peer_addrs.clone(),
generation: self.generation.clone(),
inbox_tx: self.inbox_tx.clone(),
inbox_rx: self.inbox_rx.clone(),
requeue_buffer: self.requeue_buffer.clone(),
Expand Down Expand Up @@ -94,6 +99,8 @@ impl QuicDirectTransport {
node_id,
endpoint: Arc::new(RwLock::new(None)),
peers: Arc::new(RwLock::new(HashMap::new())),
peer_addrs: Arc::new(RwLock::new(HashMap::new())),
generation: Arc::new(AtomicU64::new(0)),
inbox_tx,
inbox_rx,
requeue_buffer: Arc::new(Mutex::new(VecDeque::new())),
Expand Down Expand Up @@ -175,9 +182,10 @@ impl QuicDirectTransport {
let notify = self.message_notify.clone();
let local_node = self.node_id;
let peers = self.peers.clone();
let generation = self.generation.clone();

tokio::spawn(async move {
acceptor_task(endpoint, inbox_tx, notify, local_node, peers).await;
acceptor_task(endpoint, inbox_tx, notify, local_node, peers, generation).await;
});

Ok(())
Expand All @@ -192,6 +200,8 @@ impl QuicDirectTransport {
peer_id: NodeId,
peer_addr: SocketAddr,
) -> Result<(), TransportError> {
self.peer_addrs.write().await.insert(peer_id, peer_addr);

let endpoint_guard = self.endpoint.read().await;
let endpoint = endpoint_guard
.as_ref()
Expand Down Expand Up @@ -250,14 +260,23 @@ impl QuicDirectTransport {

let (control_tx, control_rx) = flume::bounded(PEER_CONTROL_QUEUE_CAPACITY);
let (bulk_tx, bulk_rx) = flume::bounded(PEER_BULK_QUEUE_CAPACITY);
let generation = self.generation.fetch_add(1, Ordering::SeqCst) + 1;
let peer_conn = PeerConnection {
_connection: connection.clone(),
control_tx,
bulk_tx,
generation,
};
self.peers.write().await.insert(peer_id, peer_conn);

tokio::spawn(peer_writer_task(send_stream, control_rx, bulk_rx, peer_id));
tokio::spawn(peer_writer_task(
send_stream,
control_rx,
bulk_rx,
peer_id,
self.peers.clone(),
generation,
));

let inbox_tx = self.inbox_tx.clone();
let notify = self.message_notify.clone();
Expand All @@ -270,6 +289,30 @@ impl QuicDirectTransport {
Ok(())
}

/// Re-dial every configured peer that is not currently alive.
///
/// Driven by the heartbeat liveness view rather than the transport peer
/// map, so a stale entry that has not yet been removed by a send failure
/// does not stop reconnection. `connect_to_peer` replaces any stale entry.
pub async fn redial_unlinked(&self, alive: &[NodeId]) {
let targets: Vec<(NodeId, SocketAddr)> = {
let addrs = self.peer_addrs.read().await;
addrs
.iter()
.filter(|&(node, _)| *node != self.node_id && !alive.contains(node))
.map(|(node, addr)| (*node, *addr))
.collect()
};

for (node, addr) in targets {
if let Err(e) = self.connect_to_peer(node, addr).await {
debug!(peer = node.get(), addr = %addr, error = %e, "redial to peer failed, will retry");
} else {
info!(peer = node.get(), addr = %addr, "redialled peer");
}
}
}

#[must_use]
pub fn inbox_rx(&self) -> flume::Receiver<InboundMessage> {
self.inbox_rx.clone()
Expand Down Expand Up @@ -459,6 +502,7 @@ async fn acceptor_task(
notify: Arc<Notify>,
local_node: NodeId,
peers: Arc<RwLock<HashMap<NodeId, PeerConnection>>>,
generation: Arc<AtomicU64>,
) {
info!(node = local_node.get(), "QUIC acceptor task started");

Expand All @@ -474,10 +518,13 @@ async fn acceptor_task(
let inbox_tx = inbox_tx.clone();
let notify = notify.clone();
let peers = peers.clone();
let generation = generation.clone();

tokio::spawn(async move {
if let Err(e) =
handle_incoming_connection(connection, inbox_tx, notify, local_node, peers).await
if let Err(e) = handle_incoming_connection(
connection, inbox_tx, notify, local_node, peers, generation,
)
.await
{
debug!(error = %e, "incoming connection handler failed");
}
Expand All @@ -491,6 +538,7 @@ async fn handle_incoming_connection(
notify: Arc<Notify>,
local_node: NodeId,
peers: Arc<RwLock<HashMap<NodeId, PeerConnection>>>,
generation: Arc<AtomicU64>,
) -> Result<(), TransportError> {
let (send_stream, mut recv_stream) = connection
.accept_bi()
Expand All @@ -511,11 +559,13 @@ async fn handle_incoming_connection(

let (control_tx, control_rx) = flume::bounded(PEER_CONTROL_QUEUE_CAPACITY);
let (bulk_tx, bulk_rx) = flume::bounded(PEER_BULK_QUEUE_CAPACITY);
let peer_generation = generation.fetch_add(1, Ordering::SeqCst) + 1;
{
let peer_conn = PeerConnection {
_connection: connection.clone(),
control_tx,
bulk_tx,
generation: peer_generation,
};
peers.write().await.insert(peer_node, peer_conn);
}
Expand All @@ -525,6 +575,8 @@ async fn handle_incoming_connection(
control_rx,
bulk_rx,
peer_node,
peers.clone(),
peer_generation,
));

receiver_task(recv_stream, peer_node, inbox_tx, notify, local_node).await;
Expand All @@ -536,6 +588,8 @@ async fn peer_writer_task(
control_rx: flume::Receiver<Vec<u8>>,
bulk_rx: flume::Receiver<Vec<u8>>,
peer_node: NodeId,
peers: Arc<RwLock<HashMap<NodeId, PeerConnection>>>,
generation: u64,
) {
trace!(peer = peer_node.get(), "peer writer task started");

Expand All @@ -553,14 +607,29 @@ async fn peer_writer_task(
};

if let Err(e) = send_stream.write_all(&frame).await {
warn!(peer = peer_node.get(), error = %e, "peer writer failed, tearing down stream");
warn!(peer = peer_node.get(), error = %e, "peer writer failed, removing dead peer");
remove_peer_generation(&peers, peer_node, generation).await;
break;
}
}

debug!(peer = peer_node.get(), "peer writer task ended");
}

async fn remove_peer_generation(
peers: &Arc<RwLock<HashMap<NodeId, PeerConnection>>>,
peer_node: NodeId,
generation: u64,
) {
let mut map = peers.write().await;
if map
.get(&peer_node)
.is_some_and(|peer| peer.generation == generation)
{
map.remove(&peer_node);
}
}

async fn receiver_task(
mut recv_stream: RecvStream,
peer_node: NodeId,
Expand Down
13 changes: 13 additions & 0 deletions crates/mqdb-cluster/src/cluster_agent/event_loop.rs
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,7 @@ impl ClusteredAgent {
}
_ = mesh_check_interval.tick() => {
self.warn_on_unlinked_nodes().await;
self.redial_disconnected_peers().await;
}
_ = retained_sync_cleanup_interval.tick() => {
Self::handle_retained_sync_cleanup(&synced_retained_topics).await;
Expand Down Expand Up @@ -557,6 +558,18 @@ impl ClusteredAgent {
}
}

async fn redial_disconnected_peers(&self) {
let (quic, alive) = {
let ctrl = self.controller.read().await;
(ctrl.transport().as_quic().cloned(), ctrl.alive_nodes())
};
if let Some(quic) = quic {
tokio::spawn(async move {
quic.redial_unlinked(&alive).await;
});
}
}

async fn warn_on_unlinked_nodes(&self) {
let unlinked = {
let ctrl = self.controller.read().await;
Expand Down
99 changes: 96 additions & 3 deletions crates/mqdb-cluster/tests/quic_transport_framing.rs
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ fn generate_certs() -> Certs {
}
}

fn stalling_server_endpoint(certs: &Certs) -> Endpoint {
fn server_endpoint(certs: &Certs, stream_window: u32) -> Endpoint {
let key = PrivateKeyDer::try_from(certs.leaf_key_der.clone()).unwrap();
let server_crypto = rustls::ServerConfig::builder()
.with_no_client_auth()
Expand All @@ -55,7 +55,7 @@ fn stalling_server_endpoint(certs: &Certs) -> Endpoint {
));

let mut transport = TransportConfig::default();
transport.stream_receive_window(VarInt::from_u32(STALL_WINDOW));
transport.stream_receive_window(VarInt::from_u32(stream_window));
server_config.transport_config(Arc::new(transport));

Endpoint::server(server_config, "127.0.0.1:0".parse().unwrap()).unwrap()
Expand Down Expand Up @@ -121,7 +121,7 @@ async fn setup_stalled_peer() -> StalledPeer {
std::fs::write(&leaf_path, &certs.leaf_pem).unwrap();
std::fs::write(&leaf_key_path, &certs.leaf_key_pem).unwrap();

let endpoint = stalling_server_endpoint(&certs);
let endpoint = server_endpoint(&certs, STALL_WINDOW);
let far_addr: SocketAddr = endpoint.local_addr().unwrap();

let (release_tx, release_rx) = oneshot::channel();
Expand Down Expand Up @@ -309,3 +309,96 @@ async fn control_plane_survives_a_full_bulk_queue() {

far_handle.abort();
}

async fn reconnecting_far_end(
endpoint: Endpoint,
close_first: oneshot::Receiver<()>,
accepted_tx: flume::Sender<()>,
) {
let conn1: Connection = endpoint.accept().await.unwrap().await.unwrap();
let (_s1, mut r1): (SendStream, RecvStream) = conn1.accept_bi().await.unwrap();
let mut header = [0u8; 2];
r1.read_exact(&mut header).await.unwrap();
accepted_tx.send(()).ok();

close_first.await.ok();
conn1.close(0u32.into(), b"test-close");
drop(r1);
drop(conn1);

let conn2: Connection = endpoint.accept().await.unwrap().await.unwrap();
let (_s2, mut r2): (SendStream, RecvStream) = conn2.accept_bi().await.unwrap();
let mut header2 = [0u8; 2];
r2.read_exact(&mut header2).await.unwrap();
accepted_tx.send(()).ok();

std::future::pending::<()>().await;
drop(conn2);
}

#[tokio::test]
async fn dead_peer_is_removed_and_redialled() {
let _ = rustls::crypto::ring::default_provider().install_default();

let certs = generate_certs();
let dir = tempfile::tempdir().unwrap();
let ca_path = dir.path().join("ca.pem");
let leaf_path = dir.path().join("leaf.pem");
let leaf_key_path = dir.path().join("leaf.key");
std::fs::write(&ca_path, &certs.ca_pem).unwrap();
std::fs::write(&leaf_path, &certs.leaf_pem).unwrap();
std::fs::write(&leaf_key_path, &certs.leaf_key_pem).unwrap();

let endpoint = server_endpoint(&certs, 1024 * 1024);
let far_addr: SocketAddr = endpoint.local_addr().unwrap();

let (close_tx, close_rx) = oneshot::channel();
let (accepted_tx, accepted_rx) = flume::unbounded();
let far_handle = tokio::spawn(reconnecting_far_end(endpoint, close_rx, accepted_tx));

let local = NodeId::validated(1).unwrap();
let peer = NodeId::validated(2).unwrap();
let transport = QuicDirectTransport::new(local);
transport.set_ca_file(ca_path);
transport
.bind("127.0.0.1:0".parse().unwrap(), &leaf_path, &leaf_key_path)
.await
.unwrap();
transport.connect_to_peer(peer, far_addr).await.unwrap();

tokio::time::timeout(Duration::from_secs(10), accepted_rx.recv_async())
.await
.expect("far end did not accept the initial connection")
.unwrap();
assert!(
transport.direct_peers().await.unwrap().contains(&peer),
"peer should be linked after the initial connect"
);

close_tx.send(()).ok();
let mut removed = false;
for _ in 0..200 {
let _ = transport.send(peer, heartbeat(local, 1)).await;
if !transport.direct_peers().await.unwrap().contains(&peer) {
removed = true;
break;
}
tokio::time::sleep(Duration::from_millis(25)).await;
}
assert!(
removed,
"a send-side failure should remove the dead peer from the map"
);

transport.redial_unlinked(&[]).await;
tokio::time::timeout(Duration::from_secs(10), accepted_rx.recv_async())
.await
.expect("far end did not accept the redial")
.unwrap();
assert!(
transport.direct_peers().await.unwrap().contains(&peer),
"peer should be re-linked after redial_disconnected"
);

far_handle.abort();
}
16 changes: 16 additions & 0 deletions specs/ClusterRedial.cfg
Original file line number Diff line number Diff line change
@@ -0,0 +1,16 @@
SPECIFICATION Spec
CONSTANTS
Nodes = {1, 2}
MaxGen = 5
MaxBreaks = 2
Guarded = TRUE
MonotonicId = TRUE
Redial = TRUE
RedialOverLive = FALSE
INVARIANTS
TypeOK
InvNoLiveDrop
PROPERTIES
Converge

CHECK_DEADLOCK FALSE
Loading
Loading