use core::alloc::Layout;
use core::fmt;
use core::mem::{align_of, size_of};
use core::ops::Range;
use core::slice::{self, SliceIndex};
use crate::error::{Error, ErrorKind};
use crate::owned_buf::is_aligned_to;
use crate::r#ref::Ref;
use crate::r#unsized::Unsized;
use crate::slice::Slice;
use crate::zero_copy::{UnsizedZeroCopy, ZeroCopy};
pub trait BufMut {
fn extend_from_slice(&mut self, bytes: &[u8]) -> Result<(), Error>;
fn write<T: ?Sized>(&mut self, value: &T) -> Result<(), Error>
where
T: ZeroCopy;
}
impl<B: ?Sized> BufMut for &mut B
where
B: BufMut,
{
fn extend_from_slice(&mut self, bytes: &[u8]) -> Result<(), Error> {
(**self).extend_from_slice(bytes)
}
fn write<T: ?Sized>(&mut self, value: &T) -> Result<(), Error>
where
T: ZeroCopy,
{
(**self).write(value)
}
}
pub unsafe trait AnyRef {
type Target: ?Sized;
fn read_from<'buf>(&self, buf: &'buf Buf) -> Result<&'buf Self::Target, Error>;
}
pub trait AnyValue {
type Target: ?Sized;
fn visit<V, O>(&self, buf: &Buf, visitor: V) -> Result<O, Error>
where
V: FnOnce(&Self::Target) -> O;
}
impl<T: ?Sized> AnyValue for T
where
T: AnyRef,
{
type Target = T::Target;
fn visit<V, O>(&self, buf: &Buf, visitor: V) -> Result<O, Error>
where
V: FnOnce(&Self::Target) -> O,
{
let value = buf.load(self)?;
Ok(visitor(value))
}
}
unsafe impl<T: ?Sized> AnyRef for &T
where
T: AnyRef,
{
type Target = T::Target;
#[inline]
fn read_from<'buf>(&self, buf: &'buf Buf) -> Result<&'buf Self::Target, Error> {
T::read_from(self, buf)
}
}
unsafe impl<T: ?Sized> AnyRef for Unsized<T>
where
T: UnsizedZeroCopy,
{
type Target = T;
fn read_from<'buf>(&self, buf: &'buf Buf) -> Result<&'buf Self::Target, Error> {
buf.load_unsized(*self)
}
}
unsafe impl<T: ?Sized> AnyRef for Ref<T>
where
T: ZeroCopy,
{
type Target = T;
fn read_from<'buf>(&self, buf: &'buf Buf) -> Result<&'buf Self::Target, Error> {
buf.load_sized(*self)
}
}
unsafe impl<T> AnyRef for Slice<T>
where
T: ZeroCopy,
{
type Target = [T];
fn read_from<'buf>(&self, buf: &'buf Buf) -> Result<&'buf Self::Target, Error> {
buf.load_slice(*self)
}
}
#[repr(transparent)]
pub struct Buf {
data: [u8],
}
impl Buf {
pub unsafe fn new_unchecked<T>(data: &T) -> &Buf
where
T: ?Sized + AsRef<[u8]>,
{
unsafe { &*(data.as_ref() as *const _ as *const Buf) }
}
pub fn as_ptr(&self) -> *const u8 {
self.data.as_ptr()
}
pub fn as_bytes(&self) -> &[u8] {
&self.data
}
pub fn range(&self) -> Range<usize> {
let range = self.data.as_ptr_range();
range.start as usize..range.end as usize
}
pub fn is_compatible(&self, layout: Layout) -> bool {
self.is_aligned_to(layout.align()) && self.data.len() == layout.size()
}
pub fn is_aligned_to(&self, align: usize) -> bool {
is_aligned_to(self.data.as_ptr(), align)
}
pub fn is_zeroed(&self) -> bool {
self.data.iter().all(|b| *b == 0)
}
pub fn len(&self) -> usize {
self.data.len()
}
pub fn is_empty(&self) -> bool {
self.data.is_empty()
}
pub fn get<I>(&self, range: I, align: usize) -> Result<&Buf, Error>
where
I: SliceIndex<[u8], Output = [u8]>,
{
let Some(data) = self.data.get(range) else {
return Err(Error::new(ErrorKind::OutOfBounds {
len: self.data.len(),
}));
};
if !is_aligned_to(data.as_ptr(), align) {
return Err(Error::new(ErrorKind::BadAlignment {
ptr: data.as_ptr() as usize,
align,
}));
}
Ok(unsafe { Buf::new_unchecked(data) })
}
pub fn load_unsized<T: ?Sized>(&self, ptr: Unsized<T>) -> Result<&T, Error>
where
T: UnsizedZeroCopy,
{
let start = ptr.ptr().offset();
let end = start.wrapping_add(ptr.size());
T::read_from(self.get(start..end, T::ALIGN)?)
}
pub fn load_sized<T: ?Sized>(&self, ptr: Ref<T>) -> Result<&T, Error>
where
T: ZeroCopy,
{
let start = ptr.ptr().offset();
let end = start.wrapping_add(size_of::<T>());
T::read_from(self.get(start..end, align_of::<T>())?)
}
pub fn load_slice<T>(&self, ptr: Slice<T>) -> Result<&[T], Error>
where
T: ZeroCopy,
{
let start = ptr.ptr().offset();
let end = start.wrapping_add(ptr.len().wrapping_mul(size_of::<T>()));
let buf = self.get(start..end, align_of::<T>())?;
validate_array::<T>(buf, ptr.len())?;
Ok(unsafe { slice::from_raw_parts(buf.as_ptr().cast(), ptr.len()) })
}
pub fn load<T>(&self, ptr: T) -> Result<&T::Target, Error>
where
T: AnyRef,
{
ptr.read_from(self)
}
pub unsafe fn cast<T>(&self) -> &T {
&*self.data.as_ptr().cast()
}
pub fn validate<T>(&self) -> Result<Validator<'_>, Error> {
if !self.is_compatible(Layout::new::<T>()) {
return Err(Error::new(ErrorKind::LayoutMismatch {
layout: Layout::new::<T>(),
buf: self.range(),
}));
}
Ok(Validator { data: &self.data })
}
pub unsafe fn validate_aligned(&self) -> Result<Validator<'_>, Error> {
Ok(Validator { data: &self.data })
}
}
pub(crate) fn validate_array<T>(buf: &Buf, len: usize) -> Result<(), Error>
where
T: ZeroCopy,
{
let layout =
Layout::array::<T>(len).map_err(|error| Error::new(ErrorKind::LayoutError { error }))?;
if !buf.is_compatible(layout) {
return Err(Error::new(ErrorKind::LayoutMismatch {
layout,
buf: buf.range(),
}));
}
validate_array_aligned::<T>(buf)?;
Ok(())
}
pub(crate) fn validate_array_aligned<T>(buf: &Buf) -> Result<(), Error>
where
T: ZeroCopy,
{
if !T::ANY_BITS && size_of::<T>() > 0 {
for chunk in buf.as_bytes().chunks_exact(size_of::<T>()) {
unsafe {
T::validate_aligned(Buf::new_unchecked(chunk))?;
}
}
}
Ok(())
}
impl fmt::Debug for Buf {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_tuple("Buf").field(&self.data.len()).finish()
}
}
#[must_use = "Must call `Validator::end` when validation is completed"]
pub struct Validator<'a> {
data: &'a [u8],
}
impl Validator<'_> {
pub fn field<T>(&mut self) -> Result<(), Error>
where
T: ZeroCopy,
{
let rem = (self.data.as_ptr() as usize) % align_of::<T>();
if rem != 0 {
let start = align_of::<T>() - rem;
let Some(d) = self.data.get(start..) else {
return Err(Error::new(ErrorKind::OutOfStartBound {
start,
len: self.data.len(),
}));
};
self.data = d;
}
if size_of::<T>() > self.data.len() {
return Err(Error::new(ErrorKind::OutOfStartBound {
start: size_of::<T>(),
len: self.data.len(),
}));
}
let (head, tail) = self.data.split_at(size_of::<T>());
unsafe {
T::validate_aligned(Buf::new_unchecked(head))?;
};
self.data = tail;
Ok(())
}
pub fn end(self) -> Result<(), Error> {
Ok(())
}
}