use arrayvec::ArrayVec;
use core::ops::Deref;
use s2n_quic_core::{ensure, inet::ExplicitCongestionNotification};
use std::io::IoSlice;
pub const MAX_COUNT: usize = if cfg!(target_os = "linux") { 8 } else { 1 };
pub const MAX_TOTAL: u16 = u16::MAX - 50;
type Segments<'a> = ArrayVec<IoSlice<'a>, MAX_COUNT>;
pub struct Batch<'a> {
segments: Segments<'a>,
ecn: ExplicitCongestionNotification,
}
impl<'a> Deref for Batch<'a> {
type Target = [IoSlice<'a>];
#[inline]
fn deref(&self) -> &Self::Target {
&self.segments
}
}
impl<'a> Batch<'a> {
#[inline]
pub fn new<Q>(queue: Q) -> Self
where
Q: IntoIterator<Item = (ExplicitCongestionNotification, &'a [u8])>,
{
let mut ecn = ExplicitCongestionNotification::Ect0;
let mut total_len = 0u16;
let mut segments = Segments::new();
for segment in queue {
let packet_len = segment.1.len();
debug_assert!(
packet_len <= u16::MAX as usize,
"segments should not exceed the maximum datagram size"
);
let packet_len = packet_len as u16;
let Some(new_total_len) = total_len.checked_add(packet_len) else {
break;
};
ensure!(new_total_len < MAX_TOTAL, break);
let mut undersized_segment = false;
if let Some(first_segment) = segments.first() {
ensure!(first_segment.len() >= packet_len as usize, break);
undersized_segment = first_segment.len() > packet_len as usize;
ensure!(ecn == segment.0, break);
} else {
ecn = segment.0;
}
total_len = new_total_len;
let iovec = std::io::IoSlice::new(segment.1);
segments.push(iovec);
ensure!(!undersized_segment, break);
ensure!(!segments.is_full(), break);
}
Self { segments, ecn }
}
#[inline]
pub fn ecn(&self) -> ExplicitCongestionNotification {
self.ecn
}
}