use std::{mem::MaybeUninit, sync::Arc};
use hybrid_array::{Array, ArraySize};
use super::{packed_size, read_packed_le_bytes, write_packed_le_bytes, InPlaceCodec};
use crate::errors::PrimitiveError;
unsafe impl<T: InPlaceCodec> InPlaceCodec for Arc<T> {
const ENCODED_SIZE: usize = T::ENCODED_SIZE;
fn write_le_bytes(&self, out: &mut [MaybeUninit<u8>]) {
(**self).write_le_bytes(out);
}
fn read_le_bytes(bytes: &[u8]) -> Result<Self, PrimitiveError> {
T::read_le_bytes(bytes).map(Arc::new)
}
}
unsafe impl<A: InPlaceCodec, B: InPlaceCodec> InPlaceCodec for (A, B) {
const ENCODED_SIZE: usize = A::ENCODED_SIZE + B::ENCODED_SIZE;
fn write_le_bytes(&self, out: &mut [MaybeUninit<u8>]) {
let (a, b) = out.split_at_mut(A::ENCODED_SIZE);
self.0.write_le_bytes(a);
self.1.write_le_bytes(b);
}
fn read_le_bytes(bytes: &[u8]) -> Result<Self, PrimitiveError> {
let (a, b) = bytes.split_at(A::ENCODED_SIZE);
Ok((A::read_le_bytes(a)?, B::read_le_bytes(b)?))
}
}
unsafe impl<T: InPlaceCodec, const N: usize> InPlaceCodec for [T; N] {
const ENCODED_SIZE: usize = packed_size::<T>(N);
fn write_le_bytes(&self, out: &mut [MaybeUninit<u8>]) {
write_packed_le_bytes(self, out);
}
fn read_le_bytes(bytes: &[u8]) -> Result<Self, PrimitiveError> {
let mut data: [MaybeUninit<T>; N] = std::array::from_fn(|_| MaybeUninit::uninit());
read_packed_le_bytes(bytes, &mut data)?;
Ok(data.map(|slot| unsafe { slot.assume_init() }))
}
}
unsafe impl<T: InPlaceCodec, N: ArraySize> InPlaceCodec for Array<T, N> {
const ENCODED_SIZE: usize = packed_size::<T>(N::USIZE);
fn write_le_bytes(&self, out: &mut [MaybeUninit<u8>]) {
write_packed_le_bytes(self, out);
}
fn read_le_bytes(bytes: &[u8]) -> Result<Self, PrimitiveError> {
let mut data = Array::<MaybeUninit<T>, N>::uninit();
read_packed_le_bytes(bytes, &mut data)?;
Ok(unsafe { data.assume_init() })
}
}
#[cfg(test)]
mod tests {
use hybrid_array::sizes::{U16, U20};
use super::*;
use crate::{
algebra::field::{binary::Gf2, mersenne::Mersenne107},
random::{test_rng, Random},
};
#[test]
fn array_pack_roundtrip_and_layout() {
let mut rng = test_rng();
assert_eq!(
<Array<Gf2, U20> as InPlaceCodec>::ENCODED_SIZE,
20usize.div_ceil(8)
);
let gf2 = Array::<Gf2, U20>::from_fn(|_| Gf2::random(&mut rng));
let gf2_bytes = gf2.to_inplace_bytes();
assert_eq!(gf2_bytes.len(), 3);
assert_eq!(
Array::<Gf2, U20>::from_inplace_bytes(&gf2_bytes).unwrap(),
gf2
);
let mers16 = Array::<Mersenne107, U16>::from_fn(|_| Mersenne107::random(&mut rng));
let mers16_bytes = mers16.to_inplace_bytes();
let reference: Vec<u8> = mers16.iter().flat_map(|e| e.to_inplace_bytes()).collect();
assert_eq!(
mers16_bytes, reference,
"packed Mersenne107 must be byte-identical to the per-element encoding"
);
assert_eq!(
Array::<Mersenne107, U16>::from_inplace_bytes(&mers16_bytes).unwrap(),
mers16
);
assert_eq!(
<Array<Mersenne107, U20> as InPlaceCodec>::ENCODED_SIZE,
20 * 14
);
let mers20 = Array::<Mersenne107, U20>::from_fn(|_| Mersenne107::random(&mut rng));
let mers20_bytes = mers20.to_inplace_bytes();
assert_eq!(mers20_bytes.len(), 20 * 14);
assert_eq!(
Array::<Mersenne107, U20>::from_inplace_bytes(&mers20_bytes).unwrap(),
mers20
);
}
#[test]
fn fixed_array_pack_roundtrip_and_layout() {
let mut rng = test_rng();
assert_eq!(
<[Gf2; 20] as InPlaceCodec>::ENCODED_SIZE,
20usize.div_ceil(8)
);
let gf2: [Gf2; 20] = std::array::from_fn(|_| Gf2::random(&mut rng));
let gf2_bytes = gf2.to_inplace_bytes();
assert_eq!(gf2_bytes.len(), 3);
assert_eq!(<[Gf2; 20]>::from_inplace_bytes(&gf2_bytes).unwrap(), gf2);
let mers16: [Mersenne107; 16] = std::array::from_fn(|_| Mersenne107::random(&mut rng));
let mers16_bytes = mers16.to_inplace_bytes();
let reference: Vec<u8> = mers16.iter().flat_map(|e| e.to_inplace_bytes()).collect();
assert_eq!(
mers16_bytes, reference,
"packed Mersenne107 must be byte-identical to the per-element encoding"
);
assert_eq!(
<[Mersenne107; 16]>::from_inplace_bytes(&mers16_bytes).unwrap(),
mers16
);
assert_eq!(<[Mersenne107; 20] as InPlaceCodec>::ENCODED_SIZE, 20 * 14);
let mers20: [Mersenne107; 20] = std::array::from_fn(|_| Mersenne107::random(&mut rng));
let mers20_bytes = mers20.to_inplace_bytes();
assert_eq!(mers20_bytes.len(), 20 * 14);
assert_eq!(
<[Mersenne107; 20]>::from_inplace_bytes(&mers20_bytes).unwrap(),
mers20
);
}
}