use std::{
array::from_fn,
io::{self, Cursor, Read, Result, Write},
};
pub use openvm_codec_derive::{Decode, Encode};
use p3_field::{BasedVectorSpace, PrimeField32};
use crate::StarkProtocolConfig;
pub(crate) const DECODE_PREALLOC_CAP: usize = 1024;
pub(crate) fn vec_with_capped_capacity<T>(len: usize) -> Vec<T> {
Vec::with_capacity(len.min(DECODE_PREALLOC_CAP))
}
pub trait Encode {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()>;
fn encode_to_vec(&self) -> Result<Vec<u8>> {
let mut buffer = Vec::new();
self.encode(&mut buffer)?;
Ok(buffer)
}
}
pub trait Decode: Sized {
fn decode<R: Read>(reader: &mut R) -> Result<Self>;
fn decode_from_bytes(bytes: &[u8]) -> Result<Self> {
let mut reader = Cursor::new(bytes);
let value = Self::decode(&mut reader)?;
if reader.position() != bytes.len() as u64 {
return Err(io::Error::other("trailing bytes after decoded value"));
}
Ok(value)
}
}
pub trait EncodableConfig: StarkProtocolConfig {
fn encode_base_field<W: Write>(val: &Self::F, writer: &mut W) -> Result<()>;
fn encode_extension_field<W: Write>(val: &Self::EF, writer: &mut W) -> Result<()>;
fn encode_digest<W: Write>(val: &Self::Digest, writer: &mut W) -> Result<()>;
fn encode_base_field_iter<'a, W: Write>(
iter: impl Iterator<Item = &'a Self::F>,
writer: &mut W,
) -> Result<()>
where
Self::F: 'a,
{
for val in iter {
Self::encode_base_field(val, writer)?;
}
Ok(())
}
fn encode_extension_field_iter<'a, W: Write>(
iter: impl Iterator<Item = &'a Self::EF>,
writer: &mut W,
) -> Result<()>
where
Self::EF: 'a,
{
for val in iter {
Self::encode_extension_field(val, writer)?;
}
Ok(())
}
fn encode_digest_iter<'a, W: Write>(
iter: impl Iterator<Item = &'a Self::Digest>,
writer: &mut W,
) -> Result<()>
where
Self::Digest: 'a,
{
for val in iter {
Self::encode_digest(val, writer)?;
}
Ok(())
}
fn encode_base_field_slice<W: Write>(vals: &[Self::F], writer: &mut W) -> Result<()> {
vals.len().encode(writer)?;
Self::encode_base_field_iter(vals.iter(), writer)
}
fn encode_extension_field_slice<W: Write>(vals: &[Self::EF], writer: &mut W) -> Result<()> {
vals.len().encode(writer)?;
Self::encode_extension_field_iter(vals.iter(), writer)
}
fn encode_digest_slice<W: Write>(vals: &[Self::Digest], writer: &mut W) -> Result<()> {
vals.len().encode(writer)?;
Self::encode_digest_iter(vals.iter(), writer)
}
}
pub trait DecodableConfig: StarkProtocolConfig {
fn decode_base_field<R: Read>(reader: &mut R) -> Result<Self::F>;
fn decode_extension_field<R: Read>(reader: &mut R) -> Result<Self::EF>;
fn decode_digest<R: Read>(reader: &mut R) -> Result<Self::Digest>;
fn decode_base_field_n<R: Read>(reader: &mut R, n: usize) -> Result<Vec<Self::F>> {
let mut vec = vec_with_capped_capacity(n);
for _ in 0..n {
vec.push(Self::decode_base_field(reader)?);
}
Ok(vec)
}
fn decode_extension_field_n<R: Read>(reader: &mut R, n: usize) -> Result<Vec<Self::EF>> {
let mut vec = vec_with_capped_capacity(n);
for _ in 0..n {
vec.push(Self::decode_extension_field(reader)?);
}
Ok(vec)
}
fn decode_digest_n<R: Read>(reader: &mut R, n: usize) -> Result<Vec<Self::Digest>> {
let mut vec = vec_with_capped_capacity(n);
for _ in 0..n {
vec.push(Self::decode_digest(reader)?);
}
Ok(vec)
}
fn decode_base_field_vec<R: Read>(reader: &mut R) -> Result<Vec<Self::F>> {
let len = usize::decode(reader)?;
Self::decode_base_field_n(reader, len)
}
fn decode_extension_field_vec<R: Read>(reader: &mut R) -> Result<Vec<Self::EF>> {
let len = usize::decode(reader)?;
Self::decode_extension_field_n(reader, len)
}
fn decode_digest_vec<R: Read>(reader: &mut R) -> Result<Vec<Self::Digest>> {
let len = usize::decode(reader)?;
Self::decode_digest_n(reader, len)
}
}
impl Encode for bool {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
writer.write_all(&[*self as u8])?;
Ok(())
}
}
impl Encode for u8 {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
writer.write_all(&[*self])
}
}
impl Encode for u32 {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
writer.write_all(&self.to_le_bytes())
}
}
impl Encode for usize {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
let x: u32 = (*self).try_into().map_err(io::Error::other)?;
x.encode(writer)
}
}
impl Encode for String {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
encode_slice(self.as_bytes(), writer)
}
}
pub fn encode_prime_field32<F: PrimeField32, W: Write>(val: &F, writer: &mut W) -> Result<()> {
writer.write_all(&val.as_canonical_u32().to_le_bytes())
}
pub fn decode_prime_field32<F: PrimeField32, R: Read>(reader: &mut R) -> Result<F> {
let mut bytes = [0u8; 4];
reader.read_exact(&mut bytes)?;
let value = u32::from_le_bytes(bytes);
if value < F::ORDER_U32 {
Ok(F::from_u32(value))
} else {
Err(io::Error::other(format!(
"Attempted read of {} into F >= F::ORDER_U32 {}",
value,
F::ORDER_U32
)))
}
}
pub fn encode_extension_field32<F: PrimeField32, EF: BasedVectorSpace<F>, W: Write>(
val: &EF,
writer: &mut W,
) -> Result<()> {
let base_slice: &[F] = val.as_basis_coefficients_slice();
for v in base_slice {
encode_prime_field32(v, writer)?;
}
Ok(())
}
pub fn decode_extension_field32<F: PrimeField32, EF: BasedVectorSpace<F>, R: Read>(
reader: &mut R,
) -> Result<EF> {
let d = <EF as BasedVectorSpace<F>>::DIMENSION;
let mut base_vec = Vec::with_capacity(d);
for _ in 0..d {
base_vec.push(decode_prime_field32(reader)?);
}
EF::from_basis_coefficients_slice(&base_vec)
.ok_or(io::Error::other("from_basis_coefficients_slice failed"))
}
pub fn encode_slice<T: Encode, W: Write>(slice: &[T], writer: &mut W) -> Result<()> {
slice.len().encode(writer)?;
for elt in slice {
elt.encode(writer)?;
}
Ok(())
}
pub fn encode_iter<'a, T: Encode + 'a, W: Write>(
iter: impl Iterator<Item = &'a T>,
writer: &mut W,
) -> Result<()> {
for elt in iter {
elt.encode(writer)?;
}
Ok(())
}
impl<T: Encode, const N: usize> Encode for [T; N] {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
for val in self {
val.encode(writer)?;
}
Ok(())
}
}
impl<T: Encode> Encode for Vec<T> {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
encode_slice(self, writer)
}
}
impl<S: Encode, T: Encode> Encode for (S, T) {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
self.0.encode(writer)?;
self.1.encode(writer)
}
}
impl<T: Encode> Encode for Option<T> {
fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
self.is_some().encode(writer)?;
if let Some(val) = self {
val.encode(writer)?;
}
Ok(())
}
}
impl Decode for bool {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let mut bytes = [0u8; 1];
reader.read_exact(&mut bytes)?;
Ok(bytes[0] != 0)
}
}
impl Decode for u8 {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let mut bytes = [0u8; 1];
reader.read_exact(&mut bytes)?;
Ok(bytes[0])
}
}
impl Decode for u32 {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let mut bytes = [0u8; 4];
reader.read_exact(&mut bytes)?;
Ok(u32::from_le_bytes(bytes))
}
}
impl Decode for usize {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let val = u32::decode(reader)?;
Ok(val as usize)
}
}
impl Decode for String {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let bytes = Vec::<u8>::decode(reader)?;
String::from_utf8(bytes).map_err(|err| io::Error::new(io::ErrorKind::InvalidData, err))
}
}
pub fn decode_into_vec<T: Decode, R: Read>(reader: &mut R, len: usize) -> Result<Vec<T>> {
let mut vec = vec_with_capped_capacity(len);
for _ in 0..len {
vec.push(T::decode(reader)?);
}
Ok(vec)
}
impl<T: Decode + Default, const N: usize> Decode for [T; N] {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let mut result = from_fn(|_| T::default());
for val in &mut result {
*val = T::decode(reader)?;
}
Ok(result)
}
}
impl<T: Decode> Decode for Vec<T> {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let len = usize::decode(reader)?;
let mut vec = vec_with_capped_capacity(len);
for _ in 0..len {
vec.push(T::decode(reader)?);
}
Ok(vec)
}
}
impl<S: Decode, T: Decode> Decode for (S, T) {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
Ok((S::decode(reader)?, T::decode(reader)?))
}
}
impl<T: Decode> Decode for Option<T> {
fn decode<R: Read>(reader: &mut R) -> Result<Self> {
let is_some = bool::decode(reader)?;
if is_some {
Ok(Some(T::decode(reader)?))
} else {
Ok(None)
}
}
}