#![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 const COPY_OUT_THRESHOLD: usize = 4 * 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 = {
let spare = buf.as_uninit();
read(spare).await
};
match outcome {
Ok(n) => {
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 small_frame_copy_releases_the_slab() {
use bytes::Bytes;
let mut stash = BytesMut::new();
stash.resize(READ_SLAB_SIZE, 0);
let slab_ptr = stash.as_ptr();
let mut a = unsafe { take_read_buffer(&mut stash, 8192) };
a.truncate(10);
let a_ptr = a.as_ptr();
assert_eq!(
a_ptr, slab_ptr,
"chunk should be carved from the slab front"
);
let frozen = a.freeze();
assert_eq!(
frozen.as_ptr(),
slab_ptr,
"freeze() shares the slab allocation"
);
let mut b = unsafe { take_read_buffer(&mut stash, 8192) };
b.truncate(10);
let b_ptr = b.as_ptr();
let copied = Bytes::copy_from_slice(&b);
assert_ne!(
copied.as_ptr(),
b_ptr,
"copy_from_slice must allocate off the slab, not alias it"
);
let _ = COPY_OUT_THRESHOLD;
}
#[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);
}
}