#![allow(unsafe_code)]
use bytes::BytesMut;
use compio_buf::{BufResult, IoBufMut, IoVectoredBuf};
use smallvec::SmallVec;
use std::io::{self, IoSlice};
use std::mem::MaybeUninit;
pub const READ_SLAB_SIZE: usize = 64 * 1024;
pub unsafe fn take_read_buffer(stash: &mut BytesMut, read_size: usize) -> BytesMut {
if stash.capacity() < read_size {
*stash = BytesMut::with_capacity(read_size.max(READ_SLAB_SIZE));
}
if stash.len() < read_size {
unsafe { stash.set_len(read_size) };
}
let tail = stash.split_off(read_size);
std::mem::replace(stash, tail)
}
pub async fn fill_read<B, F>(mut buf: B, read: F) -> BufResult<usize, B>
where
B: IoBufMut,
F: AsyncFnOnce(&mut [MaybeUninit<u8>]) -> io::Result<usize>,
{
let (outcome, spare_len) = {
let spare = buf.as_uninit();
let spare_len = spare.len();
(read(spare).await, spare_len)
};
match outcome {
Ok(n) => {
debug_assert!(
n <= spare_len,
"read reported {n} bytes but the spare slice was only {spare_len}"
);
unsafe {
buf.set_len(n);
}
BufResult(Ok(n), buf)
}
Err(e) => BufResult(Err(e), buf),
}
}
pub fn with_vectored_slices<B, R>(buf: &B, f: impl FnOnce(&[IoSlice<'_>]) -> R) -> R
where
B: IoVectoredBuf,
{
let slices: SmallVec<[IoSlice<'_>; 16]> = buf.iter_slice().map(IoSlice::new).collect();
f(&slices)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn reclaim_makes_slab_use_track_bytes_read_and_freezes_only_read_bytes() {
let mut stash = BytesMut::new();
stash.resize(READ_SLAB_SIZE, 0);
let slab_ptr = stash.as_ptr();
let mut buf = unsafe { take_read_buffer(&mut stash, 8192) };
assert_eq!(buf.as_ptr(), slab_ptr, "carved from the slab front");
buf[..10].copy_from_slice(b"0123456789");
buf.truncate(10);
let mut reclaimed = buf.split_off(10);
unsafe {
reclaimed.set_len(reclaimed.capacity());
}
reclaimed.unsplit(std::mem::take(&mut stash));
stash = reclaimed;
let frozen = buf.freeze();
assert_eq!(frozen.as_ref(), b"0123456789");
assert_eq!(
frozen.as_ptr(),
slab_ptr,
"frozen frame aliases the slab head"
);
let next = unsafe { take_read_buffer(&mut stash, 8192) };
assert_eq!(
next.as_ptr() as usize,
slab_ptr as usize + 10,
"the next carve reuses the reclaimed space (slab use tracks bytes read)"
);
}
#[test]
fn take_read_buffer_over_slab_size_does_not_panic_or_corrupt() {
let mut stash = BytesMut::with_capacity(READ_SLAB_SIZE);
let read_size = READ_SLAB_SIZE + 1024;
let mut buf = unsafe { take_read_buffer(&mut stash, read_size) };
assert_eq!(buf.len(), read_size);
buf.truncate(read_size);
assert_eq!(buf.freeze().len(), read_size);
}
#[test]
fn take_read_buffer_splits_front_and_keeps_tail() {
let mut stash = BytesMut::with_capacity(READ_SLAB_SIZE);
let buf = unsafe { take_read_buffer(&mut stash, 256) };
assert_eq!(buf.len(), 256);
assert_eq!(stash.len(), 0);
assert!(stash.capacity() >= READ_SLAB_SIZE - 256);
}
#[test]
fn take_read_buffer_reuses_one_slab_across_reads() {
let mut stash = BytesMut::with_capacity(READ_SLAB_SIZE);
let _first = unsafe { take_read_buffer(&mut stash, 4096) };
assert_eq!(stash.capacity(), READ_SLAB_SIZE - 4096);
let _second = unsafe { take_read_buffer(&mut stash, 4096) };
assert_eq!(stash.capacity(), READ_SLAB_SIZE - 8192);
}
}