Skip to main content

openvm_stark_backend/
codec.rs

1use std::{
2    array::from_fn,
3    io::{self, Cursor, Read, Result, Write},
4};
5
6pub use openvm_codec_derive::{Decode, Encode};
7use p3_field::{BasedVectorSpace, PrimeField32};
8
9use crate::StarkProtocolConfig;
10
11/// Upper bound on the capacity eagerly reserved while decoding untrusted
12/// length-prefixed collections.
13pub(crate) const DECODE_PREALLOC_CAP: usize = 1024;
14
15/// Allocates a `Vec` for decoding, capped to avoid attacker-controlled
16/// preallocation from length prefixes.
17pub(crate) fn vec_with_capped_capacity<T>(len: usize) -> Vec<T> {
18    Vec::with_capacity(len.min(DECODE_PREALLOC_CAP))
19}
20
21/// Hardware and language independent encoding.
22/// Uses the Writer pattern for more efficient encoding without intermediate buffers.
23// @dev Trait just for implementation sanity
24pub trait Encode {
25    /// Writes the encoded representation of `self` to the given writer.
26    fn encode<W: Write>(&self, writer: &mut W) -> Result<()>;
27
28    /// Convenience method to encode into a `Vec<u8>`
29    fn encode_to_vec(&self) -> Result<Vec<u8>> {
30        let mut buffer = Vec::new();
31        self.encode(&mut buffer)?;
32        Ok(buffer)
33    }
34}
35
36/// Hardware and language independent decoding.
37/// Uses the Reader pattern for efficient decoding.
38pub trait Decode: Sized {
39    /// Reads and decodes a value from the given reader.
40    fn decode<R: Read>(reader: &mut R) -> Result<Self>;
41    fn decode_from_bytes(bytes: &[u8]) -> Result<Self> {
42        let mut reader = Cursor::new(bytes);
43        let value = Self::decode(&mut reader)?;
44        if reader.position() != bytes.len() as u64 {
45            return Err(io::Error::other("trailing bytes after decoded value"));
46        }
47        Ok(value)
48    }
49}
50
51/// [StarkProtocolConfig] that has encodable associated types.
52/// This is a separate trait to avoid Rust's orphan rule.
53pub trait EncodableConfig: StarkProtocolConfig {
54    fn encode_base_field<W: Write>(val: &Self::F, writer: &mut W) -> Result<()>;
55
56    fn encode_extension_field<W: Write>(val: &Self::EF, writer: &mut W) -> Result<()>;
57
58    fn encode_digest<W: Write>(val: &Self::Digest, writer: &mut W) -> Result<()>;
59
60    /// Encode each element from an iterator (no length prefix).
61    fn encode_base_field_iter<'a, W: Write>(
62        iter: impl Iterator<Item = &'a Self::F>,
63        writer: &mut W,
64    ) -> Result<()>
65    where
66        Self::F: 'a,
67    {
68        for val in iter {
69            Self::encode_base_field(val, writer)?;
70        }
71        Ok(())
72    }
73
74    /// Encode each element from an iterator (no length prefix).
75    fn encode_extension_field_iter<'a, W: Write>(
76        iter: impl Iterator<Item = &'a Self::EF>,
77        writer: &mut W,
78    ) -> Result<()>
79    where
80        Self::EF: 'a,
81    {
82        for val in iter {
83            Self::encode_extension_field(val, writer)?;
84        }
85        Ok(())
86    }
87
88    /// Encode each element from an iterator (no length prefix).
89    fn encode_digest_iter<'a, W: Write>(
90        iter: impl Iterator<Item = &'a Self::Digest>,
91        writer: &mut W,
92    ) -> Result<()>
93    where
94        Self::Digest: 'a,
95    {
96        for val in iter {
97            Self::encode_digest(val, writer)?;
98        }
99        Ok(())
100    }
101
102    /// Encode length-prefixed slice of base field elements.
103    fn encode_base_field_slice<W: Write>(vals: &[Self::F], writer: &mut W) -> Result<()> {
104        vals.len().encode(writer)?;
105        Self::encode_base_field_iter(vals.iter(), writer)
106    }
107
108    /// Encode length-prefixed slice of extension field elements.
109    fn encode_extension_field_slice<W: Write>(vals: &[Self::EF], writer: &mut W) -> Result<()> {
110        vals.len().encode(writer)?;
111        Self::encode_extension_field_iter(vals.iter(), writer)
112    }
113
114    /// Encode length-prefixed slice of digests.
115    fn encode_digest_slice<W: Write>(vals: &[Self::Digest], writer: &mut W) -> Result<()> {
116        vals.len().encode(writer)?;
117        Self::encode_digest_iter(vals.iter(), writer)
118    }
119}
120
121/// [StarkProtocolConfig] that has decodable associated types.
122/// This is a separate trait to avoid Rust's orphan rule.
123pub trait DecodableConfig: StarkProtocolConfig {
124    fn decode_base_field<R: Read>(reader: &mut R) -> Result<Self::F>;
125
126    fn decode_extension_field<R: Read>(reader: &mut R) -> Result<Self::EF>;
127
128    fn decode_digest<R: Read>(reader: &mut R) -> Result<Self::Digest>;
129
130    /// Decode `n` base field elements (known length, no length prefix).
131    fn decode_base_field_n<R: Read>(reader: &mut R, n: usize) -> Result<Vec<Self::F>> {
132        let mut vec = vec_with_capped_capacity(n);
133        for _ in 0..n {
134            vec.push(Self::decode_base_field(reader)?);
135        }
136        Ok(vec)
137    }
138
139    /// Decode `n` extension field elements (known length, no length prefix).
140    fn decode_extension_field_n<R: Read>(reader: &mut R, n: usize) -> Result<Vec<Self::EF>> {
141        let mut vec = vec_with_capped_capacity(n);
142        for _ in 0..n {
143            vec.push(Self::decode_extension_field(reader)?);
144        }
145        Ok(vec)
146    }
147
148    /// Decode `n` digests (known length, no length prefix).
149    fn decode_digest_n<R: Read>(reader: &mut R, n: usize) -> Result<Vec<Self::Digest>> {
150        let mut vec = vec_with_capped_capacity(n);
151        for _ in 0..n {
152            vec.push(Self::decode_digest(reader)?);
153        }
154        Ok(vec)
155    }
156
157    /// Decode a length-prefixed vector of base field elements.
158    fn decode_base_field_vec<R: Read>(reader: &mut R) -> Result<Vec<Self::F>> {
159        let len = usize::decode(reader)?;
160        Self::decode_base_field_n(reader, len)
161    }
162
163    /// Decode a length-prefixed vector of extension field elements.
164    fn decode_extension_field_vec<R: Read>(reader: &mut R) -> Result<Vec<Self::EF>> {
165        let len = usize::decode(reader)?;
166        Self::decode_extension_field_n(reader, len)
167    }
168
169    /// Decode a length-prefixed vector of digests.
170    fn decode_digest_vec<R: Read>(reader: &mut R) -> Result<Vec<Self::Digest>> {
171        let len = usize::decode(reader)?;
172        Self::decode_digest_n(reader, len)
173    }
174}
175
176// ==================== Encode implementations for basic types ====================
177
178impl Encode for bool {
179    fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
180        writer.write_all(&[*self as u8])?;
181        Ok(())
182    }
183}
184
185impl Encode for u8 {
186    fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
187        writer.write_all(&[*self])
188    }
189}
190
191impl Encode for u32 {
192    fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
193        writer.write_all(&self.to_le_bytes())
194    }
195}
196
197impl Encode for usize {
198    fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
199        let x: u32 = (*self).try_into().map_err(io::Error::other)?;
200        x.encode(writer)
201    }
202}
203
204impl Encode for String {
205    fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
206        encode_slice(self.as_bytes(), writer)
207    }
208}
209
210// ==================== Generic field codec helpers ====================
211
212/// Encode a `PrimeField32` element as 4 little-endian bytes of its canonical u32 value.
213pub fn encode_prime_field32<F: PrimeField32, W: Write>(val: &F, writer: &mut W) -> Result<()> {
214    writer.write_all(&val.as_canonical_u32().to_le_bytes())
215}
216
217/// Decode a `PrimeField32` element from 4 little-endian bytes.
218pub fn decode_prime_field32<F: PrimeField32, R: Read>(reader: &mut R) -> Result<F> {
219    let mut bytes = [0u8; 4];
220    reader.read_exact(&mut bytes)?;
221    let value = u32::from_le_bytes(bytes);
222    if value < F::ORDER_U32 {
223        Ok(F::from_u32(value))
224    } else {
225        Err(io::Error::other(format!(
226            "Attempted read of {} into F >= F::ORDER_U32 {}",
227            value,
228            F::ORDER_U32
229        )))
230    }
231}
232
233/// Encode an extension field element by encoding each basis coefficient.
234pub fn encode_extension_field32<F: PrimeField32, EF: BasedVectorSpace<F>, W: Write>(
235    val: &EF,
236    writer: &mut W,
237) -> Result<()> {
238    let base_slice: &[F] = val.as_basis_coefficients_slice();
239    for v in base_slice {
240        encode_prime_field32(v, writer)?;
241    }
242    Ok(())
243}
244
245/// Decode an extension field element by decoding each basis coefficient.
246pub fn decode_extension_field32<F: PrimeField32, EF: BasedVectorSpace<F>, R: Read>(
247    reader: &mut R,
248) -> Result<EF> {
249    let d = <EF as BasedVectorSpace<F>>::DIMENSION;
250    let mut base_vec = Vec::with_capacity(d);
251    for _ in 0..d {
252        base_vec.push(decode_prime_field32(reader)?);
253    }
254    EF::from_basis_coefficients_slice(&base_vec)
255        .ok_or(io::Error::other("from_basis_coefficients_slice failed"))
256}
257
258// ==================== Encode helpers ====================
259
260/// Encodes length of slice and then each element
261pub fn encode_slice<T: Encode, W: Write>(slice: &[T], writer: &mut W) -> Result<()> {
262    slice.len().encode(writer)?;
263    for elt in slice {
264        elt.encode(writer)?;
265    }
266    Ok(())
267}
268
269/// Encodes each element (no length)
270pub fn encode_iter<'a, T: Encode + 'a, W: Write>(
271    iter: impl Iterator<Item = &'a T>,
272    writer: &mut W,
273) -> Result<()> {
274    for elt in iter {
275        elt.encode(writer)?;
276    }
277    Ok(())
278}
279
280impl<T: Encode, const N: usize> Encode for [T; N] {
281    fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
282        for val in self {
283            val.encode(writer)?;
284        }
285        Ok(())
286    }
287}
288
289impl<T: Encode> Encode for Vec<T> {
290    fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
291        encode_slice(self, writer)
292    }
293}
294
295impl<S: Encode, T: Encode> Encode for (S, T) {
296    fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
297        self.0.encode(writer)?;
298        self.1.encode(writer)
299    }
300}
301
302impl<T: Encode> Encode for Option<T> {
303    fn encode<W: Write>(&self, writer: &mut W) -> Result<()> {
304        self.is_some().encode(writer)?;
305        if let Some(val) = self {
306            val.encode(writer)?;
307        }
308        Ok(())
309    }
310}
311
312// ==================== Decode implementations for basic types ====================
313
314impl Decode for bool {
315    fn decode<R: Read>(reader: &mut R) -> Result<Self> {
316        let mut bytes = [0u8; 1];
317        reader.read_exact(&mut bytes)?;
318        Ok(bytes[0] != 0)
319    }
320}
321
322impl Decode for u8 {
323    fn decode<R: Read>(reader: &mut R) -> Result<Self> {
324        let mut bytes = [0u8; 1];
325        reader.read_exact(&mut bytes)?;
326        Ok(bytes[0])
327    }
328}
329
330impl Decode for u32 {
331    fn decode<R: Read>(reader: &mut R) -> Result<Self> {
332        let mut bytes = [0u8; 4];
333        reader.read_exact(&mut bytes)?;
334        Ok(u32::from_le_bytes(bytes))
335    }
336}
337
338impl Decode for usize {
339    fn decode<R: Read>(reader: &mut R) -> Result<Self> {
340        let val = u32::decode(reader)?;
341        Ok(val as usize)
342    }
343}
344
345impl Decode for String {
346    fn decode<R: Read>(reader: &mut R) -> Result<Self> {
347        let bytes = Vec::<u8>::decode(reader)?;
348        String::from_utf8(bytes).map_err(|err| io::Error::new(io::ErrorKind::InvalidData, err))
349    }
350}
351
352// ==================== Decode helpers ====================
353
354/// Decodes into a vector given preset length
355pub fn decode_into_vec<T: Decode, R: Read>(reader: &mut R, len: usize) -> Result<Vec<T>> {
356    let mut vec = vec_with_capped_capacity(len);
357    for _ in 0..len {
358        vec.push(T::decode(reader)?);
359    }
360    Ok(vec)
361}
362
363impl<T: Decode + Default, const N: usize> Decode for [T; N] {
364    fn decode<R: Read>(reader: &mut R) -> Result<Self> {
365        let mut result = from_fn(|_| T::default());
366        for val in &mut result {
367            *val = T::decode(reader)?;
368        }
369        Ok(result)
370    }
371}
372
373impl<T: Decode> Decode for Vec<T> {
374    fn decode<R: Read>(reader: &mut R) -> Result<Self> {
375        let len = usize::decode(reader)?;
376        let mut vec = vec_with_capped_capacity(len);
377        for _ in 0..len {
378            vec.push(T::decode(reader)?);
379        }
380        Ok(vec)
381    }
382}
383
384impl<S: Decode, T: Decode> Decode for (S, T) {
385    fn decode<R: Read>(reader: &mut R) -> Result<Self> {
386        Ok((S::decode(reader)?, T::decode(reader)?))
387    }
388}
389
390impl<T: Decode> Decode for Option<T> {
391    fn decode<R: Read>(reader: &mut R) -> Result<Self> {
392        let is_some = bool::decode(reader)?;
393        if is_some {
394            Ok(Some(T::decode(reader)?))
395        } else {
396            Ok(None)
397        }
398    }
399}