use std::{collections::VecDeque, ops::Range};
use bytes::{Buf, BufMut, Bytes, BytesMut};
use crate::{VarInt, connection::streams::BytesOrSlice, range_set::ArrayRangeSet};
#[derive(Default, Debug)]
pub(super) struct SendBuffer {
data: SendBufferData,
unsent: u64,
acks: ArrayRangeSet,
retransmits: ArrayRangeSet,
}
const MAX_COMBINE: usize = 1452;
#[derive(Default, Debug)]
struct SendBufferData {
offset: u64,
len: usize,
segments: VecDeque<Bytes>,
last_segment: BytesMut,
}
impl SendBufferData {
fn len(&self) -> usize {
self.len
}
#[inline(always)]
fn range(&self) -> Range<u64> {
self.offset..self.offset + self.len as u64
}
fn append<'a>(&'a mut self, data: impl BytesOrSlice<'a>) {
self.len += data.len();
if data.len() > MAX_COMBINE {
if !self.last_segment.is_empty() {
self.segments.push_back(self.last_segment.split().freeze());
}
self.segments.push_back(data.into_bytes());
} else {
let rest = if self.last_segment.len() + data.len() > MAX_COMBINE
&& !self.last_segment.is_empty()
{
let capacity = MAX_COMBINE.saturating_sub(self.last_segment.len());
let (curr, rest) = data.as_ref().split_at(capacity);
self.last_segment.put_slice(curr);
self.segments.push_back(self.last_segment.split().freeze());
rest
} else {
data.as_ref()
};
self.last_segment.extend_from_slice(rest);
}
}
fn pop_front(&mut self, n: usize) {
let mut n = n.min(self.len);
self.len -= n;
self.offset += n as u64;
while n > 0 {
let Some(front) = self.segments.front_mut() else {
break;
};
if front.len() <= n {
n -= front.len();
self.segments.pop_front();
} else {
front.advance(n);
n = 0;
}
}
self.last_segment.advance(n);
if self.segments.len() * 4 < self.segments.capacity() {
self.segments.shrink_to_fit();
}
}
fn segments_iter(&self) -> impl Iterator<Item = &[u8]> {
self.segments
.iter()
.map(|x| x.as_ref())
.chain(std::iter::once(self.last_segment.as_ref()))
}
#[cfg(any(test, feature = "bench"))]
fn get(&self, offsets: Range<u64>) -> &[u8] {
assert!(
offsets.start >= self.range().start && offsets.end <= self.range().end,
"Requested range is outside of buffered data"
);
let offsets = Range {
start: (offsets.start - self.offset) as usize,
end: (offsets.end - self.offset) as usize,
};
let mut segment_offset = 0;
for segment in self.segments_iter() {
if offsets.start >= segment_offset && offsets.start < segment_offset + segment.len() {
let start = offsets.start - segment_offset;
let end = offsets.end - segment_offset;
return &segment[start..end.min(segment.len())];
}
segment_offset += segment.len();
}
unreachable!("impossible if segments and range are consistent");
}
fn get_into(&self, offsets: Range<u64>, buf: &mut impl BufMut) {
assert!(
offsets.start >= self.range().start && offsets.end <= self.range().end,
"Requested range is outside of buffered data"
);
let offsets = Range {
start: (offsets.start - self.offset) as usize,
end: (offsets.end - self.offset) as usize,
};
let mut segment_offset = 0;
for segment in self.segments_iter() {
let start = segment_offset.max(offsets.start);
let end = (segment_offset + segment.len()).min(offsets.end);
if start < end {
buf.put_slice(&segment[start - segment_offset..end - segment_offset]);
}
segment_offset += segment.len();
if segment_offset >= offsets.end {
break;
}
}
}
#[cfg(test)]
fn to_vec(&self) -> Vec<u8> {
let mut result = Vec::with_capacity(self.len);
for segment in self.segments_iter() {
result.extend_from_slice(segment);
}
result
}
}
impl SendBuffer {
pub(super) fn new() -> Self {
Self::default()
}
pub(super) fn write<'a>(&'a mut self, data: impl BytesOrSlice<'a>) {
self.data.append(data);
}
pub(super) fn ack(&mut self, mut range: Range<u64>) {
let base_offset = self.fully_acked_offset();
range.start = base_offset.max(range.start);
range.end = base_offset.max(range.end);
self.acks.insert(range);
while self.acks.min() == Some(self.fully_acked_offset()) {
let prefix = self.acks.pop_min().unwrap();
let to_advance = (prefix.end - prefix.start) as usize;
self.data.pop_front(to_advance);
}
self.retransmits.remove(0..self.fully_acked_offset());
}
pub(super) fn poll_transmit(&mut self, mut max_len: usize) -> (Range<u64>, bool) {
debug_assert!(max_len >= 8 + 8);
let mut encode_length = false;
if let Some(range) = self.retransmits.pop_min() {
if range.start != 0 {
max_len -= VarInt::size(unsafe { VarInt::from_u64_unchecked(range.start) });
}
if range.end - range.start < max_len as u64 {
encode_length = true;
max_len -= 8;
}
let end = range.end.min((max_len as u64).saturating_add(range.start));
if end != range.end {
self.retransmits.insert(end..range.end);
}
return (range.start..end, encode_length);
}
if self.unsent != 0 {
max_len -= VarInt::size(unsafe { VarInt::from_u64_unchecked(self.unsent) });
}
if self.offset() - self.unsent < max_len as u64 {
encode_length = true;
max_len -= 8;
}
let end = self
.offset()
.min((max_len as u64).saturating_add(self.unsent));
let result = self.unsent..end;
self.unsent = end;
(result, encode_length)
}
#[cfg(any(test, feature = "bench"))]
pub(super) fn get(&self, offsets: Range<u64>) -> &[u8] {
self.data.get(offsets)
}
pub(super) fn get_into(&self, offsets: Range<u64>, buf: &mut impl BufMut) {
self.data.get_into(offsets, buf)
}
pub(super) fn retransmit(&mut self, mut range: Range<u64>) {
debug_assert!(range.end <= self.unsent, "unsent data can't be lost");
range.start = range.start.max(self.fully_acked_offset());
self.retransmits.insert(range);
}
pub(super) fn retransmit_all_for_0rtt(&mut self) {
debug_assert_eq!(self.fully_acked_offset(), 0);
self.unsent = 0;
}
fn fully_acked_offset(&self) -> u64 {
self.data.range().start
}
pub(super) fn offset(&self) -> u64 {
self.data.range().end
}
pub(super) fn is_fully_acked(&self) -> bool {
self.data.len() == 0
}
pub(super) fn has_unsent_data(&self) -> bool {
self.unsent != self.offset() || !self.retransmits.is_empty()
}
pub(super) fn unacked(&self) -> u64 {
self.data.len() as u64 - self.acks.iter().map(|x| x.end - x.start).sum::<u64>()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn fragment_with_length() {
let mut buf = SendBuffer::new();
const MSG: &[u8] = b"Hello, world!";
buf.write(MSG);
assert_eq!(buf.poll_transmit(19), (0..11, true));
assert_eq!(
buf.poll_transmit(MSG.len() + 16 - 11),
(11..MSG.len() as u64, true)
);
assert_eq!(
buf.poll_transmit(58),
(MSG.len() as u64..MSG.len() as u64, true)
);
}
#[test]
fn fragment_without_length() {
let mut buf = SendBuffer::new();
const MSG: &[u8] = b"Hello, world with some extra data!";
buf.write(MSG);
assert_eq!(buf.poll_transmit(19), (0..19, false));
assert_eq!(
buf.poll_transmit(MSG.len() - 19 + 1),
(19..MSG.len() as u64, false)
);
assert_eq!(
buf.poll_transmit(58),
(MSG.len() as u64..MSG.len() as u64, true)
);
}
#[test]
fn reserves_encoded_offset() {
let mut buf = SendBuffer::new();
let chunk: Bytes = Bytes::from_static(&[0; 1024 * 1024]);
for _ in 0..1025 {
buf.write(chunk.clone());
}
const SIZE1: u64 = 64;
const SIZE2: u64 = 16 * 1024;
const SIZE3: u64 = 1024 * 1024 * 1024;
assert_eq!(buf.poll_transmit(16), (0..16, false));
buf.retransmit(0..16);
assert_eq!(buf.poll_transmit(16), (0..16, false));
let mut transmitted = 16u64;
assert_eq!(
buf.poll_transmit((SIZE1 - transmitted + 1) as usize),
(transmitted..SIZE1, false)
);
buf.retransmit(transmitted..SIZE1);
assert_eq!(
buf.poll_transmit((SIZE1 - transmitted + 1) as usize),
(transmitted..SIZE1, false)
);
transmitted = SIZE1;
assert_eq!(
buf.poll_transmit((SIZE2 - transmitted + 2) as usize),
(transmitted..SIZE2, false)
);
buf.retransmit(transmitted..SIZE2);
assert_eq!(
buf.poll_transmit((SIZE2 - transmitted + 2) as usize),
(transmitted..SIZE2, false)
);
transmitted = SIZE2;
assert_eq!(
buf.poll_transmit((SIZE3 - transmitted + 4) as usize),
(transmitted..SIZE3, false)
);
buf.retransmit(transmitted..SIZE3);
assert_eq!(
buf.poll_transmit((SIZE3 - transmitted + 4) as usize),
(transmitted..SIZE3, false)
);
transmitted = SIZE3;
assert_eq!(
buf.poll_transmit(chunk.len() + 8),
(transmitted..transmitted + chunk.len() as u64, false)
);
buf.retransmit(transmitted..transmitted + chunk.len() as u64);
assert_eq!(
buf.poll_transmit(chunk.len() + 8),
(transmitted..transmitted + chunk.len() as u64, false)
);
}
#[test]
fn multiple_large_segments() {
const N: usize = 2000;
const K: u64 = N as u64;
fn dup(data: &[u8]) -> Bytes {
let mut buf = BytesMut::with_capacity(data.len() * N);
for c in data {
for _ in 0..N {
buf.put_u8(*c);
}
}
buf.freeze()
}
fn same(a: &[u8], b: &[u8]) -> bool {
std::ptr::eq(a.as_ptr(), b.as_ptr())
}
let mut buf = SendBuffer::new();
let msg: Bytes = dup(b"Hello, world!");
let msg_len: u64 = msg.len() as u64;
let seg1: Bytes = dup(b"He");
buf.write(seg1.clone());
let seg2: Bytes = dup(b"llo,");
buf.write(seg2.clone());
let seg3: Bytes = dup(b" w");
buf.write(seg3.clone());
let seg4: Bytes = dup(b"o");
buf.write(seg4.clone());
let seg5: Bytes = dup(b"rld!");
buf.write(seg5.clone());
assert_eq!(aggregate_unacked(&buf), msg);
assert!(same(buf.get(0..5 * K), &seg1));
assert!(same(buf.get(2 * K..8 * K), &seg2));
assert!(same(buf.get(6 * K..8 * K), &seg3));
assert!(same(buf.get(8 * 2000..msg_len), &seg4));
assert!(same(buf.get(9 * 2000..msg_len), &seg5));
buf.ack(0..K);
assert_eq!(aggregate_unacked(&buf), &msg[N..]);
buf.ack(0..3 * K);
assert_eq!(aggregate_unacked(&buf), &msg[3 * N..]);
buf.ack(3 * K..5 * K);
assert_eq!(aggregate_unacked(&buf), &msg[5 * N..]);
buf.ack(7 * K..9 * K);
assert_eq!(aggregate_unacked(&buf), &msg[5 * N..]);
buf.ack(4 * K..7 * K);
assert_eq!(aggregate_unacked(&buf), &msg[9 * N..]);
buf.ack(0..msg_len);
assert_eq!(aggregate_unacked(&buf), &[] as &[u8]);
}
#[test]
fn retransmit() {
let mut buf = SendBuffer::new();
const MSG: &[u8] = b"Hello, world with extra data!";
buf.write(MSG);
assert_eq!(buf.poll_transmit(16), (0..16, false));
assert_eq!(buf.poll_transmit(16), (16..23, true));
buf.retransmit(0..16);
assert_eq!(buf.poll_transmit(16), (0..16, false));
assert_eq!(buf.poll_transmit(16), (23..MSG.len() as u64, true));
buf.retransmit(16..23);
assert_eq!(buf.poll_transmit(16), (16..23, true));
}
#[test]
fn ack() {
let mut buf = SendBuffer::new();
const MSG: &[u8] = b"Hello, world!";
buf.write(MSG);
assert_eq!(buf.poll_transmit(16), (0..8, true));
buf.ack(0..8);
assert_eq!(aggregate_unacked(&buf), &MSG[8..]);
}
#[test]
fn reordered_ack() {
let mut buf = SendBuffer::new();
const MSG: &[u8] = b"Hello, world with extra data!";
buf.write(MSG);
assert_eq!(buf.poll_transmit(16), (0..16, false));
assert_eq!(buf.poll_transmit(16), (16..23, true));
buf.ack(16..23);
assert_eq!(aggregate_unacked(&buf), MSG);
buf.ack(0..16);
assert_eq!(aggregate_unacked(&buf), &MSG[23..]);
assert!(buf.acks.is_empty());
}
fn aggregate_unacked(buf: &SendBuffer) -> Vec<u8> {
buf.data.to_vec()
}
#[test]
#[should_panic(expected = "Requested range is outside of buffered data")]
fn send_buffer_get_out_of_range() {
let data = SendBufferData::default();
data.get(0..1);
}
#[test]
#[should_panic(expected = "Requested range is outside of buffered data")]
fn send_buffer_get_into_out_of_range() {
let data = SendBufferData::default();
let mut buf = Vec::new();
data.get_into(0..1, &mut buf);
}
}
#[cfg(all(test, not(target_family = "wasm")))]
mod proptests {
use super::*;
use proptest::prelude::*;
use test_strategy::{Arbitrary, proptest};
use crate::tests::subscribe;
use tracing::trace;
#[derive(Debug, Clone, Arbitrary)]
enum Op {
Write(#[strategy(proptest::collection::vec(any::<u8>(), 0..1024))] Vec<u8>),
Ack(Range<u64>),
Retransmit(Range<u64>),
PollTransmit(#[strategy(16usize..1024)] usize),
}
fn map_range(input: Range<u64>, target: Range<u64>) -> Range<u64> {
if target.is_empty() {
return target;
}
let size = target.end - target.start;
let a = target.start + (input.start % size);
let b = target.start + (input.end % size);
a.min(b)..a.max(b)
}
#[proptest]
fn send_buffer_matches_reference(
#[strategy(proptest::collection::vec(any::<Op>(), 1..100))] ops: Vec<Op>,
) {
let _guard = subscribe();
let mut sb = SendBuffer::new();
let mut buf = Vec::new();
let mut max_send_offset = 0u64;
let mut max_full_send_offset = 0u64;
trace!("");
for op in ops {
match op {
Op::Write(data) => {
trace!("Op::Write({})", data.len());
buf.extend_from_slice(&data);
sb.write(Bytes::from(data));
}
Op::Ack(range) => {
let range = map_range(range, 0..max_send_offset);
if range.contains(&max_full_send_offset) {
max_full_send_offset = range.end;
}
trace!("Op::Ack({:?})", range);
sb.ack(range);
}
Op::Retransmit(range) => {
let range = map_range(range, 0..max_send_offset);
trace!("Op::Retransmit({:?})", range);
sb.retransmit(range);
}
Op::PollTransmit(max_len) => {
trace!("Op::PollTransmit({})", max_len);
let (range, _partial) = sb.poll_transmit(max_len);
max_send_offset = max_send_offset.max(range.end);
assert!(
range.start >= max_full_send_offset,
"poll_transmit returned already fully acked data: range={:?}, max_full_send_offset={}",
range,
max_full_send_offset
);
let mut t1 = Vec::new();
sb.get_into(range.clone(), &mut t1);
let mut t2 = Vec::new();
t2.extend_from_slice(&buf[range.start as usize..range.end as usize]);
assert_eq!(t1, t2, "Data mismatch for range {:?}", range);
}
}
}
trace!("Op::Retransmit({:?})", 0..max_send_offset);
sb.retransmit(0..max_send_offset);
loop {
trace!("Op::PollTransmit({})", 1024);
let (range, _partial) = sb.poll_transmit(1024);
if range.is_empty() {
break;
}
trace!("Op::Ack({:?})", range);
sb.ack(range);
}
assert!(
sb.is_fully_acked(),
"SendBuffer not fully acked at end of ops"
);
}
}
#[cfg(feature = "bench")]
pub mod send_buffer_benches {
use bytes::Bytes;
use criterion::Criterion;
use super::SendBuffer;
pub fn get_into_many_segments(criterion: &mut Criterion) {
let mut group = criterion.benchmark_group("get_into_many_segments");
let mut buf = SendBuffer::new();
const SEGMENTS: u64 = 10000;
const SEGMENT_SIZE: u64 = 10;
const PACKET_SIZE: u64 = 1200;
const BYTES: u64 = SEGMENTS * SEGMENT_SIZE;
for i in 0..SEGMENTS {
buf.write(Bytes::from(vec![i as u8; SEGMENT_SIZE as usize]));
}
let mut tgt = Vec::with_capacity(PACKET_SIZE as usize);
group.bench_function("get_into", |b| {
b.iter(|| {
tgt.clear();
buf.get_into(BYTES - PACKET_SIZE..BYTES, std::hint::black_box(&mut tgt));
});
});
}
pub fn get_loop_many_segments(criterion: &mut Criterion) {
let mut group = criterion.benchmark_group("get_loop_many_segments");
let mut buf = SendBuffer::new();
const SEGMENTS: u64 = 10000;
const SEGMENT_SIZE: u64 = 10;
const PACKET_SIZE: u64 = 1200;
const BYTES: u64 = SEGMENTS * SEGMENT_SIZE;
for i in 0..SEGMENTS {
buf.write(Bytes::from(vec![i as u8; SEGMENT_SIZE as usize]));
}
let mut tgt = Vec::with_capacity(PACKET_SIZE as usize);
group.bench_function("get_loop", |b| {
b.iter(|| {
tgt.clear();
let mut range = BYTES - PACKET_SIZE..BYTES;
while range.start < range.end {
let slice = std::hint::black_box(buf.get(range.clone()));
range.start += slice.len() as u64;
tgt.extend_from_slice(slice);
}
});
});
}
}