use std::{alloc::Layout, marker::PhantomData, ptr::NonNull};
use crate::num::{Align, Bytes};
#[derive(Debug)]
pub(crate) struct Buffer {
ptr: NonNull<u8>,
stride: Bytes,
entries: usize,
layout: Layout,
}
impl Buffer {
pub(crate) fn new(
entries: usize,
bytes_per_entry: Bytes,
align: Align,
) -> Result<Self, BufferError> {
let bytes = bytes_per_entry.checked_mul(entries).ok_or(BufferError)?;
let layout = std::alloc::Layout::from_size_align(bytes.value(), align.value())
.map_err(|_: std::alloc::LayoutError| BufferError)?;
let ptr = if layout.size() == 0 {
std::ptr::dangling_mut()
} else {
unsafe { std::alloc::alloc_zeroed(layout) }
};
let ptr = match NonNull::new(ptr) {
Some(ptr) => ptr,
None => std::alloc::handle_alloc_error(layout),
};
Ok(Self {
ptr,
stride: bytes_per_entry,
entries,
layout,
})
}
#[inline]
pub(crate) fn len(&self) -> usize {
self.entries
}
#[inline]
pub(crate) fn stride(&self) -> Bytes {
self.stride
}
#[inline]
pub(crate) unsafe fn get_unchecked(&self, i: usize) -> RawSlice<'_> {
debug_assert!(i < self.entries);
let ptr = unsafe { self.ptr.add(self.stride().value() * i) };
RawSlice {
ptr,
len: self.stride,
_lifetime: PhantomData,
}
}
#[cfg(test)]
pub(crate) fn get(&self, i: usize) -> Option<RawSlice<'_>> {
if i >= self.entries {
None
} else {
Some(unsafe { self.get_unchecked(i) })
}
}
#[cfg(test)]
fn as_ptr(&self) -> *const u8 {
self.ptr.as_ptr().cast_const()
}
}
impl Drop for Buffer {
fn drop(&mut self) {
if self.layout.size() != 0 {
unsafe { std::alloc::dealloc(self.ptr.as_ptr(), self.layout) }
}
}
}
#[derive(Debug)]
#[non_exhaustive]
pub(crate) struct BufferError;
impl std::fmt::Display for BufferError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("requested allocation exceeds `isize::MAX`")
}
}
impl std::error::Error for BufferError {}
unsafe impl Send for Buffer {}
unsafe impl Sync for Buffer {}
#[derive(Debug)]
pub(crate) struct RawSlice<'a> {
ptr: NonNull<u8>,
len: Bytes,
_lifetime: PhantomData<&'a ()>,
}
impl<'a> RawSlice<'a> {
unsafe fn new(ptr: NonNull<u8>, len: Bytes) -> Self {
Self {
ptr,
len,
_lifetime: PhantomData,
}
}
#[inline]
pub(crate) fn truncate(&self, n: Bytes) -> RawSlice<'a> {
unsafe { self.truncate_unchecked(self.len.min(n)) }
}
#[inline]
pub(crate) unsafe fn truncate_unchecked(&self, n: Bytes) -> RawSlice<'a> {
debug_assert!(n <= self.len);
unsafe { Self::new(self.ptr, n) }
}
#[inline]
pub(crate) fn split(&self, n: Bytes) -> (RawSlice<'a>, RawSlice<'a>) {
unsafe { self.split_unchecked(self.len.min(n)) }
}
#[inline]
pub(crate) unsafe fn split_unchecked(&self, n: Bytes) -> (RawSlice<'a>, RawSlice<'a>) {
debug_assert!(n <= self.len);
unsafe {
(
Self::new(self.ptr, n),
Self::new(self.ptr.add(n.value()), self.len.unchecked_sub(n)),
)
}
}
#[inline]
pub(crate) fn len(&self) -> Bytes {
self.len
}
pub(crate) fn as_non_null(&self) -> NonNull<u8> {
self.ptr
}
pub(crate) fn as_ptr(&self) -> *const u8 {
self.ptr.as_ptr().cast_const()
}
pub(crate) fn as_mut_ptr(&self) -> *mut u8 {
self.ptr.as_ptr()
}
#[inline]
pub(crate) unsafe fn as_slice(&self) -> &'a [u8] {
unsafe { std::slice::from_raw_parts(self.ptr.as_ptr(), self.len.value()) }
}
#[inline]
pub(crate) unsafe fn as_mut_slice(&mut self) -> &'a mut [u8] {
unsafe { std::slice::from_raw_parts_mut(self.ptr.as_ptr(), self.len.value()) }
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::{sync::Barrier, thread};
#[derive(Debug)]
struct Ctx {
entries: usize,
bytes_per_entry: Bytes,
align: Align,
}
impl std::fmt::Display for Ctx {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"entries = {}, bytes_per_entry = {}, align = {}",
self.entries, self.bytes_per_entry, self.align
)
}
}
fn test_buffer_inner(entries: usize, bytes_per_entry: Bytes, align: Align) {
let ctx = Ctx {
entries,
bytes_per_entry,
align,
};
let mut buffer = Buffer::new(entries, bytes_per_entry, align).unwrap();
assert_eq!(buffer.len(), entries, "{}", ctx);
assert_eq!(buffer.stride(), bytes_per_entry, "{}", ctx);
if entries != 0 && !bytes_per_entry.is_zero() {
let addr = buffer.as_ptr() as usize;
assert!(
addr.is_multiple_of(align.value()),
"pointer address {:#x} must be a multiple of the requested alignment: {}",
addr,
ctx,
);
}
assert_is_zeroed(&mut buffer, &ctx);
check_slice_methods(&mut buffer, &ctx);
zero(&mut buffer);
check_threaded(&mut buffer, &ctx);
}
fn zero(buffer: &mut Buffer) {
for i in 0..buffer.len() {
let mut raw_slice = buffer.get(i).unwrap();
assert_eq!(raw_slice.len(), buffer.stride());
let slice = unsafe { raw_slice.as_mut_slice() };
assert_eq!(slice.len(), buffer.stride().value());
slice.fill(0);
}
}
fn assert_is_zeroed(buffer: &mut Buffer, ctx: &Ctx) {
for i in 0..buffer.len() {
let raw_slice = buffer.get(i).unwrap();
assert_eq!(raw_slice.len(), buffer.stride());
assert_eq!(raw_slice.as_non_null().as_ptr(), raw_slice.as_mut_ptr());
assert_eq!(
raw_slice.as_non_null().as_ptr().cast_const(),
raw_slice.as_ptr()
);
assert_eq!(
raw_slice.as_ptr(),
buffer
.as_ptr()
.wrapping_add(buffer.stride().checked_mul(i).unwrap().value()),
"stride mismatch - {}",
ctx
);
let slice = unsafe { raw_slice.as_slice() };
assert_eq!(slice.len(), buffer.stride().value());
assert!(slice.iter().all(|&i| i == 0), "{}", ctx);
}
assert!(buffer.get(buffer.len()).is_none(), "{}", ctx);
}
fn check_slice_methods(buffer: &mut Buffer, ctx: &Ctx) {
if buffer.len() == 0 {
return;
}
let mut raw = buffer.get(0).unwrap();
let base: u8 = 5;
let base_usize: usize = base.into();
iota(unsafe { raw.as_mut_slice() }, base);
for i in 0..raw.len().value() + base_usize {
let expected = i.min(raw.len().value());
let truncated = raw.truncate(Bytes::new(i));
assert_eq!(truncated.len().value(), expected, "{}", ctx);
assert!(is_iota(unsafe { truncated.as_slice() }, base), "{}", ctx);
}
for i in 0..raw.len().value() + base_usize {
let first = i.min(raw.len().value());
let last = raw.len().value() - first;
let (mut prefix, mut suffix) = raw.split(Bytes::new(i));
assert_eq!(prefix.len().value(), first, "{}", ctx);
assert_eq!(suffix.len().value(), last, "{}", ctx);
assert!(is_iota(unsafe { prefix.as_slice() }, base), "{}", ctx);
assert!(
is_iota(unsafe { suffix.as_slice() }, base.wrapping_add(i as u8)),
"{}",
ctx
);
{
let prefix = unsafe { prefix.as_mut_slice() };
let suffix = unsafe { suffix.as_mut_slice() };
suffix.fill(0);
prefix.fill(0);
}
assert!(unsafe { raw.as_slice() }.iter().all(|i| *i == 0), "{}", ctx);
iota(unsafe { raw.as_mut_slice() }, base);
}
}
fn check_threaded(buffer: &mut Buffer, ctx: &Ctx) {
let spawns = buffer.len();
let pre = &Barrier::new(spawns);
let post = &Barrier::new(spawns);
{
let borrowed: &Buffer = buffer;
thread::scope(|s| {
for i in 0..spawns {
s.spawn(move || {
let slice = unsafe { borrowed.get(i).unwrap().as_mut_slice() };
pre.wait();
iota(slice, i as u8);
post.wait();
});
}
});
}
for i in 0..spawns {
let slice = unsafe { buffer.get(i).unwrap().as_slice() };
assert!(is_iota(slice, i as u8), "i = {} -- {}", i, ctx);
}
}
fn iota(x: &mut [u8], base: u8) {
for (i, v) in x.iter_mut().enumerate() {
*v = base.wrapping_add(i as u8);
}
}
#[must_use]
fn is_iota(x: &[u8], base: u8) -> bool {
for (i, v) in x.iter().enumerate() {
if *v != base.wrapping_add(i as u8) {
return false;
}
}
true
}
#[test]
fn test_buffer() {
let entries = [0, 1, 2, 5];
let bytes_per_entry = [0, 1, 2, 5, 10].map(Bytes::new);
let align = [Align::_1, Align::_64];
for entries in entries {
for bytes_per_entry in bytes_per_entry {
for align in align {
test_buffer_inner(entries, bytes_per_entry, align);
}
}
}
}
#[test]
fn test_buffer_overflow_mul() {
let result = Buffer::new(usize::MAX, Bytes::new(2), Align::_1);
assert!(result.is_err());
}
#[test]
fn test_buffer_overflow_layout() {
let result = Buffer::new(isize::MAX as usize, Bytes::new(2), Align::_1);
assert!(result.is_err());
}
}