use core::fmt;
const FIXED_ARRAY_SCHEMA_PREFIX: u32 = 0x0100_0000;
const FIXED_ARRAY_SCHEMA_MAX_LEN: usize = 0x00ff_ffff;
#[inline(always)]
pub(crate) const fn fixed_array_schema_id(len: usize) -> u32 {
if len > FIXED_ARRAY_SCHEMA_MAX_LEN {
crate::invariant();
}
FIXED_ARRAY_SCHEMA_PREFIX | len as u32
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CodecError {
Truncated,
Malformed,
}
#[inline]
pub(crate) fn require_exact_len(actual: usize, expected: usize) -> Result<(), CodecError> {
if actual < expected {
Err(CodecError::Truncated)
} else if actual == expected {
Ok(())
} else {
Err(CodecError::Malformed)
}
}
pub trait WireEncode {
fn encode_into(&self, out: &mut [u8]) -> Result<usize, CodecError>;
}
#[inline(always)]
pub(crate) const fn erased_encoder<P: WireEncode>()
-> unsafe fn(*const (), &mut [u8]) -> Result<usize, CodecError> {
encode_erased::<P>
}
#[inline(always)]
unsafe fn encode_erased<P: WireEncode>(
ptr: *const (),
scratch: &mut [u8],
) -> Result<usize, CodecError> {
let payload = unsafe { &*ptr.cast::<P>() };
payload.encode_into(scratch)
}
pub trait WirePayload {
const SCHEMA_ID: u32;
type Decoded<'a>;
fn validate_payload(input: Payload<'_>) -> Result<(), CodecError>;
fn decode_validated_payload<'a>(input: Payload<'a>) -> Self::Decoded<'a>;
}
impl WireEncode for () {
fn encode_into(&self, out: &mut [u8]) -> Result<usize, CodecError> {
Ok(out[..0].len())
}
}
impl WirePayload for () {
const SCHEMA_ID: u32 = 0;
type Decoded<'a> = Self;
fn validate_payload(input: Payload<'_>) -> Result<(), CodecError> {
require_exact_len(input.as_bytes().len(), 0)
}
fn decode_validated_payload<'a>(input: Payload<'a>) -> Self::Decoded<'a> {
if !input.as_bytes().is_empty() {
crate::invariant();
}
}
}
impl WireEncode for bool {
fn encode_into(&self, out: &mut [u8]) -> Result<usize, CodecError> {
if out.is_empty() {
return Err(CodecError::Truncated);
}
out[0] = if *self { 1 } else { 0 };
Ok(1)
}
}
impl WirePayload for bool {
const SCHEMA_ID: u32 = 1;
type Decoded<'a> = Self;
fn validate_payload(input: Payload<'_>) -> Result<(), CodecError> {
let bytes = input.as_bytes();
require_exact_len(bytes.len(), 1)?;
match bytes[0] {
0 | 1 => Ok(()),
_ => Err(CodecError::Malformed),
}
}
fn decode_validated_payload<'a>(input: Payload<'a>) -> Self::Decoded<'a> {
input.as_bytes()[0] != 0
}
}
macro_rules! impl_wire_for_int {
($ty:ty, $len:expr, $schema:expr) => {
impl WireEncode for $ty {
fn encode_into(&self, out: &mut [u8]) -> Result<usize, CodecError> {
if out.len() < $len {
return Err(CodecError::Truncated);
}
out[..$len].copy_from_slice(&self.to_be_bytes());
Ok($len)
}
}
impl WirePayload for $ty {
const SCHEMA_ID: u32 = $schema;
type Decoded<'a> = Self;
fn validate_payload(input: Payload<'_>) -> Result<(), CodecError> {
require_exact_len(input.as_bytes().len(), $len)
}
fn decode_validated_payload<'a>(input: Payload<'a>) -> Self::Decoded<'a> {
let bytes = input.as_bytes();
let mut buf = [0u8; $len];
buf.copy_from_slice(&bytes[..$len]);
<$ty>::from_be_bytes(buf)
}
}
};
}
impl_wire_for_int!(u8, 1, 2);
impl_wire_for_int!(i8, 1, 3);
impl_wire_for_int!(u16, 2, 4);
impl_wire_for_int!(i16, 2, 5);
impl_wire_for_int!(u32, 4, 6);
impl_wire_for_int!(i32, 4, 7);
impl_wire_for_int!(u64, 8, 8);
impl_wire_for_int!(i64, 8, 9);
impl_wire_for_int!(u128, 16, 10);
impl_wire_for_int!(i128, 16, 11);
impl WireEncode for &[u8] {
fn encode_into(&self, out: &mut [u8]) -> Result<usize, CodecError> {
if out.len() < self.len() {
return Err(CodecError::Truncated);
}
out[..self.len()].copy_from_slice(self);
Ok(self.len())
}
}
impl WirePayload for &[u8] {
const SCHEMA_ID: u32 = 12;
type Decoded<'a> = &'a [u8];
fn validate_payload(input: Payload<'_>) -> Result<(), CodecError> {
input.as_bytes();
Ok(())
}
fn decode_validated_payload<'a>(input: Payload<'a>) -> Self::Decoded<'a> {
input.as_bytes()
}
}
impl<const N: usize> WireEncode for [u8; N] {
fn encode_into(&self, out: &mut [u8]) -> Result<usize, CodecError> {
if out.len() < N {
return Err(CodecError::Truncated);
}
out[..N].copy_from_slice(self);
Ok(N)
}
}
impl<const N: usize> WirePayload for [u8; N] {
const SCHEMA_ID: u32 = fixed_array_schema_id(N);
type Decoded<'a> = Self;
fn validate_payload(input: Payload<'_>) -> Result<(), CodecError> {
require_exact_len(input.as_bytes().len(), N)
}
fn decode_validated_payload<'a>(input: Payload<'a>) -> Self::Decoded<'a> {
let bytes = input.as_bytes();
let mut buf = [0u8; N];
buf.copy_from_slice(&bytes[..N]);
buf
}
}
#[cfg(test)]
mod tests {
use super::*;
fn decode<'a, P: WirePayload>(input: &'a [u8]) -> Result<P::Decoded<'a>, CodecError> {
let payload = Payload::new(input);
P::validate_payload(payload)?;
Ok(P::decode_validated_payload(payload))
}
#[test]
fn fixed_payload_decoders_reject_trailing_bytes() {
assert_eq!(decode::<()>(&[]), Ok(()));
assert_eq!(decode::<()>(&[1]), Err(CodecError::Malformed));
assert_eq!(decode::<bool>(&[1]), Ok(true));
assert_eq!(decode::<bool>(&[1, 0]), Err(CodecError::Malformed));
assert_eq!(decode::<u16>(&[0x12, 0x34]), Ok(0x1234));
assert_eq!(
decode::<u16>(&[0x12, 0x34, 0x56]),
Err(CodecError::Malformed)
);
assert_eq!(decode::<[u8; 2]>(&[7, 9]), Ok([7, 9]));
assert_eq!(decode::<[u8; 2]>(&[7, 9, 11]), Err(CodecError::Malformed));
}
#[test]
fn borrowed_byte_slice_remains_variable_length() {
let bytes = [1, 2, 3];
assert_eq!(decode::<&[u8]>(&bytes), Ok(&bytes[..]));
}
#[test]
fn builtin_payload_schemas_are_pairwise_distinct() {
let schemas = [
<() as WirePayload>::SCHEMA_ID,
<bool as WirePayload>::SCHEMA_ID,
<u8 as WirePayload>::SCHEMA_ID,
<i8 as WirePayload>::SCHEMA_ID,
<u16 as WirePayload>::SCHEMA_ID,
<i16 as WirePayload>::SCHEMA_ID,
<u32 as WirePayload>::SCHEMA_ID,
<i32 as WirePayload>::SCHEMA_ID,
<u64 as WirePayload>::SCHEMA_ID,
<i64 as WirePayload>::SCHEMA_ID,
<u128 as WirePayload>::SCHEMA_ID,
<i128 as WirePayload>::SCHEMA_ID,
<&[u8] as WirePayload>::SCHEMA_ID,
<[u8; 0] as WirePayload>::SCHEMA_ID,
<[u8; 4] as WirePayload>::SCHEMA_ID,
];
for (index, schema) in schemas.iter().enumerate() {
assert!(!schemas[..index].contains(schema));
}
assert_eq!(<[u8; 0] as WirePayload>::SCHEMA_ID, 0x0100_0000);
assert_eq!(<[u8; 4] as WirePayload>::SCHEMA_ID, 0x0100_0004);
assert_eq!(<[u8; 0x00ff_ffff] as WirePayload>::SCHEMA_ID, 0x01ff_ffff);
}
#[test]
fn fixed_array_schema_identity_is_exact_at_its_domain_boundary() {
assert_eq!(fixed_array_schema_id(0), 0x0100_0000);
assert_eq!(fixed_array_schema_id(0x00ff_ffff), 0x01ff_ffff);
}
#[test]
#[should_panic]
fn fixed_array_schema_identity_rejects_the_first_colliding_width() {
let _ = fixed_array_schema_id(0x0100_0000);
}
}
#[derive(Clone, Copy)]
pub struct Payload<'a> {
data: &'a [u8],
}
impl<'a> Payload<'a> {
#[inline]
pub const fn new(data: &'a [u8]) -> Self {
Self { data }
}
#[inline]
pub fn as_bytes(&self) -> &'a [u8] {
self.data
}
}
impl<'a> fmt::Debug for Payload<'a> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let bytes = self.as_bytes();
let preview_len = if bytes.len() > 32 { 32 } else { bytes.len() };
f.debug_struct("Payload")
.field("len", &bytes.len())
.field("preview", &&bytes[..preview_len])
.finish()
}
}
#[cfg(kani)]
mod kani;