use std::io;
use bytes::BytesMut;
use tokio::io::{AsyncWrite, AsyncWriteExt};
use tokio_util::sync::CancellationToken;
use velo_ext::MessageType;
use super::tcp::framing::{
COALESCE_THRESHOLD, DIRECT_PREFIX_CAP, MIN_HEADER_SIZE, TcpFrameCodec, stage_direct_prefix,
};
const DEFAULT_MAX_BATCH_BYTES: usize = COALESCE_THRESHOLD;
const DEFAULT_MAX_BATCH_FRAMES: usize = 1024;
const FLUSH_FAILED: &str = "batch flush failed";
pub(crate) trait Coalescable: Sized {
type FailureToken: Send;
fn msg_type(&self) -> MessageType;
fn header(&self) -> &[u8];
fn payload(&self) -> &[u8];
fn into_failure_token(self) -> Self::FailureToken;
fn fail(token: Self::FailureToken, reason: &str);
fn is_terminal(&self) -> bool {
false
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum WriterFailure {
Write,
Encode,
}
pub(crate) trait WriterObserver {
fn on_flush(&self, _frames: usize) {}
fn on_failure(&self, _kind: WriterFailure, _err: &io::Error, _frames: usize) {}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Staging {
Stage,
FlushThenStage,
WriteDirect,
}
impl Staging {
#[inline]
fn needs_flush_first(self) -> bool {
!matches!(self, Staging::Stage)
}
}
struct FrameBatchBuffer {
buf: BytesMut,
frames: usize,
max_bytes: usize,
max_frames: usize,
}
impl FrameBatchBuffer {
fn new() -> Self {
Self::with_limits(DEFAULT_MAX_BATCH_BYTES, DEFAULT_MAX_BATCH_FRAMES)
}
fn with_limits(max_bytes: usize, max_frames: usize) -> Self {
Self {
buf: BytesMut::with_capacity(8 * 1024),
frames: 0,
max_bytes,
max_frames,
}
}
#[inline]
fn frame_count(&self) -> usize {
self.frames
}
#[inline]
fn classify(&self, header_len: usize, payload_len: usize) -> Staging {
if header_len.saturating_add(payload_len) > COALESCE_THRESHOLD {
return Staging::WriteDirect;
}
if self.frames == 0 {
return Staging::Stage;
}
if self.frames + 1 > self.max_frames
|| self.buf.len() + MIN_HEADER_SIZE + header_len + payload_len > self.max_bytes
{
return Staging::FlushThenStage;
}
Staging::Stage
}
#[inline]
fn push(&mut self, msg_type: MessageType, header: &[u8], payload: &[u8]) -> io::Result<()> {
TcpFrameCodec::append_frame(&mut self.buf, msg_type, header, payload)?;
self.frames += 1;
Ok(())
}
async fn flush_to<W: AsyncWrite + Unpin>(&mut self, writer: &mut W) -> io::Result<()> {
if self.frames == 0 {
return Ok(());
}
let result = writer.write_all(&self.buf).await;
self.buf.clear();
self.frames = 0;
result
}
#[cfg(test)]
fn capacity(&self) -> usize {
self.buf.capacity()
}
}
pub(crate) async fn run_coalescing_writer<W, I, T, O>(
writer: &mut W,
rx: &flume::Receiver<I>,
wrap: impl Fn(I) -> T,
cancel: Option<&CancellationToken>,
observer: &O,
) where
W: AsyncWrite + Unpin,
T: Coalescable,
O: WriterObserver,
{
let mut batch = FrameBatchBuffer::new();
let mut staged: Vec<T::FailureToken> = Vec::new();
'writer: loop {
let first = tokio::select! {
biased;
_ = wait_cancelled(cancel) => break 'writer,
recv = rx.recv_async() => match recv {
Ok(item) => wrap(item),
Err(_) => break 'writer,
},
};
let mut pending = Some(first);
let mut terminal = false;
while let Some(item) = pending.take() {
let staging = batch.classify(item.header().len(), item.payload().len());
if staging.needs_flush_first()
&& !flush::<_, T, _>(&mut batch, &mut staged, writer, observer).await
{
T::fail(item.into_failure_token(), FLUSH_FAILED);
break 'writer;
}
if staging == Staging::WriteDirect {
if let Err((kind, e)) = write_frame_direct(writer, &item).await {
observer.on_failure(kind, &e, 1);
T::fail(item.into_failure_token(), &e.to_string());
break 'writer;
}
observer.on_flush(1);
if item.is_terminal() {
terminal = true;
break;
}
} else {
if let Err(e) = batch.push(item.msg_type(), item.header(), item.payload()) {
observer.on_failure(WriterFailure::Encode, &e, 1);
T::fail(item.into_failure_token(), &e.to_string());
flush::<_, T, _>(&mut batch, &mut staged, writer, observer).await;
break 'writer;
}
let is_terminal = item.is_terminal();
staged.push(item.into_failure_token());
if is_terminal {
terminal = true;
break;
}
}
pending = if is_cancelled(cancel) {
None
} else {
rx.try_recv().ok().map(&wrap)
};
}
if !flush::<_, T, _>(&mut batch, &mut staged, writer, observer).await || terminal {
break 'writer;
}
}
debug_assert!(
staged.is_empty(),
"every exit path flushes or reports the staged batch"
);
}
async fn flush<W, T, O>(
batch: &mut FrameBatchBuffer,
staged: &mut Vec<T::FailureToken>,
writer: &mut W,
observer: &O,
) -> bool
where
W: AsyncWrite + Unpin,
T: Coalescable,
O: WriterObserver,
{
debug_assert_eq!(
staged.len(),
batch.frame_count(),
"the writer must hold one failure token per staged frame"
);
let frames = batch.frame_count();
if frames == 0 {
return true;
}
match batch.flush_to(writer).await {
Ok(()) => {
staged.clear();
observer.on_flush(frames);
true
}
Err(e) => {
observer.on_failure(WriterFailure::Write, &e, frames);
let reason = e.to_string();
for token in staged.drain(..) {
T::fail(token, &reason);
}
false
}
}
}
async fn write_frame_direct<W, T>(
writer: &mut W,
item: &T,
) -> Result<(), (WriterFailure, io::Error)>
where
W: AsyncWrite + Unpin,
T: Coalescable,
{
let header = item.header();
let payload = item.payload();
let lengths = u32::try_from(header.len())
.and_then(|h| u32::try_from(payload.len()).map(|p| (h, p)))
.map_err(|_| {
(
WriterFailure::Encode,
io::Error::new(io::ErrorKind::InvalidData, "Frame length exceeds u32"),
)
});
let (header_len, payload_len) = lengths?;
let preamble = TcpFrameCodec::build_preamble(item.msg_type(), header_len, payload_len)
.map_err(|e| (WriterFailure::Encode, e))?;
let mut prefix = [0u8; DIRECT_PREFIX_CAP];
let segments: &[&[u8]] = match stage_direct_prefix(&preamble, header, &mut prefix) {
Some(len) => &[&prefix[..len], payload],
None => &[&preamble[..], header, payload],
};
for segment in segments {
writer
.write_all(segment)
.await
.map_err(|e| (WriterFailure::Write, e))?;
}
Ok(())
}
async fn wait_cancelled(cancel: Option<&CancellationToken>) {
match cancel {
Some(token) => token.cancelled().await,
None => std::future::pending::<()>().await,
}
}
#[inline]
fn is_cancelled(cancel: Option<&CancellationToken>) -> bool {
cancel.is_some_and(|token| token.is_cancelled())
}
#[cfg(test)]
mod tests;