use std::sync::Arc;
use bytes::{Bytes, BytesMut};
use velo_ext::{InstanceId, WorkerId};
use super::super::MuxConfig;
use super::super::protocol::{BATCH_HEADER_LEN, BatchEncoder, EncodeError, MAX_RECORDS_PER_BATCH};
use crate::messenger::{FireResult, Messenger};
use crate::observability::{MuxDirection, MuxMetricsHandle};
use crate::streaming::messenger_mux::STREAM_BATCH_HANDLER;
use crate::transports::tcp::framing::COALESCE_THRESHOLD;
pub(super) const MIN_BATCH_CAP: usize = BATCH_HEADER_LEN + 13;
pub(super) const fn batch_cap(configured: usize, eager: usize) -> usize {
let clamped = if configured < eager {
configured
} else {
eager
};
let clamped = if clamped < COALESCE_THRESHOLD {
clamped
} else {
COALESCE_THRESHOLD
};
if clamped > MIN_BATCH_CAP {
clamped
} else {
MIN_BATCH_CAP
}
}
#[derive(Debug)]
pub(super) struct FlushFailed(pub(super) anyhow::Error);
pub(super) struct BatchWriter {
messenger: Arc<Messenger>,
peer: WorkerId,
peer_instance: Option<InstanceId>,
config: MuxConfig,
metrics: Option<MuxMetricsHandle>,
epoch: u64,
next_batch_seq: u32,
cap: usize,
encoder: Option<BatchEncoder>,
buffer: BytesMut,
}
impl BatchWriter {
pub(super) fn new(
messenger: Arc<Messenger>,
peer: WorkerId,
config: MuxConfig,
metrics: Option<MuxMetricsHandle>,
epoch: u64,
) -> Self {
Self {
messenger,
peer,
peer_instance: None,
config,
metrics,
epoch,
next_batch_seq: 0,
cap: MIN_BATCH_CAP,
encoder: None,
buffer: BytesMut::new(),
}
}
pub(super) const fn epoch(&self) -> u64 {
self.epoch
}
pub(super) fn reset_epoch(&mut self, epoch: u64) {
self.epoch = epoch;
self.next_batch_seq = 0;
self.encoder = None;
}
pub(super) fn ensure_batch(&mut self) -> usize {
if self.encoder.is_none() {
self.cap = self.compute_cap();
let batch_seq = self.take_batch_seq();
let buffer = std::mem::take(&mut self.buffer);
self.encoder = Some(BatchEncoder::with_buffer(buffer, self.epoch, batch_seq));
}
self.cap
}
pub(super) fn fits(&self, bytes: usize, records: u16) -> bool {
let Some(encoder) = self.encoder.as_ref() else {
return true;
};
encoder.encoded_len().saturating_add(bytes) <= self.cap
&& u32::from(encoder.record_count()) + u32::from(records)
<= u32::from(MAX_RECORDS_PER_BATCH)
}
pub(super) fn encoder(&mut self) -> Option<&mut BatchEncoder> {
self.encoder.as_mut()
}
fn compute_cap(&mut self) -> usize {
let eager = self.peer_instance().map_or(usize::MAX, |instance| {
self.messenger
.effective_eager_payload(instance, STREAM_BATCH_HANDLER, None)
});
batch_cap(self.config.max_batch_bytes, eager)
}
fn peer_instance(&mut self) -> Option<InstanceId> {
if self.peer_instance.is_none() {
self.peer_instance = self
.messenger
.backend()
.try_translate_worker_id(self.peer)
.ok();
}
self.peer_instance
}
fn take_batch_seq(&mut self) -> u32 {
let seq = self.next_batch_seq;
self.next_batch_seq = self.next_batch_seq.wrapping_add(1);
seq
}
pub(super) async fn flush(&mut self) -> Result<(), FlushFailed> {
let Some(encoder) = self.encoder.take() else {
return Ok(());
};
if encoder.is_empty() {
self.next_batch_seq = self.next_batch_seq.wrapping_sub(1);
self.buffer = encoder.finish();
return Ok(());
}
let records = usize::from(encoder.record_count());
let mut finished = encoder.finish();
let payload = finished.split().freeze();
self.buffer = finished;
if let Some(metrics) = &self.metrics {
metrics.batch(MuxDirection::Sent, records);
}
self.dispatch(payload).await.map_err(FlushFailed)
}
pub(super) fn dispatch_singleton(
&mut self,
write: impl FnOnce(&mut BatchEncoder) -> Result<(), EncodeError>,
) -> Option<FireResult> {
let batch_seq = self.take_batch_seq();
let mut encoder = BatchEncoder::new(self.epoch, batch_seq);
if let Err(error) = write(&mut encoder) {
tracing::error!(%error, "messenger mux: dropping unencodable singleton");
return None;
}
let records = usize::from(encoder.record_count());
let payload = encoder.finish().freeze();
if let Some(metrics) = &self.metrics {
metrics.rendezvous_singleton();
metrics.batch(MuxDirection::Sent, records);
}
match self.messenger.am_send_streaming(STREAM_BATCH_HANDLER) {
Ok(builder) => Some(builder.raw_payload(payload).worker(self.peer).send()),
Err(error) => {
tracing::error!(%error, "messenger mux: could not build a singleton send");
None
}
}
}
async fn dispatch(&self, payload: Bytes) -> anyhow::Result<()> {
self.messenger
.am_send_streaming(STREAM_BATCH_HANDLER)?
.raw_payload(payload)
.worker(self.peer)
.send()
.await
}
}