use std::alloc::{alloc, dealloc, realloc, Layout};
use std::ptr::NonNull;
pub(crate) struct AlignedBuffer {
ptr: NonNull<u8>,
len: usize,
cap: usize,
item_size: usize,
align: usize,
}
unsafe impl Send for AlignedBuffer {}
unsafe impl Sync for AlignedBuffer {}
impl AlignedBuffer {
#[inline]
pub(crate) fn new(item_size: usize, align: usize) -> Self {
debug_assert!(align.is_power_of_two(), "align must be a power of two");
debug_assert!(item_size > 0, "item_size must be > 0");
AlignedBuffer {
ptr: NonNull::dangling(),
len: 0,
cap: 0,
item_size,
align,
}
}
pub(crate) fn with_capacity(item_size: usize, align: usize, capacity: usize) -> Self {
let mut buf = Self::new(item_size, align);
if capacity > 0 {
buf.grow_to(capacity);
}
buf
}
#[inline]
pub(crate) fn len(&self) -> usize {
self.len
}
#[inline]
pub(crate) fn is_empty(&self) -> bool {
self.len == 0
}
#[inline]
pub(crate) fn clear(&mut self) {
self.len = 0;
}
#[inline]
pub(crate) unsafe fn push<T: Copy>(&mut self, value: T) {
debug_assert_eq!(std::mem::size_of::<T>(), self.item_size);
debug_assert_eq!(std::mem::align_of::<T>(), self.align);
if self.len == self.cap {
self.grow();
}
unsafe {
let dst = self.ptr.as_ptr().add(self.len * self.item_size) as *mut T;
dst.write(value);
}
self.len += 1;
}
#[inline]
pub(crate) unsafe fn set_len(&mut self, new_len: usize) {
debug_assert!(new_len <= self.cap);
self.len = new_len;
}
#[inline]
pub(crate) unsafe fn as_slice<T>(&self) -> &[T] {
debug_assert_eq!(std::mem::size_of::<T>(), self.item_size);
debug_assert_eq!(std::mem::align_of::<T>(), self.align);
if self.len == 0 {
return &[];
}
unsafe { std::slice::from_raw_parts(self.ptr.as_ptr() as *const T, self.len) }
}
#[inline]
pub(crate) unsafe fn as_ptr_at(&self, index: usize) -> *const u8 {
debug_assert!(index < self.len);
unsafe { self.ptr.as_ptr().add(index * self.item_size) }
}
#[inline]
pub(crate) unsafe fn as_mut_ptr_at(&mut self, index: usize) -> *mut u8 {
debug_assert!(index < self.cap);
unsafe { self.ptr.as_ptr().add(index * self.item_size) }
}
pub(crate) fn reserve(&mut self, new_cap: usize) {
if new_cap > self.cap {
self.grow_to(new_cap);
}
}
pub(crate) fn extend_from(&mut self, other: &AlignedBuffer) {
debug_assert_eq!(self.item_size, other.item_size);
debug_assert_eq!(self.align, other.align);
if other.len == 0 {
return;
}
let needed = self.len + other.len;
if needed > self.cap {
self.grow_to(needed);
}
unsafe {
let dst = self.ptr.as_ptr().add(self.len * self.item_size);
let src = other.ptr.as_ptr();
std::ptr::copy_nonoverlapping(src, dst, other.len * other.item_size);
}
self.len += other.len;
}
#[cfg_attr(not(feature = "messaging_gpu"), allow(dead_code))]
#[inline]
pub(crate) fn as_bytes(&self) -> &[u8] {
if self.len == 0 {
return &[];
}
unsafe { std::slice::from_raw_parts(self.ptr.as_ptr(), self.len * self.item_size) }
}
#[cfg_attr(not(feature = "messaging_gpu"), allow(dead_code))]
pub(crate) fn extend_from_bytes(&mut self, bytes: &[u8]) {
debug_assert_eq!(bytes.len() % self.item_size, 0);
if bytes.is_empty() {
return;
}
let items = bytes.len() / self.item_size;
let needed = self.len + items;
if needed > self.cap {
self.grow_to(needed);
}
unsafe {
let dst = self.ptr.as_ptr().add(self.len * self.item_size);
std::ptr::copy_nonoverlapping(bytes.as_ptr(), dst, bytes.len());
}
self.len += items;
}
fn grow(&mut self) {
let new_cap = if self.cap == 0 { 4 } else { self.cap * 2 };
self.grow_to(new_cap);
}
fn grow_to(&mut self, new_cap: usize) {
debug_assert!(new_cap > self.cap);
let new_layout = Layout::from_size_align(new_cap * self.item_size, self.align)
.expect("AlignedBuffer: layout overflow");
let new_ptr = if self.cap == 0 {
unsafe { alloc(new_layout) }
} else {
let old_layout = Layout::from_size_align(self.cap * self.item_size, self.align)
.expect("AlignedBuffer: old layout overflow");
unsafe { realloc(self.ptr.as_ptr(), old_layout, new_layout.size()) }
};
self.ptr = NonNull::new(new_ptr).expect("AlignedBuffer: allocation failed");
self.cap = new_cap;
}
}
impl Drop for AlignedBuffer {
fn drop(&mut self) {
if self.cap > 0 {
let layout = Layout::from_size_align(self.cap * self.item_size, self.align)
.expect("AlignedBuffer: drop layout overflow");
unsafe { dealloc(self.ptr.as_ptr(), layout) };
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn push_and_read_u32() {
let mut buf = AlignedBuffer::new(std::mem::size_of::<u32>(), std::mem::align_of::<u32>());
for i in 0u32..8 {
unsafe { buf.push(i) };
}
assert_eq!(buf.len(), 8);
let slice: &[u32] = unsafe { buf.as_slice() };
assert_eq!(slice, &[0, 1, 2, 3, 4, 5, 6, 7]);
}
#[test]
fn push_and_read_aligned_struct() {
#[repr(C, align(16))]
#[derive(Copy, Clone, Debug, PartialEq)]
struct Wide {
x: f64,
y: f64,
}
let mut buf = AlignedBuffer::new(std::mem::size_of::<Wide>(), std::mem::align_of::<Wide>());
unsafe { buf.push(Wide { x: 1.0, y: 2.0 }) };
unsafe { buf.push(Wide { x: 3.0, y: 4.0 }) };
let slice: &[Wide] = unsafe { buf.as_slice() };
assert_eq!(slice[0], Wide { x: 1.0, y: 2.0 });
assert_eq!(slice[1], Wide { x: 3.0, y: 4.0 });
}
#[test]
fn extend_from() {
let mut a = AlignedBuffer::new(4, 4);
let mut b = AlignedBuffer::new(4, 4);
for i in 0u32..4 {
unsafe { a.push(i) }
}
for i in 4u32..8 {
unsafe { b.push(i) }
}
a.extend_from(&b);
assert_eq!(a.len(), 8);
let slice: &[u32] = unsafe { a.as_slice() };
assert_eq!(slice, &[0, 1, 2, 3, 4, 5, 6, 7]);
}
#[test]
fn clear_and_reuse() {
let mut buf = AlignedBuffer::new(4, 4);
for i in 0u32..4 {
unsafe { buf.push(i) }
}
buf.clear();
assert_eq!(buf.len(), 0);
unsafe { buf.push(99u32) };
let slice: &[u32] = unsafe { buf.as_slice() };
assert_eq!(slice, &[99]);
}
}