use super::super::tcp::framing::DEFAULT_MAX_FRAME_SIZE;
use super::*;
use parking_lot::Mutex;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::task::{Context, Poll};
use tokio_util::codec::Decoder;
#[derive(Default)]
struct RecordingSink {
data: Vec<u8>,
poll_writes: usize,
fail_at: Option<usize>,
max_per_write: Option<usize>,
cancel_on_write: Option<CancellationToken>,
live_items: Option<Arc<AtomicUsize>>,
live_at_write: Vec<usize>,
}
impl RecordingSink {
fn failing_at(idx: usize) -> Self {
Self {
fail_at: Some(idx),
..Default::default()
}
}
fn decode_frames(&self) -> Vec<(MessageType, Vec<u8>, Vec<u8>)> {
let mut codec = TcpFrameCodec::new();
let mut buf = BytesMut::from(&self.data[..]);
let mut out = Vec::new();
while let Some((t, h, p)) = codec.decode(&mut buf).expect("decode") {
out.push((t, h.to_vec(), p.to_vec()));
}
assert!(buf.is_empty(), "decoder left {} bytes behind", buf.len());
out
}
}
impl AsyncWrite for RecordingSink {
fn poll_write(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
let idx = self.poll_writes;
self.poll_writes += 1;
if let Some(live) = &self.live_items {
let n = live.load(Ordering::SeqCst);
self.live_at_write.push(n);
}
if let Some(token) = self.cancel_on_write.take() {
token.cancel();
}
if self.fail_at == Some(idx) {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"sink failure",
)));
}
let n = self.max_per_write.map_or(buf.len(), |m| m.min(buf.len()));
self.data.extend_from_slice(&buf[..n]);
Poll::Ready(Ok(n))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}
#[derive(Default)]
struct TestObserver {
flushes: Mutex<Vec<usize>>,
failures: Mutex<Vec<(WriterFailure, usize)>>,
}
impl TestObserver {
fn flushes(&self) -> Vec<usize> {
self.flushes.lock().clone()
}
fn frames_written(&self) -> usize {
self.flushes.lock().iter().sum()
}
fn failures(&self) -> Vec<(WriterFailure, usize)> {
self.failures.lock().clone()
}
}
impl WriterObserver for TestObserver {
fn on_flush(&self, frames: usize) {
self.flushes.lock().push(frames);
}
fn on_failure(&self, kind: WriterFailure, _err: &io::Error, frames: usize) {
self.failures.lock().push((kind, frames));
}
}
struct TestItem {
tag: String,
header: Vec<u8>,
payload: Vec<u8>,
terminal: bool,
errors: Arc<Mutex<Vec<String>>>,
}
struct TestToken {
tag: String,
errors: Arc<Mutex<Vec<String>>>,
}
impl Coalescable for TestItem {
type FailureToken = TestToken;
fn msg_type(&self) -> MessageType {
MessageType::Message
}
fn header(&self) -> &[u8] {
&self.header
}
fn payload(&self) -> &[u8] {
&self.payload
}
fn is_terminal(&self) -> bool {
self.terminal
}
fn into_failure_token(self) -> TestToken {
TestToken {
tag: self.tag,
errors: self.errors,
}
}
fn fail(token: TestToken, reason: &str) {
token.errors.lock().push(format!("{}: {reason}", token.tag));
}
}
struct ItemFactory {
errors: Arc<Mutex<Vec<String>>>,
}
impl ItemFactory {
fn new() -> Self {
Self {
errors: Arc::new(Mutex::new(Vec::new())),
}
}
fn item(&self, tag: &str, payload: Vec<u8>) -> TestItem {
TestItem {
tag: tag.to_string(),
header: Vec::new(),
payload,
terminal: false,
errors: Arc::clone(&self.errors),
}
}
fn terminal(&self, tag: &str, payload: Vec<u8>) -> TestItem {
TestItem {
terminal: true,
..self.item(tag, payload)
}
}
fn item_with_header(&self, tag: &str, header: Vec<u8>, payload: Vec<u8>) -> TestItem {
TestItem {
header,
..self.item(tag, payload)
}
}
fn errors(&self) -> Vec<String> {
self.errors.lock().clone()
}
fn reports_for(&self, tag: &str) -> usize {
let prefix = format!("{tag}: ");
self.errors
.lock()
.iter()
.filter(|e| e.starts_with(&prefix))
.count()
}
}
async fn run_with(items: Vec<TestItem>, sink: &mut RecordingSink, observer: &TestObserver) {
let (tx, rx) = flume::unbounded::<TestItem>();
for item in items {
tx.send(item).expect("queue item");
}
drop(tx);
run_coalescing_writer(sink, &rx, std::convert::identity, None, observer).await;
}
#[tokio::test]
async fn batch_bytes_identical_to_sequential_writes() {
let frames: Vec<(MessageType, Vec<u8>, Vec<u8>)> = vec![
(
MessageType::Message,
b"h1".to_vec(),
b"payload-one".to_vec(),
),
(MessageType::Response, Vec::new(), b"two".to_vec()),
(MessageType::Event, b"hdr3".to_vec(), Vec::new()),
(MessageType::Ack, Vec::new(), Vec::new()),
];
let mut sequential = Vec::new();
for (t, h, p) in &frames {
TcpFrameCodec::encode_frame_sync(&mut sequential, *t, h, p).unwrap();
}
let mut batch = FrameBatchBuffer::new();
for (t, h, p) in &frames {
batch.push(*t, h, p).unwrap();
}
assert_eq!(batch.frame_count(), frames.len());
let mut batched = Vec::new();
batch.flush_to(&mut batched).await.unwrap();
assert_eq!(batched, sequential, "batched bytes must match sequential");
assert_eq!(batch.frame_count(), 0, "flush resets the buffer");
}
#[tokio::test]
async fn coalesced_batch_decodes_frame_by_frame() {
let mut batch = FrameBatchBuffer::new();
for i in 0..32u8 {
batch
.push(MessageType::Message, &[], &[i; 24])
.expect("push");
}
let mut wire = RecordingSink::default();
batch.flush_to(&mut wire).await.unwrap();
let decoded = wire.decode_frames();
assert_eq!(decoded.len(), 32);
for (i, (msg_type, header, payload)) in decoded.iter().enumerate() {
assert_eq!(*msg_type, MessageType::Message);
assert!(header.is_empty());
assert_eq!(payload.as_slice(), &[i as u8; 24]);
}
}
#[tokio::test]
async fn short_writes_preserve_framing() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink {
max_per_write: Some(7),
..Default::default()
};
let items = (0..16u8)
.map(|i| factory.item(&format!("i{i}"), vec![i; 40]))
.collect();
run_with(items, &mut sink, &observer).await;
assert!(
sink.poll_writes > 1,
"the sink must have forced write_all to loop"
);
assert_eq!(observer.flushes(), vec![16], "still one logical flush");
let decoded = sink.decode_frames();
assert_eq!(decoded.len(), 16);
for (i, (_, _, payload)) in decoded.iter().enumerate() {
assert_eq!(payload.as_slice(), &[i as u8; 40]);
}
assert!(factory.errors().is_empty(), "nothing failed");
}
#[tokio::test]
async fn queued_items_coalesce_into_one_flush() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::default();
let items = (0..8u8)
.map(|i| factory.item(&format!("i{i}"), vec![i; 16]))
.collect();
run_with(items, &mut sink, &observer).await;
assert_eq!(
observer.flushes(),
vec![8],
"eight queued items must leave in one write"
);
assert_eq!(sink.decode_frames().len(), 8);
}
#[test]
fn classify_respects_byte_cap() {
let mut batch = FrameBatchBuffer::with_limits(128, 64);
assert_eq!(batch.classify(0, 10), Staging::Stage, "empty batch stages");
batch.push(MessageType::Message, &[], &[0u8; 64]).unwrap();
assert_eq!(batch.classify(0, 8), Staging::Stage);
assert_eq!(batch.classify(0, 128), Staging::FlushThenStage);
}
#[test]
fn classify_respects_frame_cap() {
let mut batch = FrameBatchBuffer::with_limits(1 << 20, 3);
for _ in 0..3 {
assert_eq!(batch.classify(0, 1), Staging::Stage);
batch.push(MessageType::Message, &[], &[0u8; 1]).unwrap();
}
assert_eq!(batch.classify(0, 1), Staging::FlushThenStage);
}
#[tokio::test]
async fn large_frames_never_enter_the_staging_buffer() {
let mut batch = FrameBatchBuffer::new();
assert_eq!(
batch.classify(0, COALESCE_THRESHOLD + 1),
Staging::WriteDirect,
"an oversized frame must not be staged even into an empty batch"
);
for _ in 0..8 {
while batch.classify(0, 4096) == Staging::Stage {
batch.push(MessageType::Message, &[], &[7u8; 4096]).unwrap();
}
let mut sink = RecordingSink::default();
batch.flush_to(&mut sink).await.unwrap();
}
assert!(
batch.capacity() <= 4 * DEFAULT_MAX_BATCH_BYTES,
"staging buffer grew to {} bytes; it should stay near the {}-byte batch cap",
batch.capacity(),
DEFAULT_MAX_BATCH_BYTES
);
}
#[tokio::test]
async fn staging_a_large_frame_would_retain_its_capacity() {
const BIG: usize = 4 * 1024 * 1024;
let mut batch = FrameBatchBuffer::new();
batch
.push(MessageType::Message, &[], &vec![0u8; BIG])
.unwrap();
let mut sink = RecordingSink::default();
batch.flush_to(&mut sink).await.unwrap();
assert_eq!(batch.frame_count(), 0, "flush resets the frame count");
assert!(
batch.capacity() >= BIG,
"flushing released {} bytes of capacity; if BytesMut started \
shrinking on clear, the WriteDirect routing's rationale changed",
batch.capacity()
);
}
#[tokio::test]
async fn large_frame_flushes_staged_frames_before_writing_direct() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::default();
let mut items: Vec<TestItem> = (0..3u8)
.map(|i| factory.item(&format!("small{i}"), vec![i; 8]))
.collect();
items.push(factory.item("big", vec![0xAB; COALESCE_THRESHOLD + 1]));
items.push(factory.item("after", vec![0xCD; 8]));
run_with(items, &mut sink, &observer).await;
assert_eq!(
observer.flushes(),
vec![3, 1, 1],
"three staged, then the large frame alone, then the tail"
);
assert_eq!(
sink.poll_writes, 4,
"the large frame must take the segmented direct path"
);
let decoded = sink.decode_frames();
assert_eq!(decoded.len(), 5);
for (i, (_, _, payload)) in decoded.iter().take(3).enumerate() {
assert_eq!(payload.as_slice(), &[i as u8; 8]);
}
assert_eq!(decoded[3].2.len(), COALESCE_THRESHOLD + 1);
assert!(decoded[3].2.iter().all(|&b| b == 0xAB));
assert_eq!(decoded[4].2.as_slice(), &[0xCD; 8]);
assert!(factory.errors().is_empty());
}
#[tokio::test]
async fn direct_write_stages_preamble_and_header_into_one_write() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::default();
let header = vec![0x11; 64];
let items =
vec![factory.item_with_header("big", header.clone(), vec![0xAB; COALESCE_THRESHOLD + 1])];
run_with(items, &mut sink, &observer).await;
assert_eq!(
sink.poll_writes, 2,
"staged prefix + payload; preamble and header must not write separately"
);
let decoded = sink.decode_frames();
assert_eq!(decoded.len(), 1);
assert_eq!(decoded[0].1, header);
assert_eq!(decoded[0].2.len(), COALESCE_THRESHOLD + 1);
assert!(factory.errors().is_empty());
}
#[tokio::test]
async fn direct_write_oversized_header_falls_back_to_three_segments() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::default();
let header = vec![0x22; DIRECT_PREFIX_CAP - MIN_HEADER_SIZE + 1];
let items =
vec![factory.item_with_header("big", header.clone(), vec![0xAB; COALESCE_THRESHOLD + 1])];
run_with(items, &mut sink, &observer).await;
assert_eq!(
sink.poll_writes, 3,
"preamble, header, and payload each take their own write"
);
let decoded = sink.decode_frames();
assert_eq!(decoded.len(), 1);
assert_eq!(decoded[0].1, header);
assert_eq!(decoded[0].2.len(), COALESCE_THRESHOLD + 1);
assert!(factory.errors().is_empty());
}
#[tokio::test]
async fn direct_write_reports_one_batch_despite_several_writes() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::default();
let items = vec![factory.item("big", vec![0xAB; COALESCE_THRESHOLD + 1])];
run_with(items, &mut sink, &observer).await;
assert_eq!(
observer.flushes(),
vec![1],
"one batch carrying one frame, whatever it cost to write"
);
assert!(
sink.poll_writes > 1,
"the direct path splits the frame across writes ({} here), which is \
why the counter must be described as batches, not syscalls",
sink.poll_writes
);
}
#[tokio::test]
async fn terminal_flushes_its_batch_then_stops() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::default();
let items = vec![
factory.item("a", vec![1; 8]),
factory.item("b", vec![2; 8]),
factory.terminal("fin", vec![3; 8]),
factory.item("after-terminal", vec![4; 8]),
];
run_with(items, &mut sink, &observer).await;
let decoded = sink.decode_frames();
assert_eq!(
decoded.len(),
3,
"frames staged ahead of the terminal must be written, and nothing after it"
);
assert_eq!(decoded[2].2.as_slice(), &[3; 8]);
assert_eq!(observer.flushes(), vec![3]);
}
#[tokio::test]
async fn channel_close_drains_remaining_items() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::default();
let items = (0..5u8)
.map(|i| factory.item(&format!("i{i}"), vec![i; 8]))
.collect();
run_with(items, &mut sink, &observer).await;
assert_eq!(sink.decode_frames().len(), 5);
assert!(factory.errors().is_empty());
}
#[tokio::test]
async fn cancellation_interrupts_a_refilled_queue() {
const QUEUED: usize = 200;
const PAYLOAD: usize = 8 * 1024;
let factory = ItemFactory::new();
let observer = TestObserver::default();
let cancel = CancellationToken::new();
let mut sink = RecordingSink {
cancel_on_write: Some(cancel.clone()),
..Default::default()
};
let (tx, rx) = flume::unbounded::<TestItem>();
for i in 0..QUEUED {
tx.send(factory.item(&format!("i{i}"), vec![0u8; PAYLOAD]))
.expect("queue");
}
let _tx = tx;
tokio::time::timeout(
std::time::Duration::from_secs(5),
run_coalescing_writer(
&mut sink,
&rx,
std::convert::identity,
Some(&cancel),
&observer,
),
)
.await
.expect("cancellation must stop the writer promptly");
assert!(
observer.frames_written() < QUEUED,
"cancellation must interrupt the drain, but all {QUEUED} items were written"
);
assert!(
!rx.is_empty(),
"items should remain queued for the caller's own drain to report"
);
}
struct LiveGuard {
live: Arc<AtomicUsize>,
}
impl LiveGuard {
fn new(live: &Arc<AtomicUsize>) -> Self {
live.fetch_add(1, Ordering::SeqCst);
Self {
live: Arc::clone(live),
}
}
}
impl Drop for LiveGuard {
fn drop(&mut self) {
self.live.fetch_sub(1, Ordering::SeqCst);
}
}
struct RetainingItem {
payload: Vec<u8>,
guard: LiveGuard,
}
impl Coalescable for RetainingItem {
type FailureToken = LiveGuard;
fn msg_type(&self) -> MessageType {
MessageType::Message
}
fn header(&self) -> &[u8] {
&[]
}
fn payload(&self) -> &[u8] {
&self.payload
}
fn into_failure_token(self) -> LiveGuard {
self.guard
}
fn fail(_token: LiveGuard, _reason: &str) {}
}
struct DiscardingItem {
payload: Vec<u8>,
_guard: LiveGuard,
}
impl Coalescable for DiscardingItem {
type FailureToken = ();
fn msg_type(&self) -> MessageType {
MessageType::Message
}
fn header(&self) -> &[u8] {
&[]
}
fn payload(&self) -> &[u8] {
&self.payload
}
fn into_failure_token(self) {}
fn fail(_token: (), _reason: &str) {}
}
async fn live_guards_at_flush<T: Coalescable>(
make: impl Fn(&Arc<AtomicUsize>, Vec<u8>) -> T,
) -> Vec<usize> {
let live = Arc::new(AtomicUsize::new(0));
let observer = TestObserver::default();
let mut sink = RecordingSink {
live_items: Some(Arc::clone(&live)),
..Default::default()
};
let (tx, rx) = flume::unbounded::<T>();
for i in 0..8u8 {
assert!(tx.send(make(&live, vec![i; 16])).is_ok(), "queue");
}
drop(tx);
run_coalescing_writer(&mut sink, &rx, std::convert::identity, None, &observer).await;
assert_eq!(observer.flushes(), vec![8], "all eight in one flush");
assert_eq!(live.load(Ordering::SeqCst), 0, "everything dropped by exit");
sink.live_at_write
}
#[tokio::test]
async fn tokens_survive_until_their_batch_is_written() {
let counts = live_guards_at_flush(|live, payload| RetainingItem {
payload,
guard: LiveGuard::new(live),
})
.await;
assert_eq!(
counts,
vec![8],
"all eight tokens must still be alive when the batch is written"
);
}
#[tokio::test]
async fn a_unit_token_retains_nothing_past_staging() {
let counts = live_guards_at_flush(|live, payload| DiscardingItem {
payload,
_guard: LiveGuard::new(live),
})
.await;
assert_eq!(
counts,
vec![0],
"items must be dropped at staging time, not held until flush"
);
}
#[tokio::test]
async fn write_failure_reports_every_staged_item() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::failing_at(0);
let items = (0..5u8)
.map(|i| factory.item(&format!("i{i}"), vec![i; 8]))
.collect();
run_with(items, &mut sink, &observer).await;
let errors = factory.errors();
assert_eq!(errors.len(), 5, "all five items reported: {errors:?}");
for i in 0..5 {
assert_eq!(
factory.reports_for(&format!("i{i}")),
1,
"item i{i} must be reported exactly once: {errors:?}"
);
}
assert_eq!(observer.failures(), vec![(WriterFailure::Write, 5)]);
}
#[tokio::test]
async fn flush_failure_before_staging_reports_the_held_item_once() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::failing_at(0);
let items = vec![
factory.item("staged", vec![1; 8]),
factory.item("held", vec![2; COALESCE_THRESHOLD + 1]),
];
run_with(items, &mut sink, &observer).await;
let errors = factory.errors();
assert_eq!(errors.len(), 2, "both items reported: {errors:?}");
assert_eq!(factory.reports_for("staged"), 1);
assert_eq!(factory.reports_for("held"), 1);
assert!(
errors.iter().any(|e| e == &format!("held: {FLUSH_FAILED}")),
"the held item must carry the flush-failure reason: {errors:?}"
);
}
#[tokio::test]
async fn encode_failure_flushes_staged_frames_and_reports_the_offender() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::default();
let items = vec![
factory.item("good", vec![1; 8]),
factory.item("bad", vec![0u8; (DEFAULT_MAX_FRAME_SIZE as usize) + 1]),
];
run_with(items, &mut sink, &observer).await;
assert_eq!(
sink.decode_frames().len(),
1,
"the frame staged before the bad one must still be written"
);
let errors = factory.errors();
assert_eq!(errors.len(), 1, "only the offender is reported: {errors:?}");
assert_eq!(factory.reports_for("bad"), 1);
assert_eq!(observer.failures(), vec![(WriterFailure::Encode, 1)]);
assert_eq!(observer.flushes(), vec![1], "the good frame flushed");
}
#[tokio::test]
async fn direct_write_failure_reports_only_that_item() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::failing_at(0);
let items = vec![factory.item("big", vec![0xAB; COALESCE_THRESHOLD + 1])];
run_with(items, &mut sink, &observer).await;
assert_eq!(factory.errors().len(), 1, "{:?}", factory.errors());
assert_eq!(factory.reports_for("big"), 1);
assert_eq!(observer.failures(), vec![(WriterFailure::Write, 1)]);
}
struct WrappedFrame(Vec<u8>);
impl Coalescable for WrappedFrame {
type FailureToken = ();
fn msg_type(&self) -> MessageType {
MessageType::Message
}
fn header(&self) -> &[u8] {
&[]
}
fn payload(&self) -> &[u8] {
&self.0
}
fn into_failure_token(self) {}
fn fail(_token: (), _reason: &str) {}
fn is_terminal(&self) -> bool {
self.0.first() == Some(&0xFF)
}
}
#[tokio::test]
async fn wrapped_channel_items_keep_their_terminal_semantics() {
let observer = TestObserver::default();
let mut sink = RecordingSink::default();
let (tx, rx) = flume::unbounded::<Vec<u8>>();
tx.send(vec![1u8; 8]).expect("queue");
tx.send(vec![0xFFu8; 8]).expect("queue terminal");
tx.send(vec![2u8; 8]).expect("queue after terminal");
let _tx = tx;
tokio::time::timeout(
std::time::Duration::from_secs(5),
run_coalescing_writer(&mut sink, &rx, WrappedFrame, None, &observer),
)
.await
.expect("the terminal frame must stop the writer");
let decoded = sink.decode_frames();
assert_eq!(
decoded.len(),
2,
"the frame before the terminal and the terminal itself, nothing after"
);
assert_eq!(decoded[1].2.as_slice(), &[0xFF; 8]);
assert_eq!(observer.flushes(), vec![2]);
assert_eq!(
rx.len(),
1,
"the frame queued behind the terminal must be left alone"
);
}
#[tokio::test]
async fn successful_writes_report_nothing() {
let factory = ItemFactory::new();
let observer = TestObserver::default();
let mut sink = RecordingSink::default();
let items = (0..12u8)
.map(|i| factory.item(&format!("i{i}"), vec![i; 32]))
.collect();
run_with(items, &mut sink, &observer).await;
assert_eq!(sink.decode_frames().len(), 12);
assert!(factory.errors().is_empty());
assert!(observer.failures().is_empty());
}