use alloc::collections::VecDeque;
use alloc::vec::Vec;
use core::mem;
use std::io;
use crate::crypto::cipher::OutboundPlain;
pub(crate) struct ChunkVecBuffer {
prefix_used: usize,
chunks: VecDeque<Vec<u8>>,
limit: Option<usize>,
spare: Option<Vec<Vec<u8>>>,
}
impl ChunkVecBuffer {
pub(crate) fn new(limit: Option<usize>) -> Self {
Self {
prefix_used: 0,
chunks: VecDeque::new(),
limit,
spare: None,
}
}
pub(crate) fn new_recycling(limit: Option<usize>) -> Self {
Self {
prefix_used: 0,
chunks: VecDeque::new(),
limit,
spare: Some(Vec::new()),
}
}
pub(crate) fn take_spare(&mut self) -> Vec<u8> {
self.spare
.as_mut()
.and_then(|spare| spare.pop())
.unwrap_or_default()
}
pub(crate) fn set_limit(&mut self, new_limit: Option<usize>) {
self.limit = new_limit;
}
pub(crate) fn is_empty(&self) -> bool {
self.chunks.is_empty()
}
pub(crate) fn len(&self) -> usize {
self.chunks
.iter()
.fold(0usize, |acc, chunk| acc + chunk.len())
- self.prefix_used
}
pub(crate) fn append(&mut self, bytes: Vec<u8>) -> usize {
let len = bytes.len();
if !bytes.is_empty() {
if self.chunks.is_empty() {
debug_assert_eq!(self.prefix_used, 0);
}
self.chunks.push_back(bytes);
}
len
}
pub(crate) fn take(&mut self) -> Vec<Vec<u8>> {
if self.chunks.is_empty() {
return Vec::new();
}
let mut chunks = Vec::from(mem::take(&mut self.chunks));
let prefix = mem::take(&mut self.prefix_used);
chunks[0].drain(0..prefix);
chunks
}
pub(crate) fn take_one_vec(&mut self) -> Vec<u8> {
let Some(mut first) = self.pop() else {
return Vec::new();
};
while let Some(chunk) = self.chunks.pop_front() {
first.extend_from_slice(&chunk);
}
first
}
pub(crate) fn pop(&mut self) -> Option<Vec<u8>> {
let mut first = self.chunks.pop_front();
if let Some(first) = &mut first {
let prefix = mem::take(&mut self.prefix_used);
first.drain(0..prefix);
}
first
}
pub(crate) fn peek(&self) -> Option<&[u8]> {
self.chunks
.front()
.map(|ch| ch.as_slice())
}
}
impl ChunkVecBuffer {
pub(crate) fn is_full(&self) -> bool {
self.limit
.map(|limit| self.len() >= limit)
.unwrap_or_default()
}
pub(crate) fn append_limited_copy(&mut self, payload: OutboundPlain<'_>) -> usize {
let take = self.apply_limit(payload.len());
self.append(payload.split_at(take).0.to_vec());
take
}
pub(crate) fn apply_limit(&self, len: usize) -> usize {
let Some(limit) = self.limit else {
return len;
};
let space = limit.saturating_sub(self.len());
Ord::min(len, space)
}
pub(crate) fn read(&mut self, buf: &mut [u8]) -> usize {
let mut offs = 0;
while offs < buf.len() && !self.is_empty() {
let chunk = &self.chunks[0][self.prefix_used..];
let used = Ord::min(chunk.len(), buf.len() - offs);
buf[offs..offs + used].copy_from_slice(&chunk[..used]);
self.consume(used);
offs += used;
}
offs
}
pub(crate) fn consume_first_chunk(&mut self, used: usize) {
assert!(
used <= self
.chunk()
.map(|ch| ch.len())
.unwrap_or_default(),
"illegal `BufRead::consume` usage",
);
self.consume(used);
}
fn consume(&mut self, used: usize) {
self.prefix_used += used;
while let Some(buf) = self.chunks.front() {
if self.prefix_used < buf.len() {
return;
}
self.prefix_used -= buf.len();
if let Some(spent) = self.chunks.pop_front() {
self.recycle(spent);
}
}
debug_assert_eq!(
self.prefix_used, 0,
"attempted to `ChunkVecBuffer::consume` more than available"
);
}
fn recycle(&mut self, spent: Vec<u8>) {
let Some(spare) = &mut self.spare else {
return;
};
if let Some(limit) = self.limit {
let retained = spare
.iter()
.map(|chunk| chunk.capacity())
.sum::<usize>();
if retained + spent.capacity() > limit {
return;
}
}
spare.push(spent);
}
pub(crate) fn write_to(&mut self, wr: &mut dyn io::Write) -> io::Result<usize> {
if self.is_empty() {
return Ok(0);
}
let mut prefix = self.prefix_used;
let mut bufs = [io::IoSlice::new(&[]); 64];
for (iov, chunk) in bufs.iter_mut().zip(self.chunks.iter()) {
*iov = io::IoSlice::new(&chunk[prefix..]);
prefix = 0;
}
let len = Ord::min(bufs.len(), self.chunks.len());
let bufs = &bufs[..len];
let used = wr.write_vectored(bufs)?;
let available_bytes = bufs.iter().map(|ch| ch.len()).sum();
if used > available_bytes {
self.consume(available_bytes);
return Err(io::Error::other(std::format!(
"illegal write_vectored return value ({used} > {available_bytes})"
)));
}
self.consume(used);
Ok(used)
}
pub(crate) fn chunk(&self) -> Option<&[u8]> {
self.chunks
.front()
.map(|chunk| &chunk[self.prefix_used..])
}
}
#[cfg(test)]
mod tests {
use alloc::vec;
use alloc::vec::Vec;
use super::ChunkVecBuffer;
#[test]
fn short_append_copy_with_limit() {
let mut cvb = ChunkVecBuffer::new(Some(12));
assert_eq!(cvb.append_limited_copy(b"hello"[..].into()), 5);
assert_eq!(cvb.append_limited_copy(b"world"[..].into()), 5);
assert_eq!(cvb.append_limited_copy(b"hello"[..].into()), 2);
assert_eq!(cvb.append_limited_copy(b"world"[..].into()), 0);
let mut buf = [0u8; 12];
assert_eq!(cvb.read(&mut buf), 12);
assert_eq!(buf.to_vec(), b"helloworldhe".to_vec());
}
#[test]
fn recycling_retains_spent_chunks() {
let mut cvb = ChunkVecBuffer::new_recycling(None);
assert!(cvb.take_spare().is_empty());
cvb.append(b"first".to_vec());
cvb.append(b"second".to_vec());
let mut buf = [0u8; 11];
assert_eq!(cvb.read(&mut buf), 11);
assert_eq!(cvb.take_spare(), b"second");
assert_eq!(cvb.take_spare(), b"first");
assert!(cvb.take_spare().is_empty());
let mut cvb = ChunkVecBuffer::new(None);
cvb.append(b"first".to_vec());
assert_eq!(cvb.read(&mut buf), 5);
assert!(cvb.take_spare().is_empty());
}
#[test]
fn recycling_bounded_by_limit() {
let mut cvb = ChunkVecBuffer::new_recycling(Some(8));
cvb.append(vec![1u8; 6]);
cvb.append(vec![2u8; 6]);
let mut buf = [0u8; 12];
assert_eq!(cvb.read(&mut buf), 12);
assert_eq!(cvb.take_spare(), vec![1u8; 6]);
assert!(cvb.take_spare().is_empty());
}
#[test]
fn take_slices_off_consumed_prefix() {
let mut cvb = ChunkVecBuffer::new(None);
cvb.append(b"hello".to_vec());
cvb.append(b"world".to_vec());
assert_eq!(cvb.read(&mut [0u8; 3]), 3);
assert_eq!(cvb.take(), [b"lo".to_vec(), b"world".to_vec()]);
assert_eq!(cvb.len(), 0);
let mut cvb = ChunkVecBuffer::new(None);
cvb.append(b"hello".to_vec());
cvb.append(b"world".to_vec());
assert_eq!(cvb.read(&mut [0u8; 3]), 3);
assert_eq!(cvb.take_one_vec(), b"loworld");
assert_eq!(cvb.len(), 0);
}
#[test]
fn read_byte_by_byte() {
let mut cvb = ChunkVecBuffer::new(None);
cvb.append(b"test fixture data".to_vec());
assert!(!cvb.is_empty());
for expect in b"test fixture data" {
let mut byte = [0];
assert_eq!(cvb.read(&mut byte), 1);
assert_eq!(byte[0], *expect);
}
assert_eq!(cvb.read(&mut [0]), 0);
}
#[test]
fn every_possible_chunk_interleaving() {
let input = (0..=0xffu8)
.cycle()
.take(4096)
.collect::<Vec<u8>>();
for input_chunk_len in 1..64usize {
for output_chunk_len in 1..65usize {
std::println!("check input={input_chunk_len} output={output_chunk_len}");
let mut cvb = ChunkVecBuffer::new(None);
for chunk in input.chunks(input_chunk_len) {
cvb.append(chunk.to_vec());
}
assert_eq!(cvb.len(), input.len());
let mut buf = vec![0u8; output_chunk_len];
for expect in input.chunks(output_chunk_len) {
assert_eq!(expect.len(), cvb.read(&mut buf));
assert_eq!(expect, &buf[..expect.len()]);
}
assert_eq!(cvb.read(&mut [0]), 0);
}
}
}
}
#[cfg(all(test, bench))]
mod benchmarks {
use alloc::vec;
use super::ChunkVecBuffer;
#[bench]
fn read_one_byte_from_large_message(b: &mut test::Bencher) {
b.iter(|| {
let mut cvb = ChunkVecBuffer::new(None);
cvb.append(vec![0u8; 16_384]);
assert_eq!(1, cvb.read(&mut [0u8]));
});
}
#[bench]
fn read_all_individual_from_large_message(b: &mut test::Bencher) {
b.iter(|| {
let mut cvb = ChunkVecBuffer::new(None);
cvb.append(vec![0u8; 16_384]);
loop {
if cvb.read(&mut [0u8]) == 0 {
break;
}
}
});
}
#[bench]
fn read_half_bytes_from_large_message(b: &mut test::Bencher) {
b.iter(|| {
let mut cvb = ChunkVecBuffer::new(None);
cvb.append(vec![0u8; 16_384]);
assert_eq!(8192, cvb.read(&mut [0u8; 8192]));
assert_eq!(8192, cvb.read(&mut [0u8; 8192]));
});
}
#[bench]
fn read_entire_large_message(b: &mut test::Bencher) {
b.iter(|| {
let mut cvb = ChunkVecBuffer::new(None);
cvb.append(vec![0u8; 16_384]);
assert_eq!(16_384, cvb.read(&mut [0u8; 16_384]));
});
}
}