arcium-primitives 0.8.1

Arcium primitives
Documentation
//! [`InPlaceCodec`] for fixed-size container types.

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;

// SAFETY: `Arc<T>` encodes exactly as its pointee `T` (same `ENCODED_SIZE`, same bytes), which is
// architecture-independent, initializes every byte, and round-trips unbiasedly; decoding allocates
// a fresh `Arc`.
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)
    }
}

// SAFETY: a 2-tuple encodes as its two elements' `InPlaceCodec` encodings, back to back. Each half
// occupies its element's `ENCODED_SIZE`; `write_le_bytes` initializes both (hence every byte) and
// the round-trip is unbiased since each element's is.
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)?))
    }
}

// Fast (de)serialization for a native fixed-size array of `InPlaceCodec` elements. The `PACK`-aware
// layout (`ceil(N / PACK)` full packs, trailing partial group padded with dummy elements) is
// shared with `HeapArray<T, M>`/`Array<T, N>`'s impls via `packed_size`/`write_packed_le_bytes`/
// `read_packed_le_bytes`.
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)?;
        // SAFETY: every slot of `data` was written above (`read_packed_le_bytes` only returns `Ok`
        // once it has filled every slot).
        Ok(data.map(|slot| unsafe { slot.assume_init() }))
    }
}

// Fast (de)serialization for `hybrid_array::Array`s of `InPlaceCodec` elements. Same `PACK`-aware
// layout as `[T; N]` above, shared via
// `packed_size`/`write_packed_le_bytes`/`read_packed_le_bytes`.
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)?;
        // SAFETY: every slot of `data` was written above (`read_packed_le_bytes` only returns `Ok`
        // once it has filled every slot).
        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},
    };

    /// Mirrors `test_heap_array_pack_roundtrip_and_layout` in `types::heap_array::array`: `Gf2`
    /// exercises sub-byte packing (8 -> 1); `Mersenne107` exercises the vectorized-but-unpacked
    /// `PACK` path (byte-identical for full packs; an unpadded per-element tail for a remainder,
    /// since padding would only add dead bytes here — see `tail_is_unpacked`).
    #[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
        );

        // No compression from packing (`PACK_BYTES == PACK * ENCODED_SIZE`), so an unpadded tail
        // keeps the total at exactly `20 * ENCODED_SIZE`, unlike `Gf2` above.
        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
        );
    }

    /// Same as `array_pack_roundtrip_and_layout`, for the native `[T; N]` impl.
    #[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
        );
    }
}