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
52 changes: 41 additions & 11 deletions crates/mqdb-cluster/src/cluster/quic_transport.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,11 +15,13 @@ use tokio::sync::{Notify, RwLock};
use tracing::{debug, error, info, trace, warn};

const INBOX_CHANNEL_CAPACITY: usize = 16384;
const PEER_SEND_QUEUE_CAPACITY: usize = 1024;
const PEER_CONTROL_QUEUE_CAPACITY: usize = 256;
const PEER_BULK_QUEUE_CAPACITY: usize = 1024;

struct PeerConnection {
_connection: Connection,
writer_tx: flume::Sender<Vec<u8>>,
control_tx: flume::Sender<Vec<u8>>,
bulk_tx: flume::Sender<Vec<u8>>,
}

const MAX_MESSAGE_SIZE: usize = 10 * 1024 * 1024;
Expand Down Expand Up @@ -246,14 +248,16 @@ impl QuicDirectTransport {
.await
.map_err(|e| TransportError::SendFailed(format!("failed to send header: {e}")))?;

let (writer_tx, writer_rx) = flume::bounded(PEER_SEND_QUEUE_CAPACITY);
let (control_tx, control_rx) = flume::bounded(PEER_CONTROL_QUEUE_CAPACITY);
let (bulk_tx, bulk_rx) = flume::bounded(PEER_BULK_QUEUE_CAPACITY);
let peer_conn = PeerConnection {
_connection: connection.clone(),
writer_tx,
control_tx,
bulk_tx,
};
self.peers.write().await.insert(peer_id, peer_conn);

tokio::spawn(peer_writer_task(send_stream, writer_rx, peer_id));
tokio::spawn(peer_writer_task(send_stream, control_rx, bulk_rx, peer_id));

let inbox_tx = self.inbox_tx.clone();
let notify = self.message_notify.clone();
Expand Down Expand Up @@ -306,7 +310,13 @@ impl QuicDirectTransport {
.get(&peer_id)
.ok_or(TransportError::NodeNotFound(peer_id))?;

match peer.writer_tx.try_send(frame) {
let lane = if message.is_control_plane() {
&peer.control_tx
} else {
&peer.bulk_tx
};

match lane.try_send(frame) {
Ok(()) => {
trace!(
from = self.node_id.get(),
Expand Down Expand Up @@ -499,29 +509,49 @@ async fn handle_incoming_connection(

info!(peer = peer_node.get(), "accepted incoming QUIC connection");

let (writer_tx, writer_rx) = flume::bounded(PEER_SEND_QUEUE_CAPACITY);
let (control_tx, control_rx) = flume::bounded(PEER_CONTROL_QUEUE_CAPACITY);
let (bulk_tx, bulk_rx) = flume::bounded(PEER_BULK_QUEUE_CAPACITY);
{
let peer_conn = PeerConnection {
_connection: connection.clone(),
writer_tx,
control_tx,
bulk_tx,
};
peers.write().await.insert(peer_node, peer_conn);
}

tokio::spawn(peer_writer_task(send_stream, writer_rx, peer_node));
tokio::spawn(peer_writer_task(
send_stream,
control_rx,
bulk_rx,
peer_node,
));

receiver_task(recv_stream, peer_node, inbox_tx, notify, local_node).await;
Ok(())
}

async fn peer_writer_task(
mut send_stream: SendStream,
writer_rx: flume::Receiver<Vec<u8>>,
control_rx: flume::Receiver<Vec<u8>>,
bulk_rx: flume::Receiver<Vec<u8>>,
peer_node: NodeId,
) {
trace!(peer = peer_node.get(), "peer writer task started");

while let Ok(frame) = writer_rx.recv_async().await {
loop {
let frame = tokio::select! {
biased;
control = control_rx.recv_async() => match control {
Ok(frame) => frame,
Err(_) => break,
},
bulk = bulk_rx.recv_async() => match bulk {
Ok(frame) => frame,
Err(_) => break,
},
};

if let Err(e) = send_stream.write_all(&frame).await {
warn!(peer = peer_node.get(), error = %e, "peer writer failed, tearing down stream");
break;
Expand Down
15 changes: 15 additions & 0 deletions crates/mqdb-cluster/src/cluster/transport.rs
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,21 @@ pub enum ClusterMessage {
}

impl ClusterMessage {
#[must_use]
pub fn is_control_plane(&self) -> bool {
matches!(
self,
Self::Heartbeat(_)
| Self::DeathNotice { .. }
| Self::DrainNotification { .. }
| Self::RequestVote(_)
| Self::RequestVoteResponse(_)
| Self::AppendEntries(_)
| Self::AppendEntriesResponse(_)
| Self::PartitionUpdate(_)
)
}

#[must_use]
pub fn message_type(&self) -> u8 {
match self {
Expand Down
136 changes: 128 additions & 8 deletions crates/mqdb-cluster/tests/quic_transport_framing.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,8 @@ use std::sync::Arc;
use std::time::Duration;

use mqdb_cluster::{
ClusterMessage, ClusterTransport, Heartbeat, NodeId, QuicDirectTransport, TransportError,
ClusterMessage, ClusterTransport, ForwardedPublish, Heartbeat, NodeId, QuicDirectTransport,
TransportError,
};
use quinn::{Connection, Endpoint, RecvStream, SendStream, ServerConfig, TransportConfig, VarInt};
use rcgen::{BasicConstraints, CertificateParams, IsCa, Issuer, KeyPair};
Expand Down Expand Up @@ -99,8 +100,16 @@ async fn collect_frame(frame_rx: &flume::Receiver<Vec<u8>>) -> Vec<u8> {
.expect("far-end frame channel closed")
}

#[tokio::test]
async fn backpressure_drops_messages_but_keeps_the_stream_framed() {
struct StalledPeer {
transport: QuicDirectTransport,
local: NodeId,
peer: NodeId,
frame_rx: flume::Receiver<Vec<u8>>,
release_tx: oneshot::Sender<()>,
far_handle: tokio::task::JoinHandle<()>,
}

async fn setup_stalled_peer() -> StalledPeer {
let _ = rustls::crypto::ring::default_provider().install_default();

let certs = generate_certs();
Expand All @@ -122,19 +131,58 @@ async fn backpressure_drops_messages_but_keeps_the_stream_framed() {
let local = NodeId::validated(1).unwrap();
let peer = NodeId::validated(2).unwrap();
let transport = QuicDirectTransport::new(local);
transport.set_ca_file(ca_path.clone());
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();

StalledPeer {
transport,
local,
peer,
frame_rx,
release_tx,
far_handle,
}
}

fn heartbeat(local: NodeId, tick: u64) -> ClusterMessage {
ClusterMessage::Heartbeat(Heartbeat::create(local, tick))
}

fn bulk_message(local: NodeId) -> ClusterMessage {
ClusterMessage::ForwardedPublish(ForwardedPublish::new(
local,
"bulk".to_string(),
0,
false,
vec![0u8; 64],
Vec::new(),
))
}

#[tokio::test]
async fn backpressure_drops_messages_but_keeps_the_stream_framed() {
let StalledPeer {
transport,
local,
peer,
frame_rx,
release_tx,
far_handle,
} = setup_stalled_peer().await;

let mut ok_count = 0u32;
let mut terminated_by_drop = false;
let mut blocked = false;
for tick in 0..SEND_ATTEMPTS {
let message = ClusterMessage::Heartbeat(Heartbeat::create(local, u64::from(tick)));
match tokio::time::timeout(Duration::from_millis(500), transport.send(peer, message)).await
match tokio::time::timeout(
Duration::from_millis(500),
transport.send(peer, heartbeat(local, u64::from(tick))),
)
.await
{
Err(_) => {
blocked = true;
Expand Down Expand Up @@ -177,8 +225,10 @@ async fn backpressure_drops_messages_but_keeps_the_stream_framed() {
);
}

let follow_up = ClusterMessage::Heartbeat(Heartbeat::create(local, u64::MAX));
transport.send(peer, follow_up).await.unwrap();
transport
.send(peer, heartbeat(local, u64::MAX))
.await
.unwrap();
let frame = collect_frame(&frame_rx).await;
assert!(frame.len() >= 3);
assert_eq!(
Expand All @@ -189,3 +239,73 @@ async fn backpressure_drops_messages_but_keeps_the_stream_framed() {

far_handle.abort();
}

#[tokio::test]
async fn control_plane_survives_a_full_bulk_queue() {
let StalledPeer {
transport,
local,
peer,
frame_rx,
release_tx,
far_handle,
} = setup_stalled_peer().await;

let mut bulk_ok = 0u32;
let mut bulk_dropped = false;
for _ in 0..SEND_ATTEMPTS {
match tokio::time::timeout(
Duration::from_millis(500),
transport.send(peer, bulk_message(local)),
)
.await
{
Err(_) => panic!("bulk send blocked under backpressure"),
Ok(Ok(())) => bulk_ok += 1,
Ok(Err(TransportError::SendQueueFull(_))) => {
bulk_dropped = true;
break;
}
Ok(Err(other)) => panic!("unexpected bulk send error: {other}"),
}
}
assert!(bulk_dropped, "bulk lane should saturate and start dropping");

let result = tokio::time::timeout(
Duration::from_millis(500),
transport.send(peer, heartbeat(local, 7)),
)
.await
.expect("control-plane send blocked under backpressure");
assert!(
matches!(result, Ok(())),
"control-plane heartbeat was dropped while the bulk lane was full: {result:?}"
);

release_tx.send(()).ok();

let heartbeat_type = heartbeat(local, 0).message_type();
let mut heartbeat_delivered = false;
for _ in 0..bulk_ok + 2 {
let frame = collect_frame(&frame_rx).await;
assert!(
frame.len() >= 3,
"frame shorter than a cluster message header"
);
assert_eq!(
&frame[0..2],
&local.get().to_be_bytes(),
"stream mis-framed under the two-lane writer"
);
if frame[2] == heartbeat_type {
heartbeat_delivered = true;
break;
}
}
assert!(
heartbeat_delivered,
"control-plane heartbeat was accepted but never delivered"
);

far_handle.abort();
}
Loading