Skip to main content

miden_serde_utils/
lib.rs

1// Copyright (c) Facebook, Inc. and its affiliates.
2//
3// This source code is licensed under the MIT license found in the
4// LICENSE file in the root directory of this source tree.
5
6#![cfg_attr(not(feature = "std"), no_std)]
7
8extern crate alloc;
9
10use alloc::{
11    collections::{BTreeMap, BTreeSet, VecDeque},
12    format,
13    string::String,
14    sync::Arc,
15    vec::Vec,
16};
17use core::mem::size_of;
18
19/// Deserializes a wincode schema and rejects trailing bytes.
20///
21/// Unlike `wincode::config::deserialize_exact`, this accepts schema adapters whose decoded type
22/// differs from the schema type.
23pub fn deserialize_schema_exact<'de, T, C>(
24    mut bytes: &'de [u8],
25    _config: C,
26) -> wincode::ReadResult<T::Dst>
27where
28    T: wincode::SchemaRead<'de, C>,
29    C: wincode::config::Config,
30{
31    use wincode::io::Reader as _;
32
33    let value = T::get(bytes.by_ref())?;
34    if bytes.is_empty() {
35        Ok(value)
36    } else {
37        Err(wincode::error::trailing_bytes())
38    }
39}
40
41#[cfg(test)]
42mod wincode_tests {
43    use serde_wincode::SerdeCompat;
44
45    use super::deserialize_schema_exact;
46
47    #[test]
48    fn exact_deserialization_rejects_trailing_bytes() {
49        let config = wincode::config::Configuration::default();
50        let mut bytes =
51            <SerdeCompat<u8> as wincode::config::Serialize<_>>::serialize(&7, config).unwrap();
52
53        assert_eq!(deserialize_schema_exact::<SerdeCompat<u8>, _>(&bytes, config).unwrap(), 7);
54
55        bytes.push(0);
56        assert!(matches!(
57            deserialize_schema_exact::<SerdeCompat<u8>, _>(&bytes, config),
58            Err(wincode::error::ReadError::TrailingBytes)
59        ));
60    }
61}
62
63// ERROR
64// ================================================================================================
65
66/// Defines errors which can occur during deserialization.
67#[derive(Clone, Debug, PartialEq, Eq)]
68pub enum DeserializationError {
69    /// Indicates that the deserialization failed because of insufficient data.
70    UnexpectedEOF,
71    /// Indicates that the deserialization failed because the value was not valid.
72    InvalidValue(String),
73    /// Indicates that deserialization failed for an unknown reason.
74    UnknownError(String),
75}
76
77impl core::error::Error for DeserializationError {}
78
79impl core::fmt::Display for DeserializationError {
80    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
81        match self {
82            Self::UnexpectedEOF => write!(f, "unexpected end of file"),
83            Self::InvalidValue(msg) => write!(f, "invalid value: {msg}"),
84            Self::UnknownError(msg) => write!(f, "unknown error: {msg}"),
85        }
86    }
87}
88
89mod byte_reader;
90#[cfg(feature = "std")]
91pub use byte_reader::ReadAdapter;
92pub use byte_reader::{BudgetedReader, ByteReader, ReadManyIter, SliceReader};
93
94mod byte_writer;
95pub use byte_writer::ByteWriter;
96
97// BOUNDED LENGTH HELPERS
98// ================================================================================================
99
100/// Reads and validates a serialized length before it is used for allocation.
101///
102/// `label` names the collection being read and is only used to build the error message.
103/// `min_element_size` is the minimum number of bytes one element occupies once serialized.
104///
105/// # Errors
106/// Returns an error if the length cannot be read, or if it fails the checks described in
107/// [`validate_bounded_len`].
108pub fn read_bounded_len<R: ByteReader>(
109    source: &mut R,
110    label: &str,
111    min_element_size: usize,
112) -> Result<usize, DeserializationError> {
113    let len = source.read_usize()?;
114    validate_bounded_len(source, label, len, min_element_size)?;
115    Ok(len)
116}
117
118/// Validates that a serialized length fits both the reader budget and remaining input.
119///
120/// This guards against malicious length prefixes that would otherwise cause a compact payload
121/// to be amplified into a large allocation.
122///
123/// # Errors
124/// Returns [`DeserializationError::InvalidValue`] if:
125/// * `len` exceeds the number of elements the reader's remaining budget allows.
126/// * `len * min_element_size` overflows.
127/// * The source does not hold `len * min_element_size` more bytes.
128pub fn validate_bounded_len<R: ByteReader>(
129    source: &R,
130    label: &str,
131    len: usize,
132    min_element_size: usize,
133) -> Result<(), DeserializationError> {
134    let max_len = source.max_alloc(min_element_size);
135    if len > max_len {
136        return Err(DeserializationError::InvalidValue(format!(
137            "{label} count {len} exceeds budget {max_len}"
138        )));
139    }
140
141    let min_bytes = len.checked_mul(min_element_size).ok_or_else(|| {
142        DeserializationError::InvalidValue(format!(
143            "{label} count {len} overflows minimum serialized size {min_element_size}"
144        ))
145    })?;
146    source.check_eor(min_bytes).map_err(|err| match err {
147        DeserializationError::UnexpectedEOF => DeserializationError::InvalidValue(format!(
148            "{label} count {len} exceeds remaining input"
149        )),
150        err => err,
151    })
152}
153
154// SERIALIZABLE TRAIT
155// ================================================================================================
156
157/// Defines how to serialize `Self` into bytes.
158pub trait Serializable {
159    // REQUIRED METHODS
160    // --------------------------------------------------------------------------------------------
161    /// Serializes `self` into bytes and writes these bytes into the `target`.
162    fn write_into<W: ByteWriter>(&self, target: &mut W);
163
164    // PROVIDED METHODS
165    // --------------------------------------------------------------------------------------------
166
167    /// Serializes `self` into a vector of bytes.
168    fn to_bytes(&self) -> Vec<u8> {
169        let mut result = Vec::with_capacity(self.get_size_hint());
170        self.write_into(&mut result);
171        result
172    }
173
174    /// Returns an estimate of how many bytes are needed to represent self.
175    ///
176    /// The default implementation returns zero.
177    fn get_size_hint(&self) -> usize {
178        0
179    }
180}
181
182impl<T: Serializable> Serializable for &T {
183    fn write_into<W: ByteWriter>(&self, target: &mut W) {
184        (*self).write_into(target)
185    }
186
187    fn get_size_hint(&self) -> usize {
188        (*self).get_size_hint()
189    }
190}
191
192impl Serializable for () {
193    fn write_into<W: ByteWriter>(&self, _target: &mut W) {}
194
195    fn get_size_hint(&self) -> usize {
196        0
197    }
198}
199
200impl<T1> Serializable for (T1,)
201where
202    T1: Serializable,
203{
204    fn write_into<W: ByteWriter>(&self, target: &mut W) {
205        self.0.write_into(target);
206    }
207
208    fn get_size_hint(&self) -> usize {
209        self.0.get_size_hint()
210    }
211}
212
213impl<T1, T2> Serializable for (T1, T2)
214where
215    T1: Serializable,
216    T2: Serializable,
217{
218    fn write_into<W: ByteWriter>(&self, target: &mut W) {
219        self.0.write_into(target);
220        self.1.write_into(target);
221    }
222
223    fn get_size_hint(&self) -> usize {
224        self.0.get_size_hint() + self.1.get_size_hint()
225    }
226}
227
228impl<T1, T2, T3> Serializable for (T1, T2, T3)
229where
230    T1: Serializable,
231    T2: Serializable,
232    T3: Serializable,
233{
234    fn write_into<W: ByteWriter>(&self, target: &mut W) {
235        self.0.write_into(target);
236        self.1.write_into(target);
237        self.2.write_into(target);
238    }
239
240    fn get_size_hint(&self) -> usize {
241        self.0.get_size_hint() + self.1.get_size_hint() + self.2.get_size_hint()
242    }
243}
244
245impl<T1, T2, T3, T4> Serializable for (T1, T2, T3, T4)
246where
247    T1: Serializable,
248    T2: Serializable,
249    T3: Serializable,
250    T4: Serializable,
251{
252    fn write_into<W: ByteWriter>(&self, target: &mut W) {
253        self.0.write_into(target);
254        self.1.write_into(target);
255        self.2.write_into(target);
256        self.3.write_into(target);
257    }
258
259    fn get_size_hint(&self) -> usize {
260        self.0.get_size_hint()
261            + self.1.get_size_hint()
262            + self.2.get_size_hint()
263            + self.3.get_size_hint()
264    }
265}
266
267impl<T1, T2, T3, T4, T5> Serializable for (T1, T2, T3, T4, T5)
268where
269    T1: Serializable,
270    T2: Serializable,
271    T3: Serializable,
272    T4: Serializable,
273    T5: Serializable,
274{
275    fn write_into<W: ByteWriter>(&self, target: &mut W) {
276        self.0.write_into(target);
277        self.1.write_into(target);
278        self.2.write_into(target);
279        self.3.write_into(target);
280        self.4.write_into(target);
281    }
282
283    fn get_size_hint(&self) -> usize {
284        self.0.get_size_hint()
285            + self.1.get_size_hint()
286            + self.2.get_size_hint()
287            + self.3.get_size_hint()
288            + self.4.get_size_hint()
289    }
290}
291
292impl<T1, T2, T3, T4, T5, T6> Serializable for (T1, T2, T3, T4, T5, T6)
293where
294    T1: Serializable,
295    T2: Serializable,
296    T3: Serializable,
297    T4: Serializable,
298    T5: Serializable,
299    T6: Serializable,
300{
301    fn write_into<W: ByteWriter>(&self, target: &mut W) {
302        self.0.write_into(target);
303        self.1.write_into(target);
304        self.2.write_into(target);
305        self.3.write_into(target);
306        self.4.write_into(target);
307        self.5.write_into(target);
308    }
309
310    fn get_size_hint(&self) -> usize {
311        self.0.get_size_hint()
312            + self.1.get_size_hint()
313            + self.2.get_size_hint()
314            + self.3.get_size_hint()
315            + self.4.get_size_hint()
316            + self.5.get_size_hint()
317    }
318}
319
320impl Serializable for u8 {
321    fn write_into<W: ByteWriter>(&self, target: &mut W) {
322        target.write_u8(*self);
323    }
324
325    fn get_size_hint(&self) -> usize {
326        size_of::<u8>()
327    }
328}
329
330impl Serializable for u16 {
331    fn write_into<W: ByteWriter>(&self, target: &mut W) {
332        target.write_u16(*self);
333    }
334
335    fn get_size_hint(&self) -> usize {
336        size_of::<u16>()
337    }
338}
339
340impl Serializable for u32 {
341    fn write_into<W: ByteWriter>(&self, target: &mut W) {
342        target.write_u32(*self);
343    }
344
345    fn get_size_hint(&self) -> usize {
346        size_of::<u32>()
347    }
348}
349
350impl Serializable for u64 {
351    fn write_into<W: ByteWriter>(&self, target: &mut W) {
352        target.write_u64(*self);
353    }
354
355    fn get_size_hint(&self) -> usize {
356        size_of::<u64>()
357    }
358}
359
360impl Serializable for u128 {
361    fn write_into<W: ByteWriter>(&self, target: &mut W) {
362        target.write_u128(*self);
363    }
364
365    fn get_size_hint(&self) -> usize {
366        size_of::<u128>()
367    }
368}
369
370impl Serializable for usize {
371    fn write_into<W: ByteWriter>(&self, target: &mut W) {
372        target.write_usize(*self)
373    }
374
375    fn get_size_hint(&self) -> usize {
376        byte_writer::usize_encoded_len(*self as u64)
377    }
378}
379
380impl<T: Serializable> Serializable for Option<T> {
381    fn write_into<W: ByteWriter>(&self, target: &mut W) {
382        match self {
383            Some(v) => {
384                target.write_bool(true);
385                v.write_into(target);
386            },
387            None => target.write_bool(false),
388        }
389    }
390
391    fn get_size_hint(&self) -> usize {
392        size_of::<bool>() + self.as_ref().map(Serializable::get_size_hint).unwrap_or(0)
393    }
394}
395
396impl<T: Serializable, const C: usize> Serializable for [T; C] {
397    fn write_into<W: ByteWriter>(&self, target: &mut W) {
398        target.write_many(self)
399    }
400
401    fn get_size_hint(&self) -> usize {
402        let mut size = 0;
403        for item in self {
404            size += item.get_size_hint();
405        }
406        size
407    }
408}
409
410impl<T: Serializable> Serializable for [T] {
411    fn write_into<W: ByteWriter>(&self, target: &mut W) {
412        target.write_usize(self.len());
413        for element in self.iter() {
414            element.write_into(target);
415        }
416    }
417
418    fn get_size_hint(&self) -> usize {
419        let mut size = self.len().get_size_hint();
420        for element in self {
421            size += element.get_size_hint();
422        }
423        size
424    }
425}
426
427impl<T: Serializable> Serializable for Vec<T> {
428    fn write_into<W: ByteWriter>(&self, target: &mut W) {
429        target.write_usize(self.len());
430        target.write_many(self);
431    }
432
433    fn get_size_hint(&self) -> usize {
434        let mut size = self.len().get_size_hint();
435        for item in self {
436            size += item.get_size_hint();
437        }
438        size
439    }
440}
441
442impl<K: Serializable, V: Serializable> Serializable for BTreeMap<K, V> {
443    fn write_into<W: ByteWriter>(&self, target: &mut W) {
444        target.write_usize(self.len());
445        target.write_many(self);
446    }
447
448    fn get_size_hint(&self) -> usize {
449        let mut size = self.len().get_size_hint();
450        for item in self {
451            size += item.get_size_hint();
452        }
453        size
454    }
455}
456
457impl<T: Serializable> Serializable for BTreeSet<T> {
458    fn write_into<W: ByteWriter>(&self, target: &mut W) {
459        target.write_usize(self.len());
460        target.write_many(self);
461    }
462
463    fn get_size_hint(&self) -> usize {
464        let mut size = self.len().get_size_hint();
465        for item in self {
466            size += item.get_size_hint();
467        }
468        size
469    }
470}
471
472impl Serializable for str {
473    fn write_into<W: ByteWriter>(&self, target: &mut W) {
474        target.write_usize(self.len());
475        target.write_many(self.as_bytes());
476    }
477
478    fn get_size_hint(&self) -> usize {
479        self.len().get_size_hint() + self.len()
480    }
481}
482
483impl Serializable for String {
484    fn write_into<W: ByteWriter>(&self, target: &mut W) {
485        self.as_str().write_into(target);
486    }
487
488    fn get_size_hint(&self) -> usize {
489        self.as_str().get_size_hint()
490    }
491}
492
493impl Serializable for Arc<str> {
494    fn write_into<W: ByteWriter>(&self, target: &mut W) {
495        self.as_ref().write_into(target);
496    }
497
498    fn get_size_hint(&self) -> usize {
499        self.as_ref().get_size_hint()
500    }
501}
502
503// DESERIALIZABLE
504// ================================================================================================
505
506/// Defines how to deserialize `Self` from bytes.
507pub trait Deserializable: Sized {
508    // REQUIRED METHODS
509    // --------------------------------------------------------------------------------------------
510
511    /// Reads a sequence of bytes from the provided `source`, attempts to deserialize these bytes
512    /// into `Self`, and returns the result.
513    ///
514    /// # Errors
515    /// Returns an error if:
516    /// * The `source` does not contain enough bytes to deserialize `Self`.
517    /// * Bytes read from the `source` do not represent a valid value for `Self`.
518    #[track_caller]
519    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError>;
520
521    /// Returns the minimum serialized size for one instance of this type.
522    ///
523    /// This is used by [`ByteReader::max_alloc`] to estimate how many elements can be
524    /// deserialized from the remaining budget, preventing denial-of-service attacks from
525    /// malicious length prefixes.
526    ///
527    /// The default implementation returns `size_of::<Self>()`, which is conservative: it may
528    /// reject valid input for types where the serialized size is smaller than the in-memory
529    /// size (e.g., structs with computed/cached fields that aren't serialized).
530    ///
531    /// Override this method for types where the serialized representation is smaller than
532    /// the in-memory representation to allow more elements to be deserialized.
533    fn min_serialized_size() -> usize {
534        size_of::<Self>()
535    }
536
537    // PROVIDED METHODS
538    // --------------------------------------------------------------------------------------------
539
540    /// Attempts to deserialize the provided `bytes` into `Self` and returns the result.
541    ///
542    /// # Errors
543    /// Returns an error if:
544    /// * The `bytes` do not contain enough information to deserialize `Self`.
545    /// * The `bytes` do not represent a valid value for `Self`.
546    ///
547    /// Note: if `bytes` contains more data than needed to deserialize `self`, no error is
548    /// returned.
549    ///
550    /// # Security
551    /// This method is for trusted input. It does not bound allocations or reject trailing bytes.
552    /// Use [`Deserializable::read_from_bytes_with_budget`] for attacker-controlled bytes.
553    #[track_caller]
554    fn read_from_bytes(bytes: &[u8]) -> Result<Self, DeserializationError> {
555        Self::read_from(&mut SliceReader::new(bytes))
556    }
557
558    /// Deserializes `Self` from bytes with a byte budget limit.
559    ///
560    /// This is the recommended method for deserializing untrusted input. The budget limits
561    /// how many bytes can be consumed during deserialization, preventing denial-of-service
562    /// attacks that exploit length fields to cause huge allocations.
563    ///
564    /// # Errors
565    /// Returns an error if:
566    /// * The budget is exhausted before deserialization completes.
567    /// * The `bytes` do not contain enough information to deserialize `Self`.
568    /// * The `bytes` do not represent a valid value for `Self`.
569    #[track_caller]
570    fn read_from_bytes_with_budget(
571        bytes: &[u8],
572        budget: usize,
573    ) -> Result<Self, DeserializationError> {
574        Self::read_from(&mut BudgetedReader::new(SliceReader::new(bytes), budget))
575    }
576}
577
578impl Deserializable for () {
579    fn read_from<R: ByteReader>(_source: &mut R) -> Result<Self, DeserializationError> {
580        Ok(())
581    }
582}
583
584impl<T1> Deserializable for (T1,)
585where
586    T1: Deserializable,
587{
588    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
589        let v1 = T1::read_from(source)?;
590        Ok((v1,))
591    }
592
593    fn min_serialized_size() -> usize {
594        T1::min_serialized_size()
595    }
596}
597
598impl<T1, T2> Deserializable for (T1, T2)
599where
600    T1: Deserializable,
601    T2: Deserializable,
602{
603    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
604        let v1 = T1::read_from(source)?;
605        let v2 = T2::read_from(source)?;
606        Ok((v1, v2))
607    }
608
609    fn min_serialized_size() -> usize {
610        T1::min_serialized_size().saturating_add(T2::min_serialized_size())
611    }
612}
613
614impl<T1, T2, T3> Deserializable for (T1, T2, T3)
615where
616    T1: Deserializable,
617    T2: Deserializable,
618    T3: Deserializable,
619{
620    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
621        let v1 = T1::read_from(source)?;
622        let v2 = T2::read_from(source)?;
623        let v3 = T3::read_from(source)?;
624        Ok((v1, v2, v3))
625    }
626
627    fn min_serialized_size() -> usize {
628        T1::min_serialized_size()
629            .saturating_add(T2::min_serialized_size())
630            .saturating_add(T3::min_serialized_size())
631    }
632}
633
634impl<T1, T2, T3, T4> Deserializable for (T1, T2, T3, T4)
635where
636    T1: Deserializable,
637    T2: Deserializable,
638    T3: Deserializable,
639    T4: Deserializable,
640{
641    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
642        let v1 = T1::read_from(source)?;
643        let v2 = T2::read_from(source)?;
644        let v3 = T3::read_from(source)?;
645        let v4 = T4::read_from(source)?;
646        Ok((v1, v2, v3, v4))
647    }
648
649    fn min_serialized_size() -> usize {
650        T1::min_serialized_size()
651            .saturating_add(T2::min_serialized_size())
652            .saturating_add(T3::min_serialized_size())
653            .saturating_add(T4::min_serialized_size())
654    }
655}
656
657impl<T1, T2, T3, T4, T5> Deserializable for (T1, T2, T3, T4, T5)
658where
659    T1: Deserializable,
660    T2: Deserializable,
661    T3: Deserializable,
662    T4: Deserializable,
663    T5: Deserializable,
664{
665    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
666        let v1 = T1::read_from(source)?;
667        let v2 = T2::read_from(source)?;
668        let v3 = T3::read_from(source)?;
669        let v4 = T4::read_from(source)?;
670        let v5 = T5::read_from(source)?;
671        Ok((v1, v2, v3, v4, v5))
672    }
673
674    fn min_serialized_size() -> usize {
675        T1::min_serialized_size()
676            .saturating_add(T2::min_serialized_size())
677            .saturating_add(T3::min_serialized_size())
678            .saturating_add(T4::min_serialized_size())
679            .saturating_add(T5::min_serialized_size())
680    }
681}
682
683impl<T1, T2, T3, T4, T5, T6> Deserializable for (T1, T2, T3, T4, T5, T6)
684where
685    T1: Deserializable,
686    T2: Deserializable,
687    T3: Deserializable,
688    T4: Deserializable,
689    T5: Deserializable,
690    T6: Deserializable,
691{
692    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
693        let v1 = T1::read_from(source)?;
694        let v2 = T2::read_from(source)?;
695        let v3 = T3::read_from(source)?;
696        let v4 = T4::read_from(source)?;
697        let v5 = T5::read_from(source)?;
698        let v6 = T6::read_from(source)?;
699        Ok((v1, v2, v3, v4, v5, v6))
700    }
701
702    fn min_serialized_size() -> usize {
703        T1::min_serialized_size()
704            .saturating_add(T2::min_serialized_size())
705            .saturating_add(T3::min_serialized_size())
706            .saturating_add(T4::min_serialized_size())
707            .saturating_add(T5::min_serialized_size())
708            .saturating_add(T6::min_serialized_size())
709    }
710}
711
712impl Deserializable for u8 {
713    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
714        source.read_u8()
715    }
716}
717
718impl Deserializable for u16 {
719    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
720        source.read_u16()
721    }
722}
723
724impl Deserializable for u32 {
725    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
726        source.read_u32()
727    }
728}
729
730impl Deserializable for u64 {
731    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
732        source.read_u64()
733    }
734}
735
736impl Deserializable for u128 {
737    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
738        source.read_u128()
739    }
740}
741
742impl Deserializable for usize {
743    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
744        source.read_usize()
745    }
746
747    fn min_serialized_size() -> usize {
748        1 // vint64 encoding: minimum 1 byte for values 0-127
749    }
750}
751
752impl<T: Deserializable> Deserializable for Option<T> {
753    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
754        if source.read_bool()? {
755            Ok(Some(T::read_from(source)?))
756        } else {
757            Ok(None)
758        }
759    }
760
761    /// Returns 1 (just the bool discriminator).
762    ///
763    /// The `Some` variant would be `1 + T::min_serialized_size()`, but we use the minimum
764    /// to allow more elements through the early check.
765    fn min_serialized_size() -> usize {
766        1
767    }
768}
769
770impl<T: Deserializable, const C: usize> Deserializable for [T; C] {
771    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
772        let data: Vec<T> = source.read_many_iter(C)?.collect::<Result<_, _>>()?;
773
774        // The iterator yields exactly C elements (or fails early), so this always succeeds
775        Ok(data.try_into().unwrap_or_else(|v: Vec<T>| {
776            panic!("Expected a Vec of length {} but it was {}", C, v.len())
777        }))
778    }
779
780    fn min_serialized_size() -> usize {
781        C.saturating_mul(T::min_serialized_size())
782    }
783}
784
785impl<T: Deserializable> Deserializable for Vec<T> {
786    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
787        let len = source.read_usize()?;
788        source.read_many_iter(len)?.collect()
789    }
790
791    /// Returns 1 (the minimum vint length prefix size).
792    ///
793    /// The actual serialized size depends on the number of elements, which we don't know
794    /// at the point this is called. Using the minimum allows more elements through the
795    /// early check; budget enforcement during actual reads provides the real protection.
796    fn min_serialized_size() -> usize {
797        1
798    }
799}
800
801impl<T: Serializable> Serializable for VecDeque<T> {
802    fn write_into<W: ByteWriter>(&self, target: &mut W) {
803        target.write_usize(self.len());
804        for item in self {
805            item.write_into(target);
806        }
807    }
808
809    fn get_size_hint(&self) -> usize {
810        let mut size = self.len().get_size_hint();
811        for item in self {
812            size += item.get_size_hint();
813        }
814        size
815    }
816}
817
818impl<T: Deserializable> Deserializable for VecDeque<T> {
819    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
820        let len = source.read_usize()?;
821        source.read_many_iter(len)?.collect()
822    }
823
824    /// Returns 1 (the minimum vint length prefix size).
825    ///
826    /// See the note on the `Vec` impl above: budget enforcement during the actual reads provides
827    /// the real protection against oversized length prefixes.
828    fn min_serialized_size() -> usize {
829        1
830    }
831}
832
833impl<K: Deserializable + Ord, V: Deserializable> Deserializable for BTreeMap<K, V> {
834    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
835        let len = source.read_usize()?;
836        let mut map = BTreeMap::new();
837        for entry in source.read_many_iter(len)? {
838            let (key, value) = entry?;
839            if map.insert(key, value).is_some() {
840                return Err(DeserializationError::InvalidValue(String::from(
841                    "duplicate key in BTreeMap encoding",
842                )));
843            }
844        }
845        Ok(map)
846    }
847
848    fn min_serialized_size() -> usize {
849        1 // minimum vint length prefix
850    }
851}
852
853impl<T: Deserializable + Ord> Deserializable for BTreeSet<T> {
854    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
855        let len = source.read_usize()?;
856        let mut set = BTreeSet::new();
857        for item in source.read_many_iter(len)? {
858            if !set.insert(item?) {
859                return Err(DeserializationError::InvalidValue(String::from(
860                    "duplicate item in BTreeSet encoding",
861                )));
862            }
863        }
864        Ok(set)
865    }
866
867    fn min_serialized_size() -> usize {
868        1 // minimum vint length prefix
869    }
870}
871
872impl Deserializable for String {
873    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
874        let len = source.read_usize()?;
875        let data: Vec<u8> = source.read_many_iter(len)?.collect::<Result<_, _>>()?;
876
877        String::from_utf8(data).map_err(|err| DeserializationError::InvalidValue(format!("{err}")))
878    }
879
880    fn min_serialized_size() -> usize {
881        1 // minimum vint length prefix
882    }
883}
884
885impl Deserializable for Arc<str> {
886    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
887        String::read_from(source).map(Arc::from)
888    }
889
890    fn min_serialized_size() -> usize {
891        1 // minimum vint length prefix
892    }
893}
894
895// PLONKY3 FIELD IMPLEMENTATIONS
896// ================================================================================================
897
898impl<F, const D: usize> Serializable for p3_field::extension::BinomialExtensionField<F, D>
899where
900    F: p3_field::Field + p3_field::extension::BinomiallyExtendable<D> + Serializable,
901{
902    fn write_into<W: ByteWriter>(&self, target: &mut W) {
903        let coefficients =
904            <Self as p3_field::BasedVectorSpace<F>>::as_basis_coefficients_slice(self);
905        target.write_many(coefficients);
906    }
907
908    fn get_size_hint(&self) -> usize {
909        <Self as p3_field::BasedVectorSpace<F>>::as_basis_coefficients_slice(self)
910            .iter()
911            .map(Serializable::get_size_hint)
912            .sum()
913    }
914}
915
916impl<F, const D: usize> Deserializable for p3_field::extension::BinomialExtensionField<F, D>
917where
918    F: p3_field::Field + p3_field::extension::BinomiallyExtendable<D> + Deserializable,
919{
920    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
921        Ok(Self::new(<[F; D]>::read_from(source)?))
922    }
923
924    fn min_serialized_size() -> usize {
925        D.saturating_mul(F::min_serialized_size())
926    }
927}
928
929impl Serializable for p3_goldilocks::Goldilocks {
930    fn write_into<W: ByteWriter>(&self, target: &mut W) {
931        use p3_field::PrimeField64;
932        target.write_u64(self.as_canonical_u64());
933    }
934
935    fn get_size_hint(&self) -> usize {
936        size_of::<u64>()
937    }
938}
939
940impl Deserializable for p3_goldilocks::Goldilocks {
941    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
942        use p3_field::integers::QuotientMap;
943
944        let value = source.read_u64()?;
945        Self::from_canonical_checked(value).ok_or_else(|| {
946            DeserializationError::InvalidValue(format!(
947                "value {value} is not a valid Goldilocks field element"
948            ))
949        })
950    }
951}
952
953#[cfg(test)]
954mod tests {
955    use alloc::{collections::VecDeque, sync::Arc};
956
957    use p3_field::extension::BinomialExtensionField;
958    use p3_goldilocks::Goldilocks;
959
960    use super::*;
961
962    #[test]
963    fn arc_str_roundtrip() {
964        let original: Arc<str> = Arc::from("hello world");
965        let bytes = original.to_bytes();
966        let deserialized = Arc::<str>::read_from_bytes(&bytes).unwrap();
967        assert_eq!(original, deserialized);
968    }
969
970    #[test]
971    fn string_roundtrip() {
972        let original = String::from("hello world");
973        let bytes = original.to_bytes();
974        let deserialized = String::read_from_bytes(&bytes).unwrap();
975        assert_eq!(original, deserialized);
976    }
977
978    #[test]
979    fn vec_deque_roundtrip_and_vec_compatibility() {
980        let original = VecDeque::from([1u32, 2, 3]);
981        let bytes = original.to_bytes();
982
983        assert_eq!(VecDeque::<u32>::read_from_bytes(&bytes).unwrap(), original);
984        assert_eq!(Vec::<u32>::read_from_bytes(&bytes).unwrap(), Vec::from([1, 2, 3]));
985        assert_eq!(
986            VecDeque::<u32>::read_from_bytes(&Vec::from([1u32, 2, 3]).to_bytes()).unwrap(),
987            original
988        );
989    }
990
991    #[test]
992    fn binomial_extension_field_roundtrip() {
993        let coefficients = [Goldilocks::new(1), Goldilocks::new(2)];
994        let original = BinomialExtensionField::<Goldilocks, 2>::new(coefficients);
995        let bytes = original.to_bytes();
996
997        assert_eq!(bytes, coefficients.to_bytes());
998        assert_eq!(
999            BinomialExtensionField::<Goldilocks, 2>::read_from_bytes(&bytes).unwrap(),
1000            original
1001        );
1002    }
1003
1004    #[test]
1005    fn empty_string_roundtrip() {
1006        let arc: Arc<str> = Arc::from("");
1007        let bytes = arc.to_bytes();
1008        let deserialized = Arc::<str>::read_from_bytes(&bytes).unwrap();
1009        assert_eq!(deserialized, Arc::from(""));
1010
1011        let string = String::from("");
1012        let bytes = string.to_bytes();
1013        let deserialized = String::read_from_bytes(&bytes).unwrap();
1014        assert_eq!(deserialized, "");
1015    }
1016
1017    #[test]
1018    fn multibyte_utf8_roundtrip() {
1019        let text = "héllo 🌍";
1020
1021        let arc: Arc<str> = Arc::from(text);
1022        let bytes = arc.to_bytes();
1023        let deserialized = Arc::<str>::read_from_bytes(&bytes).unwrap();
1024        assert_eq!(&*deserialized, text);
1025
1026        let string = String::from(text);
1027        let bytes = string.to_bytes();
1028        let deserialized = String::read_from_bytes(&bytes).unwrap();
1029        assert_eq!(deserialized, text);
1030
1031        // Cross-compat: Arc<str> bytes can be read as String and vice versa
1032        let arc_bytes = Arc::<str>::from(text).to_bytes();
1033        let string_bytes = String::from(text).to_bytes();
1034        assert_eq!(arc_bytes, string_bytes);
1035        assert_eq!(String::read_from_bytes(&arc_bytes).unwrap(), text);
1036        assert_eq!(&*Arc::<str>::read_from_bytes(&string_bytes).unwrap(), text);
1037    }
1038
1039    #[test]
1040    fn arc_str_string_cross_compat() {
1041        // Arc<str> -> bytes -> String
1042        let arc: Arc<str> = Arc::from("cross type");
1043        let bytes = arc.to_bytes();
1044        let as_string = String::read_from_bytes(&bytes).unwrap();
1045        assert_eq!(as_string, "cross type");
1046
1047        // String -> bytes -> Arc<str>
1048        let string = String::from("other direction");
1049        let bytes = string.to_bytes();
1050        let as_arc = Arc::<str>::read_from_bytes(&bytes).unwrap();
1051        assert_eq!(&*as_arc, "other direction");
1052    }
1053
1054    #[test]
1055    fn btree_map_rejects_duplicate_keys() {
1056        let mut bytes = Vec::new();
1057        bytes.extend_from_slice(&2usize.to_bytes());
1058        bytes.extend_from_slice(&7u8.to_bytes());
1059        bytes.extend_from_slice(&1u8.to_bytes());
1060        bytes.extend_from_slice(&7u8.to_bytes());
1061        bytes.extend_from_slice(&2u8.to_bytes());
1062
1063        let result = BTreeMap::<u8, u8>::read_from_bytes(&bytes);
1064
1065        assert!(matches!(result, Err(DeserializationError::InvalidValue(_))));
1066    }
1067
1068    #[test]
1069    fn btree_set_rejects_duplicate_items() {
1070        let mut bytes = Vec::new();
1071        bytes.extend_from_slice(&2usize.to_bytes());
1072        bytes.extend_from_slice(&7u8.to_bytes());
1073        bytes.extend_from_slice(&7u8.to_bytes());
1074
1075        let result = BTreeSet::<u8>::read_from_bytes(&bytes);
1076
1077        assert!(matches!(result, Err(DeserializationError::InvalidValue(_))));
1078    }
1079
1080    #[test]
1081    fn read_bounded_len_accepts_a_length_backed_by_enough_input() {
1082        let mut bytes = 3usize.to_bytes();
1083        bytes.extend_from_slice(&[0u8; 3]);
1084        let mut source = SliceReader::new(&bytes);
1085
1086        assert_eq!(read_bounded_len(&mut source, "elements", 1).unwrap(), 3);
1087    }
1088
1089    #[test]
1090    fn read_bounded_len_rejects_a_length_exceeding_remaining_input() {
1091        let mut bytes = 8usize.to_bytes();
1092        bytes.extend_from_slice(&[0u8; 3]);
1093        let mut source = SliceReader::new(&bytes);
1094
1095        assert!(matches!(
1096            read_bounded_len(&mut source, "elements", 1),
1097            Err(DeserializationError::InvalidValue(_))
1098        ));
1099    }
1100
1101    #[test]
1102    fn read_bounded_len_rejects_a_length_exceeding_the_budget() {
1103        let mut bytes = 64usize.to_bytes();
1104        bytes.extend_from_slice(&[0u8; 64]);
1105        let mut source = BudgetedReader::new(SliceReader::new(&bytes), 16);
1106
1107        assert!(matches!(
1108            read_bounded_len(&mut source, "elements", 1),
1109            Err(DeserializationError::InvalidValue(_))
1110        ));
1111    }
1112
1113    #[test]
1114    fn validate_bounded_len_rejects_a_length_overflowing_the_element_size() {
1115        let source = SliceReader::new(&[]);
1116
1117        assert!(matches!(
1118            validate_bounded_len(&source, "elements", usize::MAX, 2),
1119            Err(DeserializationError::InvalidValue(_))
1120        ));
1121    }
1122
1123    #[test]
1124    fn budgeted_vec_of_one_element_tuples_accepts_exact_budget() {
1125        let values = Vec::<(usize,)>::from([(0,), (1,), (2,)]);
1126        let bytes = values.to_bytes();
1127
1128        let decoded = Vec::<(usize,)>::read_from_bytes_with_budget(&bytes, bytes.len()).unwrap();
1129
1130        assert_eq!(decoded, values);
1131    }
1132}