use crate::{Codec, CodecFor};
use core::fmt;
pub struct RkyvCodec<const SCRATCH: usize>;
pub unsafe trait TrustedArchive: rkyv::Archive {}
macro_rules! trusted_archive_primitives {
($($ty:ty),* $(,)?) => {
$(
unsafe impl TrustedArchive for $ty {}
)*
};
}
trusted_archive_primitives!(
(),
bool,
char,
u8,
u16,
u32,
u64,
u128,
i8,
i16,
i32,
i64,
i128,
f32,
f64,
);
unsafe impl<T: TrustedArchive, const N: usize> TrustedArchive for [T; N] {}
macro_rules! impl_trusted_archive_tuple {
($($T:ident),+) => {
unsafe impl<$($T: TrustedArchive),+> TrustedArchive for ($($T,)+) {}
};
}
impl_trusted_archive_tuple!(T0);
impl_trusted_archive_tuple!(T0, T1);
impl_trusted_archive_tuple!(T0, T1, T2);
impl_trusted_archive_tuple!(T0, T1, T2, T3);
impl_trusted_archive_tuple!(T0, T1, T2, T3, T4);
impl_trusted_archive_tuple!(T0, T1, T2, T3, T4, T5);
impl_trusted_archive_tuple!(T0, T1, T2, T3, T4, T5, T6);
impl_trusted_archive_tuple!(T0, T1, T2, T3, T4, T5, T6, T7);
impl<const SCRATCH: usize> Codec for RkyvCodec<SCRATCH> {
type Error = rkyv::rancor::Error;
}
#[derive(Debug)]
struct BufferTooSmall;
impl fmt::Display for BufferTooSmall {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("destination buffer is too small")
}
}
impl core::error::Error for BufferTooSmall {}
fn buffer_too_small() -> rkyv::rancor::Error {
<rkyv::rancor::Error as rkyv::rancor::Source>::new(BufferTooSmall)
}
#[cfg(feature = "alloc")]
use rkyv::{
Archive, Serialize,
api::high::{HighSerializer, HighValidator},
bytecheck::CheckBytes,
rancor::Error as RkyvError,
ser::allocator::ArenaHandle,
util::AlignedVec,
};
#[cfg(feature = "alloc")]
impl<T, const SCRATCH: usize> CodecFor<T> for RkyvCodec<SCRATCH>
where
T: Archive + for<'a> Serialize<HighSerializer<AlignedVec, ArenaHandle<'a>, RkyvError>>,
T::Archived: for<'buf> CheckBytes<HighValidator<'buf, RkyvError>>,
{
type Decoded<'buf>
= &'buf T::Archived
where
T: 'buf;
fn encode(msg: &T, buf: &mut [u8]) -> Result<usize, Self::Error> {
let bytes = rkyv::to_bytes::<RkyvError>(msg)?;
let len = bytes.len();
if len > buf.len() {
consortium_log::trace!(
"rkyv encode needs {} bytes but buffer is {}",
len,
buf.len()
);
return Err(buffer_too_small());
}
buf[..len].copy_from_slice(&bytes);
Ok(len)
}
fn decode<'buf>(buf: &'buf [u8]) -> Result<Self::Decoded<'buf>, Self::Error>
where
T: 'buf,
{
rkyv::access::<T::Archived, RkyvError>(buf).inspect_err(|_| {
consortium_log::trace!("rkyv access/validation failed on {} bytes", buf.len());
})
}
}
#[cfg(not(feature = "alloc"))]
use rkyv::{
Serialize,
api::low::{LowSerializer, to_bytes_in_with_alloc},
rancor,
ser::{allocator::SubAllocator, writer::Buffer},
util::Align,
};
#[cfg(not(feature = "alloc"))]
impl<T, const SCRATCH: usize> CodecFor<T> for RkyvCodec<SCRATCH>
where
T: TrustedArchive
+ for<'a> Serialize<LowSerializer<Buffer<'a>, SubAllocator<'a>, rancor::Error>>,
{
type Decoded<'buf>
= &'buf T::Archived
where
T: 'buf;
fn encode(msg: &T, buf: &mut [u8]) -> Result<usize, Self::Error> {
use core::mem::MaybeUninit;
let mut output = Align([MaybeUninit::<u8>::uninit(); SCRATCH]);
let mut alloc = [MaybeUninit::<u8>::uninit(); SCRATCH];
let bytes = to_bytes_in_with_alloc::<_, _, rancor::Error>(
msg,
Buffer::from(&mut *output),
SubAllocator::new(&mut alloc),
)?;
let len = bytes.len();
if len > buf.len() {
consortium_log::trace!(
"rkyv encode needs {} bytes but buffer is {}",
len,
buf.len()
);
return Err(buffer_too_small());
}
buf[..len].copy_from_slice(&bytes);
Ok(len)
}
fn decode<'buf>(buf: &'buf [u8]) -> Result<Self::Decoded<'buf>, Self::Error>
where
T: 'buf,
{
Ok(unsafe { rkyv::access_unchecked::<T::Archived>(buf) })
}
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_trusted_archive<T: TrustedArchive>() {}
#[test]
fn fixed_size_primitives_are_trusted_archives() {
assert_trusted_archive::<()>();
assert_trusted_archive::<bool>();
assert_trusted_archive::<char>();
assert_trusted_archive::<u8>();
assert_trusted_archive::<u16>();
assert_trusted_archive::<u32>();
assert_trusted_archive::<u64>();
assert_trusted_archive::<u128>();
assert_trusted_archive::<i8>();
assert_trusted_archive::<i16>();
assert_trusted_archive::<i32>();
assert_trusted_archive::<i64>();
assert_trusted_archive::<i128>();
assert_trusted_archive::<f32>();
assert_trusted_archive::<f64>();
}
#[test]
fn round_trips_archived_value() {
let msg = 0x1234_5678u32;
let mut buf = [0u8; 16];
let len = RkyvCodec::<64>::encode(&msg, &mut buf).expect("encode should fit");
let decoded =
<RkyvCodec<64> as CodecFor<u32>>::decode(&buf[..len]).expect("decode should succeed");
assert_eq!(*decoded, msg);
}
#[test]
fn rejects_buffer_that_is_too_small() {
let msg = 0x1234_5678u32;
let mut buf = [0u8; 1];
assert!(RkyvCodec::<64>::encode(&msg, &mut buf).is_err());
}
#[cfg(feature = "alloc")]
#[test]
fn rejects_invalid_archived_bytes() {
let buf = [0u8; 1];
assert!(<RkyvCodec<64> as CodecFor<u32>>::decode(&buf).is_err());
}
}