use std::cell::UnsafeCell;
use std::io;
use std::mem::MaybeUninit;
use std::sync::Arc;
use minarrow::Vec64;
use minarrow::structs::shared_buffer::SharedBuffer;
use crate::constants::arena_capacity;
struct ArenaBacking {
data: UnsafeCell<Vec64<u8>>,
}
unsafe impl Send for ArenaBacking {}
unsafe impl Sync for ArenaBacking {}
struct BufferWindow {
backing: Arc<ArenaBacking>,
offset: usize,
len: usize,
}
impl AsRef<[u8]> for BufferWindow {
#[inline]
fn as_ref(&self) -> &[u8] {
let ptr = self.backing.data.get();
let data_ptr = unsafe { (*ptr).as_ptr() };
unsafe { std::slice::from_raw_parts(data_ptr.add(self.offset), self.len) }
}
}
unsafe impl Send for BufferWindow {}
unsafe impl Sync for BufferWindow {}
#[cfg(feature = "io_uring")]
pub struct ArenaRegion {
backing: Arc<ArenaBacking>,
offset: usize,
capacity: usize,
filled: usize,
}
#[cfg(feature = "io_uring")]
unsafe impl tokio_uring::buf::IoBuf for ArenaRegion {
fn stable_ptr(&self) -> *const u8 {
let ptr = self.backing.data.get();
unsafe { (*ptr).as_ptr().add(self.offset) }
}
fn bytes_init(&self) -> usize {
self.filled
}
fn bytes_total(&self) -> usize {
self.capacity
}
}
#[cfg(feature = "io_uring")]
unsafe impl tokio_uring::buf::IoBufMut for ArenaRegion {
fn stable_mut_ptr(&mut self) -> *mut u8 {
let ptr = self.backing.data.get();
unsafe { (*ptr).as_mut_ptr().add(self.offset) }
}
unsafe fn set_init(&mut self, pos: usize) {
if pos > self.filled {
self.filled = pos;
}
}
}
pub struct StreamArena {
backing: Arc<ArenaBacking>,
write_pos: usize,
capacity: usize,
}
impl Default for StreamArena {
fn default() -> Self {
Self::new()
}
}
impl StreamArena {
pub fn new() -> Self {
Self::with_capacity(arena_capacity())
}
pub fn with_capacity(capacity: usize) -> Self {
let v = Vec64::with_capacity(capacity);
Self {
backing: Arc::new(ArenaBacking {
data: UnsafeCell::new(v),
}),
write_pos: 0,
capacity,
}
}
#[inline]
pub fn spare_uninit(&mut self) -> &mut [MaybeUninit<u8>] {
let ptr = self.backing.data.get();
let data_ptr = unsafe { (*ptr).as_mut_ptr() };
let spare_ptr = unsafe { data_ptr.add(self.write_pos) as *mut MaybeUninit<u8> };
let spare_len = self.capacity - self.write_pos;
unsafe { std::slice::from_raw_parts_mut(spare_ptr, spare_len) }
}
pub fn extend_from_slice(&mut self, src: &[u8]) -> io::Result<()> {
if src.len() > self.capacity - self.write_pos {
return Err(io::Error::other(
"StreamArena::extend_from_slice: source exceeds remaining capacity",
));
}
let ptr = self.backing.data.get();
unsafe {
let dst = (*ptr).as_mut_ptr().add(self.write_pos);
std::ptr::copy_nonoverlapping(src.as_ptr(), dst, src.len());
}
self.write_pos += src.len();
Ok(())
}
#[inline]
pub unsafe fn advance(&mut self, n: usize) {
assert!(
n <= self.capacity - self.write_pos,
"StreamArena::advance past capacity"
);
self.write_pos += n;
}
#[inline]
pub fn align(&mut self) {
let remainder = self.write_pos % 64;
if remainder != 0 {
let padding = (64 - remainder).min(self.capacity - self.write_pos);
let ptr = self.backing.data.get();
unsafe {
(*ptr).as_mut_ptr().add(self.write_pos).write_bytes(0, padding);
}
self.write_pos += padding;
}
}
#[inline]
pub fn window(&self, offset: usize, len: usize) -> SharedBuffer {
let end = offset
.checked_add(len)
.expect("StreamArena::window range overflows");
assert!(end <= self.write_pos, "window extends past write_pos");
SharedBuffer::from_owner(BufferWindow {
backing: self.backing.clone(),
offset,
len,
})
}
#[cfg(feature = "io_uring")]
pub fn uring_region(&self, offset: usize, len: usize) -> ArenaRegion {
let end = offset
.checked_add(len)
.expect("StreamArena::uring_region range overflows");
assert!(end <= self.capacity, "uring region extends past capacity");
ArenaRegion {
backing: self.backing.clone(),
offset,
capacity: len,
filled: 0,
}
}
#[inline]
pub fn remaining(&self) -> usize {
self.capacity - self.write_pos
}
pub fn ensure_capacity(&mut self, needed: usize) {
if self.capacity - self.write_pos >= needed {
return;
}
if needed > self.capacity {
self.capacity = needed.div_ceil(64) * 64;
self.backing = Arc::new(ArenaBacking {
data: UnsafeCell::new(Vec64::with_capacity(self.capacity)),
});
self.write_pos = 0;
} else {
self.recycle_or_reset();
}
}
#[inline]
pub fn write_pos(&self) -> usize {
self.write_pos
}
#[inline]
pub fn recycle_if_free(&mut self) {
if Arc::strong_count(&self.backing) == 1 {
self.write_pos = 0;
}
}
pub fn recycle_or_reset(&mut self) {
if Arc::strong_count(&self.backing) == 1 {
self.write_pos = 0;
} else {
self.backing = Arc::new(ArenaBacking {
data: UnsafeCell::new(Vec64::with_capacity(self.capacity)),
});
self.write_pos = 0;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn write_and_window() {
let mut arena = StreamArena::with_capacity(1024);
assert_eq!(arena.remaining(), 1024);
let start = arena.write_pos();
arena.extend_from_slice(b"hello").unwrap();
let shared = arena.window(start, 5);
assert_eq!(shared.as_slice(), b"hello");
assert_eq!(arena.write_pos(), 5);
arena.align();
assert_eq!(arena.write_pos(), 64);
}
#[test]
fn multiple_windows() {
let mut arena = StreamArena::with_capacity(1024);
let start1 = arena.write_pos();
arena.extend_from_slice(b"abc").unwrap();
let w1 = arena.window(start1, 3);
arena.align();
let start2 = arena.write_pos();
assert_eq!(start2 % 64, 0);
arena.extend_from_slice(b"def").unwrap();
let w2 = arena.window(start2, 3);
arena.align();
assert_eq!(w1.as_slice(), b"abc");
assert_eq!(w2.as_slice(), b"def");
}
#[test]
fn alignment_padding_is_initialised() {
let mut arena = StreamArena::with_capacity(64);
arena.extend_from_slice(b"x").unwrap();
arena.align();
assert_eq!(arena.window(1, 63).as_slice(), &[0; 63]);
}
#[test]
fn recycle_when_all_dropped() {
let mut arena = StreamArena::with_capacity(256);
arena.extend_from_slice(&[1u8; 10]).unwrap();
let w = arena.window(0, 10);
arena.align();
assert_eq!(arena.write_pos(), 64);
arena.recycle_or_reset();
assert_eq!(arena.write_pos(), 0);
assert_eq!(w.as_slice(), &[1u8; 10]);
drop(w);
}
#[test]
fn recycle_reuses_allocation() {
let mut arena = StreamArena::with_capacity(256);
arena.extend_from_slice(&[1u8; 10]).unwrap();
{
let w = arena.window(0, 10);
assert_eq!(w.as_slice(), &[1u8; 10]);
}
let backing_ptr_before = Arc::as_ptr(&arena.backing);
arena.recycle_or_reset();
let backing_ptr_after = Arc::as_ptr(&arena.backing);
assert_eq!(
backing_ptr_before, backing_ptr_after,
"should reuse same backing"
);
assert_eq!(arena.write_pos(), 0);
}
#[test]
fn arena_fills_then_rolls_over() {
let mut arena = StreamArena::with_capacity(128);
arena.extend_from_slice(&[42u8; 64]).unwrap();
let w = arena.window(0, 64);
arena.align();
arena.extend_from_slice(&[43u8; 64]).unwrap();
assert_eq!(arena.remaining(), 0);
arena.recycle_or_reset();
assert_eq!(arena.write_pos(), 0);
assert_eq!(w.as_slice(), &[42u8; 64]);
let start = arena.write_pos();
arena.extend_from_slice(b"new!").unwrap();
let w2 = arena.window(start, 4);
assert_eq!(w2.as_slice(), b"new!");
}
#[test]
fn windows_are_64_byte_aligned() {
let mut arena = StreamArena::with_capacity(4096);
for i in 0..3 {
let start = arena.write_pos();
assert_eq!(start % 64, 0, "window {i} start not 64-byte aligned");
let data = vec![(i + 1) as u8; 100];
arena.extend_from_slice(&data).unwrap();
let w = arena.window(start, 100);
assert_eq!(w.as_slice(), &data);
arena.align();
}
}
#[test]
fn multi_read_payload_is_contiguous() {
let mut arena = StreamArena::with_capacity(4096);
let start = arena.write_pos();
arena.extend_from_slice(&[1u8; 10]).unwrap();
arena.extend_from_slice(&[2u8; 10]).unwrap();
arena.extend_from_slice(&[3u8; 10]).unwrap();
let w = arena.window(start, 30);
assert_eq!(w.len(), 30);
assert_eq!(&w.as_slice()[..10], &[1u8; 10]);
assert_eq!(&w.as_slice()[10..20], &[2u8; 10]);
assert_eq!(&w.as_slice()[20..30], &[3u8; 10]);
arena.align();
assert_eq!(arena.write_pos() % 64, 0);
}
#[test]
#[should_panic(expected = "advance past capacity")]
fn advance_past_capacity_panics() {
let mut arena = StreamArena::with_capacity(64);
unsafe { arena.advance(65) };
}
#[test]
#[should_panic(expected = "window extends past write_pos")]
fn window_past_write_pos_panics() {
let mut arena = StreamArena::with_capacity(64);
arena.extend_from_slice(&[1u8; 8]).unwrap();
let _ = arena.window(0, 9);
}
#[test]
fn extend_past_capacity_errors() {
let mut arena = StreamArena::with_capacity(8);
assert!(arena.extend_from_slice(&[0u8; 9]).is_err());
assert_eq!(arena.write_pos(), 0);
}
}