openvm_stark_backend/
codec.rs1use 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
11pub(crate) const DECODE_PREALLOC_CAP: usize = 1024;
14
15pub(crate) fn vec_with_capped_capacity<T>(len: usize) -> Vec<T> {
18 Vec::with_capacity(len.min(DECODE_PREALLOC_CAP))
19}
20
21pub trait Encode {
25 fn encode<W: Write>(&self, writer: &mut W) -> Result<()>;
27
28 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
36pub trait Decode: Sized {
39 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
51pub 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 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 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 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 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 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 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
121pub 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 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 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 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 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 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 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
176impl 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
210pub 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
217pub 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
233pub 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
245pub 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
258pub 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
269pub 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
312impl 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
352pub 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}