miden_serde_utils/byte_reader.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(feature = "std")]
7use alloc::string::ToString;
8use alloc::{format, string::String, vec::Vec};
9#[cfg(feature = "std")]
10use core::cell::{Ref, RefCell};
11#[cfg(feature = "std")]
12use std::io::BufRead;
13
14use crate::{Deserializable, DeserializationError};
15
16// BYTE READER TRAIT
17// ================================================================================================
18
19/// Defines how primitive values are to be read from `Self`.
20///
21/// Whenever data is read from the reader using any of the `read_*` functions, the reader advances
22/// to the next unread byte. If the error occurs, the reader is not rolled back to the state prior
23/// to calling any of the function.
24pub trait ByteReader {
25 // REQUIRED METHODS
26 // --------------------------------------------------------------------------------------------
27
28 /// Returns a single byte read from `self`.
29 ///
30 /// # Errors
31 /// Returns a [DeserializationError] error the reader is at EOF.
32 fn read_u8(&mut self) -> Result<u8, DeserializationError>;
33
34 /// Returns the next byte to be read from `self` without advancing the reader to the next byte.
35 ///
36 /// # Errors
37 /// Returns a [DeserializationError] error the reader is at EOF.
38 fn peek_u8(&self) -> Result<u8, DeserializationError>;
39
40 /// Returns a slice of bytes of the specified length read from `self`.
41 ///
42 /// # Errors
43 /// Returns a [DeserializationError] if a slice of the specified length could not be read
44 /// from `self`.
45 fn read_slice(&mut self, len: usize) -> Result<&[u8], DeserializationError>;
46
47 /// Returns a byte array of length `N` read from `self`.
48 ///
49 /// # Errors
50 /// Returns a [DeserializationError] if an array of the specified length could not be read
51 /// from `self`.
52 fn read_array<const N: usize>(&mut self) -> Result<[u8; N], DeserializationError>;
53
54 /// Checks if it is possible to read at least `num_bytes` bytes from this ByteReader
55 ///
56 /// # Errors
57 /// Returns an error if, when reading the requested number of bytes, we go beyond the
58 /// the data available in the reader.
59 fn check_eor(&self, num_bytes: usize) -> Result<(), DeserializationError>;
60
61 /// Returns true if there are more bytes left to be read from `self`.
62 fn has_more_bytes(&self) -> bool;
63
64 /// Returns the maximum number of elements that can be safely allocated, given each
65 /// element occupies `element_size` bytes when serialized.
66 ///
67 /// This can be used by callers to pre-validate collection lengths before iterating,
68 /// preventing denial-of-service attacks from malicious length prefixes that claim
69 /// billions of elements.
70 ///
71 /// The default implementation returns `usize::MAX`, meaning no limit is enforced.
72 /// [`BudgetedReader`] overrides this to return `remaining_budget / element_size`,
73 /// providing tight, adaptive limits based on the caller's budget. For zero-sized serialized
74 /// elements, [`BudgetedReader`] returns zero so non-empty collections are rejected.
75 ///
76 /// # Arguments
77 /// * `element_size` - The serialized size of one element, from
78 /// [`Deserializable::min_serialized_size`]. Defaults to `size_of::<D>()` but can be
79 /// overridden for types where serialized size differs from in-memory size.
80 fn max_alloc(&self, _element_size: usize) -> usize {
81 usize::MAX
82 }
83
84 // PROVIDED METHODS
85 // --------------------------------------------------------------------------------------------
86
87 /// Returns a boolean value read from `self` consuming 1 byte from the reader.
88 ///
89 /// # Errors
90 /// Returns a [DeserializationError] if a u16 value could not be read from `self`.
91 fn read_bool(&mut self) -> Result<bool, DeserializationError> {
92 let byte = self.read_u8()?;
93 match byte {
94 0 => Ok(false),
95 1 => Ok(true),
96 _ => Err(DeserializationError::InvalidValue(format!("{byte} is not a boolean value"))),
97 }
98 }
99
100 /// Returns a u16 value read from `self` in little-endian byte order.
101 ///
102 /// # Errors
103 /// Returns a [DeserializationError] if a u16 value could not be read from `self`.
104 fn read_u16(&mut self) -> Result<u16, DeserializationError> {
105 let bytes = self.read_array::<2>()?;
106 Ok(u16::from_le_bytes(bytes))
107 }
108
109 /// Returns a u32 value read from `self` in little-endian byte order.
110 ///
111 /// # Errors
112 /// Returns a [DeserializationError] if a u32 value could not be read from `self`.
113 fn read_u32(&mut self) -> Result<u32, DeserializationError> {
114 let bytes = self.read_array::<4>()?;
115 Ok(u32::from_le_bytes(bytes))
116 }
117
118 /// Returns a u64 value read from `self` in little-endian byte order.
119 ///
120 /// # Errors
121 /// Returns a [DeserializationError] if a u64 value could not be read from `self`.
122 fn read_u64(&mut self) -> Result<u64, DeserializationError> {
123 let bytes = self.read_array::<8>()?;
124 Ok(u64::from_le_bytes(bytes))
125 }
126
127 /// Returns a u128 value read from `self` in little-endian byte order.
128 ///
129 /// # Errors
130 /// Returns a [DeserializationError] if a u128 value could not be read from `self`.
131 fn read_u128(&mut self) -> Result<u128, DeserializationError> {
132 let bytes = self.read_array::<16>()?;
133 Ok(u128::from_le_bytes(bytes))
134 }
135
136 /// Returns a usize value read from `self` in [vint64](https://docs.rs/vint64/latest/vint64/)
137 /// format.
138 ///
139 /// # Errors
140 /// Returns a [DeserializationError] if:
141 /// * usize value could not be read from `self`.
142 /// * encoded value is greater than `usize` maximum value on a given platform.
143 fn read_usize(&mut self) -> Result<usize, DeserializationError> {
144 let first_byte = self.peek_u8()?;
145 let length = first_byte.trailing_zeros() as usize + 1;
146
147 let result = if length == 9 {
148 // 9-byte special case
149 self.read_u8()?;
150 let value = self.read_array::<8>()?;
151 u64::from_le_bytes(value)
152 } else {
153 let mut encoded = [0u8; 8];
154 let value = self.read_slice(length)?;
155 encoded[..length].copy_from_slice(value);
156 u64::from_le_bytes(encoded) >> length
157 };
158
159 // check if the result value is within acceptable bounds for `usize` on a given platform
160 if result > usize::MAX as u64 {
161 return Err(DeserializationError::InvalidValue(format!(
162 "Encoded value must be less than {}, but {} was provided",
163 usize::MAX,
164 result
165 )));
166 }
167
168 Ok(result as usize)
169 }
170
171 /// Returns a byte vector of the specified length read from `self`.
172 ///
173 /// # Errors
174 /// Returns a [DeserializationError] if a vector of the specified length could not be read
175 /// from `self`.
176 fn read_vec(&mut self, len: usize) -> Result<Vec<u8>, DeserializationError> {
177 let data = self.read_slice(len)?;
178 Ok(data.to_vec())
179 }
180
181 /// Returns a String of the specified length read from `self`.
182 ///
183 /// # Errors
184 /// Returns a [DeserializationError] if a String of the specified length could not be read
185 /// from `self`.
186 fn read_string(&mut self, num_bytes: usize) -> Result<String, DeserializationError> {
187 let data = self.read_vec(num_bytes)?;
188 String::from_utf8(data).map_err(|err| DeserializationError::InvalidValue(format!("{err}")))
189 }
190
191 /// Reads a deserializable value from `self`.
192 ///
193 /// # Errors
194 /// Returns a [DeserializationError] if the specified value could not be read from `self`.
195 fn read<D>(&mut self) -> Result<D, DeserializationError>
196 where
197 Self: Sized,
198 D: Deserializable,
199 {
200 D::read_from(self)
201 }
202
203 /// Returns an iterator that deserializes `num_elements` instances of `D` from this reader.
204 ///
205 /// This method validates the requested count against the reader's capacity before returning
206 /// the iterator, rejecting implausible lengths early. Each element is then deserialized
207 /// lazily as the iterator is consumed.
208 ///
209 /// # Errors
210 ///
211 /// Returns an error if `num_elements` exceeds `self.max_alloc(D::min_serialized_size())`,
212 /// indicating the reader cannot allocate that many elements.
213 ///
214 /// # Example
215 ///
216 /// ```ignore
217 /// // Collect into a Vec
218 /// let items: Vec<u64> = reader
219 /// .read_many_iter::<u64>(count)?
220 /// .collect::<Result<_, _>>()?;
221 ///
222 /// // Collect directly into a BTreeMap (no intermediate Vec)
223 /// let map: BTreeMap<K, V> = reader
224 /// .read_many_iter::<(K, V)>(count)?
225 /// .collect::<Result<_, _>>()?;
226 /// ```
227 fn read_many_iter<D>(
228 &mut self,
229 num_elements: usize,
230 ) -> Result<ReadManyIter<'_, Self, D>, DeserializationError>
231 where
232 Self: Sized,
233 D: Deserializable,
234 {
235 let max_elements = self.max_alloc(D::min_serialized_size());
236 if num_elements > max_elements {
237 return Err(DeserializationError::InvalidValue(format!(
238 "requested {num_elements} elements but reader can provide at most {max_elements}"
239 )));
240 }
241 Ok(ReadManyIter {
242 reader: self,
243 remaining: num_elements,
244 _item: core::marker::PhantomData,
245 })
246 }
247}
248
249// READ MANY ITERATOR
250// ================================================================================================
251
252/// Iterator that lazily deserializes elements from a [`ByteReader`].
253///
254/// Created by [`ByteReader::read_many_iter`]. Each call to `next()` deserializes one element.
255/// This avoids upfront allocation and naturally integrates with [`BudgetedReader`] for
256/// protection against malicious inputs.
257pub struct ReadManyIter<'reader, R: ByteReader, D: Deserializable> {
258 reader: &'reader mut R,
259 remaining: usize,
260 _item: core::marker::PhantomData<D>,
261}
262
263impl<'reader, R: ByteReader, D: Deserializable> Iterator for ReadManyIter<'reader, R, D> {
264 type Item = Result<D, DeserializationError>;
265
266 fn next(&mut self) -> Option<Self::Item> {
267 if self.remaining > 0 {
268 self.remaining -= 1;
269 Some(D::read_from(self.reader))
270 } else {
271 None
272 }
273 }
274
275 fn size_hint(&self) -> (usize, Option<usize>) {
276 (self.remaining, Some(self.remaining))
277 }
278}
279
280impl<'reader, R: ByteReader, D: Deserializable> ExactSizeIterator for ReadManyIter<'reader, R, D> {}
281
282// STANDARD LIBRARY ADAPTER
283// ================================================================================================
284
285/// An adapter of [ByteReader] to any type that implements [std::io::Read]
286///
287/// In particular, this covers things like [std::fs::File], standard input, etc.
288#[cfg(feature = "std")]
289pub struct ReadAdapter<'a> {
290 // NOTE: The [ByteReader] trait does not currently support reader implementations that require
291 // mutation during `peek_u8`, `has_more_bytes`, and `check_eor`. These (or equivalent)
292 // operations on the standard library [std::io::BufRead] trait require a mutable reference, as
293 // it may be necessary to read from the underlying input to implement them.
294 //
295 // To handle this, we wrap the underlying reader in an [RefCell], this allows us to mutate the
296 // reader if necessary during a call to one of the above-mentioned trait methods, without
297 // sacrificing safety - at the cost of enforcing Rust's borrowing semantics dynamically.
298 //
299 // This should not be a problem in practice, except in the case where `read_slice` is called,
300 // and the reference returned is from `reader` directly, rather than `buf`. If a call to one
301 // of the above-mentioned methods is made while that reference is live, and we attempt to read
302 // from `reader`, a panic will occur.
303 //
304 // Ultimately, this should be addressed by making the [ByteReader] trait align with the
305 // standard library I/O traits, so this is a temporary solution.
306 reader: RefCell<std::io::BufReader<&'a mut dyn std::io::Read>>,
307 // A temporary buffer to store chunks read from `reader` that are larger than what is required
308 // for the higher-level [ByteReader] APIs.
309 //
310 // By default we attempt to satisfy reads from `reader` directly, but that is not always
311 // possible.
312 buf: Vec<u8>,
313 // The position in `buf` at which we should start reading the next byte, when `buf` is
314 // non-empty.
315 pos: usize,
316 // This is set when we attempt to read from `reader` and get an empty buffer. This indicates
317 // that once we exhaust `buf`, we have truly reached end-of-file.
318 //
319 // We will use this to more accurately handle functions like `has_more_bytes` when this is set.
320 guaranteed_eof: bool,
321}
322
323#[cfg(feature = "std")]
324impl<'a> ReadAdapter<'a> {
325 /// Create a new [ByteReader] adapter for the given implementation of [std::io::Read]
326 pub fn new(reader: &'a mut dyn std::io::Read) -> Self {
327 Self {
328 reader: RefCell::new(std::io::BufReader::with_capacity(256, reader)),
329 buf: Default::default(),
330 pos: 0,
331 guaranteed_eof: false,
332 }
333 }
334
335 /// Get the internal adapter buffer as a (possibly empty) slice of bytes
336 #[inline(always)]
337 fn buffer(&self) -> &[u8] {
338 self.buf.get(self.pos..).unwrap_or(&[])
339 }
340
341 /// Get the internal adapter buffer as a slice of bytes, or `None` if the buffer is empty
342 #[inline(always)]
343 fn non_empty_buffer(&self) -> Option<&[u8]> {
344 self.buf.get(self.pos..).filter(|b| !b.is_empty())
345 }
346
347 /// Return the current reader buffer as a (possibly empty) slice of bytes.
348 ///
349 /// This buffer being empty _does not_ mean we're at EOF, you must call
350 /// [non_empty_reader_buffer_mut] first.
351 #[inline(always)]
352 fn reader_buffer(&self) -> Ref<'_, [u8]> {
353 Ref::map(self.reader.borrow(), |r| r.buffer())
354 }
355
356 /// Return the current reader buffer, reading from the underlying reader
357 /// if the buffer is empty.
358 ///
359 /// Returns `Ok` only if the buffer is non-empty, and no errors occurred
360 /// while filling it (if filling was needed).
361 fn non_empty_reader_buffer_mut(&mut self) -> Result<&[u8], DeserializationError> {
362 use std::io::ErrorKind;
363 let buf = self.reader.get_mut().fill_buf().map_err(|e| match e.kind() {
364 ErrorKind::UnexpectedEof => DeserializationError::UnexpectedEOF,
365 e => DeserializationError::UnknownError(e.to_string()),
366 })?;
367 if buf.is_empty() {
368 self.guaranteed_eof = true;
369 Err(DeserializationError::UnexpectedEOF)
370 } else {
371 Ok(buf)
372 }
373 }
374
375 /// Same as [non_empty_reader_buffer_mut], but with dynamically-enforced
376 /// borrow check rules so that it can be called in functions like `peek_u8`.
377 ///
378 /// This comes with overhead for the dynamic checks, so you should prefer
379 /// to call [non_empty_reader_buffer_mut] if you already have a mutable
380 /// reference to `self`
381 fn non_empty_reader_buffer(&self) -> Result<Ref<'_, [u8]>, DeserializationError> {
382 use std::io::ErrorKind;
383 let mut reader = self.reader.borrow_mut();
384 let buf = reader.fill_buf().map_err(|e| match e.kind() {
385 ErrorKind::UnexpectedEof => DeserializationError::UnexpectedEOF,
386 e => DeserializationError::UnknownError(e.to_string()),
387 })?;
388 if buf.is_empty() {
389 Err(DeserializationError::UnexpectedEOF)
390 } else {
391 // Re-borrow immutably
392 drop(reader);
393 Ok(self.reader_buffer())
394 }
395 }
396
397 /// Returns true if there is sufficient capacity remaining in `buf` to hold `n` bytes
398 #[inline]
399 fn has_remaining_capacity(&self, n: usize) -> bool {
400 let remaining = self.buf.capacity() - self.buffer().len();
401 remaining >= n
402 }
403
404 /// Takes the next byte from the input, returning an error if the operation fails
405 fn pop(&mut self) -> Result<u8, DeserializationError> {
406 if let Some(byte) = self.non_empty_buffer().map(|b| b[0]) {
407 self.pos += 1;
408 return Ok(byte);
409 }
410 let result = self.non_empty_reader_buffer_mut().map(|b| b[0]);
411 if result.is_ok() {
412 self.reader.get_mut().consume(1);
413 } else {
414 self.guaranteed_eof = true;
415 }
416 result
417 }
418
419 /// Takes the next `N` bytes from the input as an array, returning an error if the operation
420 /// fails
421 fn read_exact<const N: usize>(&mut self) -> Result<[u8; N], DeserializationError> {
422 let mut output = [0; N];
423 let buf = self.buffer();
424
425 if buf.len() >= N {
426 output.copy_from_slice(&buf[..N]);
427 self.pos += N;
428
429 if self.buffer().is_empty() {
430 unsafe {
431 self.buf.set_len(0);
432 }
433 self.pos = 0;
434 }
435
436 return Ok(output);
437 }
438
439 if buf.is_empty() {
440 let reader_buf = self.non_empty_reader_buffer_mut()?;
441 if reader_buf.len() >= N {
442 output.copy_from_slice(&reader_buf[..N]);
443 self.reader.get_mut().consume(N);
444 return Ok(output);
445 }
446 }
447
448 output.copy_from_slice(<Self as ByteReader>::read_slice(self, N)?);
449 Ok(output)
450 }
451
452 /// Fill `self.buf` with `count` bytes
453 ///
454 /// This should only be called when we can't read from the reader directly
455 fn buffer_at_least(&mut self, count: usize) -> Result<(), DeserializationError> {
456 // Read until we have at least `count` bytes, or until we reach end-of-file,
457 // which ever comes first.
458 loop {
459 // If we have successfully read `count` bytes, we're done
460 if self.buffer().len() >= count {
461 break Ok(());
462 }
463
464 // This operation will return an error if the underlying reader hits EOF
465 self.non_empty_reader_buffer_mut()?;
466
467 // Extend `self.buf` with the bytes read from the underlying reader.
468 //
469 // NOTE: We have to re-borrow the reader buffer here, since we can't get a mutable
470 // reference to `self.buf` while holding an immutable reference to the reader buffer.
471 let reader = self.reader.get_mut();
472 let buf = reader.buffer();
473 let consumed = buf.len();
474 self.buf.extend_from_slice(buf);
475 reader.consume(consumed);
476 }
477 }
478}
479
480#[cfg(feature = "std")]
481impl ByteReader for ReadAdapter<'_> {
482 #[inline(always)]
483 fn read_u8(&mut self) -> Result<u8, DeserializationError> {
484 self.pop()
485 }
486
487 /// NOTE: If we happen to not have any bytes buffered yet when this is called, then we will be
488 /// forced to try and read from the underlying reader. This requires a mutable reference, which
489 /// is obtained dynamically via [RefCell].
490 ///
491 /// <div class="warning">
492 /// Callers must ensure that they do not hold any immutable references to the buffer of this
493 /// reader when calling this function so as to avoid a situation in which the dynamic borrow
494 /// check fails. Specifically, you must not be holding a reference to the result of
495 /// [Self::read_slice] when this function is called.
496 /// </div>
497 fn peek_u8(&self) -> Result<u8, DeserializationError> {
498 if let Some(byte) = self.buffer().first() {
499 return Ok(*byte);
500 }
501 self.non_empty_reader_buffer().map(|b| b[0])
502 }
503
504 fn read_slice(&mut self, len: usize) -> Result<&[u8], DeserializationError> {
505 // Edge case
506 if len == 0 {
507 return Ok(&[]);
508 }
509
510 // If we have unused buffer, and the consumed portion is
511 // large enough, we will move the unused portion of the buffer
512 // to the start, freeing up bytes at the end for more reads
513 // before forcing a reallocation
514 let should_optimize_storage = self.pos >= 16 && !self.has_remaining_capacity(len);
515 if should_optimize_storage {
516 // We're going to optimize storage first
517 let buf = self.buffer();
518 let src = buf.as_ptr();
519 let count = buf.len();
520 let dst = self.buf.as_mut_ptr();
521 unsafe {
522 core::ptr::copy(src, dst, count);
523 self.buf.set_len(count);
524 self.pos = 0;
525 }
526 }
527
528 // Fill the buffer so we have at least `len` bytes available,
529 // this will return an error if we hit EOF first
530 self.buffer_at_least(len)?;
531
532 let slice = &self.buf[self.pos..(self.pos + len)];
533 self.pos += len;
534 Ok(slice)
535 }
536
537 #[inline]
538 fn read_array<const N: usize>(&mut self) -> Result<[u8; N], DeserializationError> {
539 if N == 0 {
540 return Ok([0; N]);
541 }
542 self.read_exact()
543 }
544
545 fn check_eor(&self, num_bytes: usize) -> Result<(), DeserializationError> {
546 // Do we have sufficient data in the local buffer?
547 let buffer_len = self.buffer().len();
548 if buffer_len >= num_bytes {
549 return Ok(());
550 }
551
552 // What about if we include what is in the local buffer and the reader's buffer?
553 let reader_buffer_len = self.non_empty_reader_buffer().map(|b| b.len())?;
554 let buffer_len = buffer_len + reader_buffer_len;
555 if buffer_len >= num_bytes {
556 return Ok(());
557 }
558
559 // We have no more input, thus can't fulfill a request of `num_bytes`
560 if self.guaranteed_eof {
561 return Err(DeserializationError::UnexpectedEOF);
562 }
563
564 // Because this function is read-only, we must optimistically assume we can read `num_bytes`
565 // from the input, and fail later if that does not hold. We know we're not at EOF yet, but
566 // that's all we can say without buffering more from the reader. We could make use of
567 // `buffer_at_least`, which would guarantee a correct result, but it would also impose
568 // additional restrictions on the use of this function, e.g. not using it while holding a
569 // reference returned from `read_slice`. Since it is not a memory safety violation to return
570 // an optimistic result here, it makes for a better tradeoff.
571 Ok(())
572 }
573
574 #[inline]
575 fn has_more_bytes(&self) -> bool {
576 !self.buffer().is_empty() || self.non_empty_reader_buffer().is_ok()
577 }
578}
579
580// CURSOR
581// ================================================================================================
582
583#[cfg(feature = "std")]
584macro_rules! cursor_remaining_buf {
585 ($cursor:ident) => {{
586 let buf = $cursor.get_ref().as_ref();
587 let start = $cursor.position().min(buf.len() as u64) as usize;
588 &buf[start..]
589 }};
590}
591
592#[cfg(feature = "std")]
593impl<T: AsRef<[u8]>> ByteReader for std::io::Cursor<T> {
594 fn read_u8(&mut self) -> Result<u8, DeserializationError> {
595 let buf = cursor_remaining_buf!(self);
596 if buf.is_empty() {
597 Err(DeserializationError::UnexpectedEOF)
598 } else {
599 let byte = buf[0];
600 self.set_position(self.position() + 1);
601 Ok(byte)
602 }
603 }
604
605 fn peek_u8(&self) -> Result<u8, DeserializationError> {
606 cursor_remaining_buf!(self)
607 .first()
608 .copied()
609 .ok_or(DeserializationError::UnexpectedEOF)
610 }
611
612 fn read_slice(&mut self, len: usize) -> Result<&[u8], DeserializationError> {
613 let pos = self.position();
614 let size = self.get_ref().as_ref().len() as u64;
615 if size.saturating_sub(pos) < len as u64 {
616 Err(DeserializationError::UnexpectedEOF)
617 } else {
618 self.set_position(pos + len as u64);
619 let start = pos.min(size) as usize;
620 Ok(&self.get_ref().as_ref()[start..(start + len)])
621 }
622 }
623
624 fn read_array<const N: usize>(&mut self) -> Result<[u8; N], DeserializationError> {
625 self.read_slice(N).map(|bytes| {
626 let mut result = [0u8; N];
627 result.copy_from_slice(bytes);
628 result
629 })
630 }
631
632 fn check_eor(&self, num_bytes: usize) -> Result<(), DeserializationError> {
633 if cursor_remaining_buf!(self).len() >= num_bytes {
634 Ok(())
635 } else {
636 Err(DeserializationError::UnexpectedEOF)
637 }
638 }
639
640 #[inline]
641 fn has_more_bytes(&self) -> bool {
642 let pos = self.position();
643 let size = self.get_ref().as_ref().len() as u64;
644 pos < size
645 }
646}
647
648// SLICE READER
649// ================================================================================================
650
651/// Implements [ByteReader] trait for a slice of bytes.
652///
653/// NOTE: If you are building with the `std` feature, you should probably prefer [std::io::Cursor]
654/// instead. However, [SliceReader] is still useful in no-std environments until stabilization of
655/// the `core_io_borrowed_buf` feature.
656pub struct SliceReader<'a> {
657 source: &'a [u8],
658 pos: usize,
659}
660
661impl<'a> SliceReader<'a> {
662 /// Creates a new slice reader from the specified slice.
663 pub fn new(source: &'a [u8]) -> Self {
664 SliceReader { source, pos: 0 }
665 }
666}
667
668impl ByteReader for SliceReader<'_> {
669 fn read_u8(&mut self) -> Result<u8, DeserializationError> {
670 self.check_eor(1)?;
671 let result = self.source[self.pos];
672 self.pos += 1;
673 Ok(result)
674 }
675
676 fn peek_u8(&self) -> Result<u8, DeserializationError> {
677 self.check_eor(1)?;
678 Ok(self.source[self.pos])
679 }
680
681 fn read_slice(&mut self, len: usize) -> Result<&[u8], DeserializationError> {
682 self.check_eor(len)?;
683 let result = &self.source[self.pos..self.pos + len];
684 self.pos += len;
685 Ok(result)
686 }
687
688 fn read_array<const N: usize>(&mut self) -> Result<[u8; N], DeserializationError> {
689 self.check_eor(N)?;
690 let mut result = [0_u8; N];
691 result.copy_from_slice(&self.source[self.pos..self.pos + N]);
692 self.pos += N;
693 Ok(result)
694 }
695
696 fn check_eor(&self, num_bytes: usize) -> Result<(), DeserializationError> {
697 self.pos
698 .checked_add(num_bytes)
699 .filter(|end| *end <= self.source.len())
700 .map(|_| ())
701 .ok_or(DeserializationError::UnexpectedEOF)
702 }
703
704 fn has_more_bytes(&self) -> bool {
705 self.pos < self.source.len()
706 }
707}
708
709// BUDGETED READER
710// ================================================================================================
711
712/// A reader wrapper that enforces a byte budget during deserialization.
713///
714/// # Threat Model
715///
716/// Malicious input can attack deserialization in two ways:
717///
718/// 1. **Fake length prefix**: Input claims `len = 2^60` elements, causing allocation of a huge
719/// `Vec` before any data is read.
720/// 2. **Oversized input**: Attacker sends gigabytes of valid-looking data to exhaust memory over
721/// time.
722///
723/// # Defense Strategy
724///
725/// Use `BudgetedReader` to limit total bytes consumed. Its [`max_alloc`](ByteReader::max_alloc)
726/// method derives a bound from the remaining budget, which
727/// [`read_many_iter`](ByteReader::read_many_iter) checks before iterating.
728///
729/// ## Problem: SliceReader alone doesn't bound allocations
730///
731/// ```
732/// use miden_serde_utils::{ByteReader, Deserializable, SliceReader};
733///
734/// // Malicious input: length prefix says 1 billion u64s, but only 16 bytes of data
735/// let mut data = Vec::new();
736/// data.push(0u8); // vint64 9-byte marker
737/// data.extend_from_slice(&1_000_000_000u64.to_le_bytes());
738/// data.extend_from_slice(&[0u8; 16]);
739///
740/// // SliceReader and read_from_bytes are unbudgeted. Use read_from_bytes_with_budget
741/// // or wrap SliceReader in BudgetedReader when reading untrusted input.
742/// let reader = SliceReader::new(&data);
743/// assert_eq!(reader.max_alloc(8), usize::MAX);
744/// ```
745///
746/// ## Solution: BudgetedReader bounds allocations via max_alloc
747///
748/// ```
749/// use miden_serde_utils::{BudgetedReader, ByteReader, Deserializable, SliceReader};
750///
751/// // Same malicious input
752/// let mut data = Vec::new();
753/// data.push(0u8);
754/// data.extend_from_slice(&1_000_000_000u64.to_le_bytes());
755/// data.extend_from_slice(&[0u8; 16]);
756///
757/// // BudgetedReader with 64-byte budget: max_alloc(8) = 64/8 = 8 elements
758/// let inner = SliceReader::new(&data);
759/// let reader = BudgetedReader::new(inner, 64);
760/// assert_eq!(reader.max_alloc(8), 8);
761///
762/// // read_many_iter rejects the 1B length since 1B > 8
763/// let result = Vec::<u64>::read_from_bytes_with_budget(&data, 64);
764/// assert!(result.is_err());
765/// ```
766///
767/// ## Best practice: Set budget to expected input size
768///
769/// ```
770/// use miden_serde_utils::{ByteWriter, Deserializable, Serializable};
771///
772/// // Legitimate input: 3 u64s, properly serialized
773/// let original = vec![1u64, 2, 3];
774/// let mut data = Vec::new();
775/// original.write_into(&mut data);
776///
777/// // Budget = data.len() bounds both fake lengths and total consumption
778/// let result = Vec::<u64>::read_from_bytes_with_budget(&data, data.len());
779/// assert_eq!(result.unwrap(), vec![1, 2, 3]);
780/// ```
781pub struct BudgetedReader<R> {
782 inner: R,
783 remaining: usize,
784}
785
786impl<R> BudgetedReader<R> {
787 /// Wraps a reader with the specified byte budget.
788 pub fn new(inner: R, budget: usize) -> Self {
789 Self { inner, remaining: budget }
790 }
791
792 /// Returns remaining budget in bytes.
793 pub fn remaining(&self) -> usize {
794 self.remaining
795 }
796
797 /// Consumes budget, returning an error if insufficient.
798 fn consume_budget(&mut self, n: usize) -> Result<(), DeserializationError> {
799 if n > self.remaining {
800 return Err(DeserializationError::InvalidValue(format!(
801 "budget exhausted: requested {n} bytes, {} remaining",
802 self.remaining
803 )));
804 }
805 self.remaining -= n;
806 Ok(())
807 }
808}
809
810impl<R: ByteReader> ByteReader for BudgetedReader<R> {
811 fn read_u8(&mut self) -> Result<u8, DeserializationError> {
812 self.consume_budget(1)?;
813 self.inner.read_u8()
814 }
815
816 fn peek_u8(&self) -> Result<u8, DeserializationError> {
817 // peek doesn't consume budget since it doesn't advance the reader
818 self.inner.peek_u8()
819 }
820
821 fn read_slice(&mut self, len: usize) -> Result<&[u8], DeserializationError> {
822 self.consume_budget(len)?;
823 self.inner.read_slice(len)
824 }
825
826 fn read_array<const N: usize>(&mut self) -> Result<[u8; N], DeserializationError> {
827 self.consume_budget(N)?;
828 self.inner.read_array()
829 }
830
831 fn check_eor(&self, num_bytes: usize) -> Result<(), DeserializationError> {
832 // check budget first, then delegate
833 if num_bytes > self.remaining {
834 return Err(DeserializationError::InvalidValue(format!(
835 "budget exhausted: requested {num_bytes} bytes, {} remaining",
836 self.remaining
837 )));
838 }
839 self.inner.check_eor(num_bytes)
840 }
841
842 fn has_more_bytes(&self) -> bool {
843 self.remaining > 0 && self.inner.has_more_bytes()
844 }
845
846 fn max_alloc(&self, element_size: usize) -> usize {
847 if element_size == 0 {
848 return 0;
849 }
850 self.remaining / element_size
851 }
852}
853
854#[cfg(all(test, feature = "std"))]
855mod tests {
856 use core::mem::size_of;
857 use std::io::{Cursor, Read};
858
859 use super::*;
860 use crate::ByteWriter;
861
862 struct ChunkedReader {
863 data: Vec<u8>,
864 pos: usize,
865 chunk_size: usize,
866 }
867
868 impl ChunkedReader {
869 fn new(data: Vec<u8>, chunk_size: usize) -> Self {
870 Self { data, pos: 0, chunk_size }
871 }
872 }
873
874 impl Read for ChunkedReader {
875 fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
876 let remaining = &self.data[self.pos..];
877 let len = remaining.len().min(buf.len()).min(self.chunk_size);
878 buf[..len].copy_from_slice(&remaining[..len]);
879 self.pos += len;
880 Ok(len)
881 }
882 }
883
884 #[test]
885 fn read_adapter_empty() {
886 let mut reader = std::io::empty();
887 let mut adapter = ReadAdapter::new(&mut reader);
888 assert!(!adapter.has_more_bytes());
889 assert_eq!(adapter.check_eor(8), Err(DeserializationError::UnexpectedEOF));
890 assert_eq!(adapter.peek_u8(), Err(DeserializationError::UnexpectedEOF));
891 assert_eq!(adapter.read_u8(), Err(DeserializationError::UnexpectedEOF));
892 assert_eq!(adapter.read_slice(0), Ok([].as_slice()));
893 assert_eq!(adapter.read_slice(1), Err(DeserializationError::UnexpectedEOF));
894 assert_eq!(adapter.read_array(), Ok([]));
895 assert_eq!(adapter.read_array::<1>(), Err(DeserializationError::UnexpectedEOF));
896 }
897
898 #[test]
899 fn read_adapter_passthrough() {
900 let mut reader = std::io::repeat(0b101);
901 let mut adapter = ReadAdapter::new(&mut reader);
902 assert!(adapter.has_more_bytes());
903 assert_eq!(adapter.check_eor(8), Ok(()));
904 assert_eq!(adapter.peek_u8(), Ok(0b101));
905 assert_eq!(adapter.read_u8(), Ok(0b101));
906 assert_eq!(adapter.read_slice(0), Ok([].as_slice()));
907 assert_eq!(adapter.read_slice(4), Ok([0b101, 0b101, 0b101, 0b101].as_slice()));
908 assert_eq!(adapter.read_array(), Ok([]));
909 assert_eq!(adapter.read_array(), Ok([0b101, 0b101]));
910 }
911
912 #[test]
913 fn read_adapter_exact() {
914 const VALUE: usize = 2048;
915 let mut reader = Cursor::new(VALUE.to_le_bytes());
916 let mut adapter = ReadAdapter::new(&mut reader);
917 assert_eq!(usize::from_le_bytes(adapter.read_array().unwrap()), VALUE);
918 assert!(!adapter.has_more_bytes());
919 assert_eq!(adapter.peek_u8(), Err(DeserializationError::UnexpectedEOF));
920 assert_eq!(adapter.read_u8(), Err(DeserializationError::UnexpectedEOF));
921 }
922
923 #[test]
924 fn read_adapter_large_array_from_chunked_reader() {
925 let data = (0..897).map(|i| (i % 251) as u8).collect::<Vec<_>>();
926 let expected: [u8; 897] = data.clone().try_into().unwrap();
927 let mut chunked = ChunkedReader::new(data, 128);
928 let mut adapter = ReadAdapter::new(&mut chunked);
929
930 assert_eq!(adapter.read_array::<897>().unwrap(), expected);
931 }
932
933 #[test]
934 fn read_adapter_large_array_after_buffered_prefix() {
935 let data = (0..700).map(|i| (i % 251) as u8).collect::<Vec<_>>();
936 let expected: [u8; 625] = data[17..642].try_into().unwrap();
937 let mut chunked = ChunkedReader::new(data.clone(), 128);
938 let mut adapter = ReadAdapter::new(&mut chunked);
939
940 assert_eq!(adapter.read_slice(17).unwrap(), &data[..17]);
941 assert_eq!(adapter.read_array::<625>().unwrap(), expected);
942 }
943
944 #[test]
945 fn read_adapter_exact_array_resets_empty_local_buffer() {
946 let data = (0..300).map(|i| (i % 251) as u8).collect::<Vec<_>>();
947 let expected: [u8; 111] = data[17..128].try_into().unwrap();
948 let mut chunked = ChunkedReader::new(data.clone(), 128);
949 let mut adapter = ReadAdapter::new(&mut chunked);
950
951 assert_eq!(adapter.read_slice(17).unwrap(), &data[..17]);
952 assert_eq!(adapter.read_array::<111>().unwrap(), expected);
953 assert_eq!(adapter.read_slice(8).unwrap(), &data[128..136]);
954 }
955
956 #[test]
957 fn read_adapter_roundtrip() {
958 const VALUE: usize = 2048;
959
960 // Write VALUE to storage
961 let mut cursor = Cursor::new([0; size_of::<usize>()]);
962 cursor.write_usize(VALUE);
963
964 // Read VALUE from storage
965 cursor.set_position(0);
966 let mut adapter = ReadAdapter::new(&mut cursor);
967
968 assert_eq!(adapter.read_usize(), Ok(VALUE));
969 }
970
971 #[test]
972 fn read_adapter_for_file() {
973 use std::fs::File;
974
975 use crate::ByteWriter;
976
977 let path = std::env::temp_dir().join("read_adapter_for_file.bin");
978
979 // Encode some data to a buffer, then write that buffer to a file
980 {
981 let mut buf = Vec::<u8>::with_capacity(256);
982 buf.write_bytes(b"MAGIC\0");
983 buf.write_bool(true);
984 buf.write_u32(0xbeef);
985 buf.write_usize(0xfeed);
986 buf.write_u16(0x5);
987
988 std::fs::write(&path, &buf).unwrap();
989 }
990
991 // Open the file, and try to decode the encoded items
992 let mut file = File::open(&path).unwrap();
993 let mut reader = ReadAdapter::new(&mut file);
994 assert_eq!(reader.peek_u8().unwrap(), b'M');
995 assert_eq!(reader.read_slice(6).unwrap(), b"MAGIC\0");
996 assert!(reader.read_bool().unwrap());
997 assert_eq!(reader.read_u32().unwrap(), 0xbeef);
998 assert_eq!(reader.read_usize().unwrap(), 0xfeed);
999 assert_eq!(reader.read_u16().unwrap(), 0x5);
1000 assert!(!reader.has_more_bytes(), "expected there to be no more data in the input");
1001 }
1002
1003 #[test]
1004 fn read_adapter_issue_383() {
1005 const STR_BYTES: &[u8] = b"just a string";
1006
1007 use std::fs::File;
1008
1009 use crate::ByteWriter;
1010
1011 let path = std::env::temp_dir().join("issue_383.bin");
1012
1013 // Encode some data to a buffer, then write that buffer to a file
1014 {
1015 let mut buf = vec![0u8; 1024];
1016 unsafe {
1017 buf.set_len(0);
1018 }
1019 buf.write_u128(2 * u64::MAX as u128);
1020 unsafe {
1021 buf.set_len(512);
1022 }
1023 buf.write_bytes(STR_BYTES);
1024 buf.write_u32(0xbeef);
1025
1026 std::fs::write(&path, &buf).unwrap();
1027 }
1028
1029 // Open the file, and try to decode the encoded items
1030 let mut file = File::open(&path).unwrap();
1031 let mut reader = ReadAdapter::new(&mut file);
1032 assert_eq!(reader.read_u128().unwrap(), 2 * u64::MAX as u128);
1033 assert_eq!(reader.buf.len(), 0);
1034 assert_eq!(reader.pos, 0);
1035 // Read to offset 512 (we're 16 bytes into the underlying file, i.e. offset of 496)
1036 reader.read_slice(496).unwrap();
1037 assert_eq!(reader.buf.len(), 496);
1038 assert_eq!(reader.pos, 496);
1039 // The byte string is 13 bytes, followed by 4 bytes containing the trailing u32 value.
1040 // We expect that the underlying reader will buffer the remaining bytes of the file when
1041 // reading STR_BYTES, so the total size of our adapter's buffer should be
1042 // 496 + STR_BYTES.len() + size_of::<u32>();
1043 assert_eq!(reader.read_slice(STR_BYTES.len()).unwrap(), STR_BYTES);
1044 assert_eq!(reader.buf.len(), 496 + STR_BYTES.len() + size_of::<u32>());
1045 // We haven't read the u32 yet
1046 assert_eq!(reader.pos, 509);
1047 assert_eq!(reader.read_u32().unwrap(), 0xbeef);
1048 // Now we have
1049 assert_eq!(reader.buf.len(), 0);
1050 assert_eq!(reader.pos, 0);
1051 assert!(!reader.has_more_bytes(), "expected there to be no more data in the input");
1052 }
1053
1054 #[test]
1055 fn budgeted_reader_basic() {
1056 let data = [1u8, 2, 3, 4, 5, 6, 7, 8];
1057 let inner = SliceReader::new(&data);
1058 let mut reader = BudgetedReader::new(inner, 4);
1059
1060 assert_eq!(reader.remaining(), 4);
1061 assert!(reader.has_more_bytes());
1062
1063 // read 4 bytes (within budget)
1064 assert_eq!(reader.read_u32().unwrap(), 0x04030201);
1065 assert_eq!(reader.remaining(), 0);
1066
1067 // budget exhausted
1068 assert!(!reader.has_more_bytes());
1069 assert!(reader.read_u8().is_err());
1070 }
1071
1072 #[test]
1073 fn budgeted_reader_peek_does_not_consume() {
1074 let data = [42u8];
1075 let inner = SliceReader::new(&data);
1076 let mut reader = BudgetedReader::new(inner, 1);
1077
1078 // peek multiple times, budget unchanged
1079 assert_eq!(reader.peek_u8().unwrap(), 42);
1080 assert_eq!(reader.peek_u8().unwrap(), 42);
1081 assert_eq!(reader.remaining(), 1);
1082
1083 // actual read consumes budget
1084 assert_eq!(reader.read_u8().unwrap(), 42);
1085 assert_eq!(reader.remaining(), 0);
1086 }
1087
1088 #[test]
1089 fn budgeted_reader_check_eor_respects_budget() {
1090 let data = [0u8; 100];
1091 let inner = SliceReader::new(&data);
1092 let reader = BudgetedReader::new(inner, 10);
1093
1094 // within budget
1095 assert!(reader.check_eor(10).is_ok());
1096
1097 // exceeds budget (even though inner has enough bytes)
1098 assert!(reader.check_eor(11).is_err());
1099 }
1100
1101 #[test]
1102 fn budgeted_reader_read_slice() {
1103 let data = [1u8, 2, 3, 4, 5];
1104 let inner = SliceReader::new(&data);
1105 let mut reader = BudgetedReader::new(inner, 3);
1106
1107 // read 3 bytes (exactly budget)
1108 assert_eq!(reader.read_slice(3).unwrap(), &[1, 2, 3]);
1109 assert_eq!(reader.remaining(), 0);
1110
1111 // can't read more
1112 assert!(reader.read_slice(1).is_err());
1113 }
1114
1115 #[test]
1116 fn budgeted_reader_read_array() {
1117 let data = [0xaau8, 0xbb, 0xcc, 0xdd];
1118 let inner = SliceReader::new(&data);
1119 let mut reader = BudgetedReader::new(inner, 2);
1120
1121 // read 2-byte array
1122 assert_eq!(reader.read_array::<2>().unwrap(), [0xaa, 0xbb]);
1123 assert_eq!(reader.remaining(), 0);
1124
1125 // budget exhausted
1126 assert!(reader.read_array::<2>().is_err());
1127 }
1128
1129 #[test]
1130 fn budgeted_reader_zero_budget() {
1131 let data = [1u8];
1132 let inner = SliceReader::new(&data);
1133 let mut reader = BudgetedReader::new(inner, 0);
1134
1135 assert!(!reader.has_more_bytes());
1136 assert!(reader.read_u8().is_err());
1137 // peek still works (doesn't consume budget)
1138 assert_eq!(reader.peek_u8().unwrap(), 1);
1139 }
1140
1141 #[test]
1142 fn budgeted_reader_max_alloc() {
1143 let data = [0u8; 100];
1144 let inner = SliceReader::new(&data);
1145 let reader = BudgetedReader::new(inner, 64);
1146
1147 // 64 bytes budget / 8 bytes per u64 = 8 elements max
1148 assert_eq!(reader.max_alloc(8), 8);
1149
1150 // 64 bytes budget / 1 byte per u8 = 64 elements max
1151 assert_eq!(reader.max_alloc(1), 64);
1152
1153 // 64 bytes budget / 16 bytes per u128 = 4 elements max
1154 assert_eq!(reader.max_alloc(16), 4);
1155
1156 // Budgeted readers reject non-empty ZST collections because they cannot charge budget.
1157 assert_eq!(reader.max_alloc(0), 0);
1158 }
1159
1160 #[test]
1161 fn unbounded_reader_max_alloc_returns_max() {
1162 let data = [0u8; 100];
1163 let reader = SliceReader::new(&data);
1164
1165 assert_eq!(reader.max_alloc(1), usize::MAX);
1166 assert_eq!(reader.max_alloc(8), usize::MAX);
1167 assert_eq!(reader.max_alloc(0), usize::MAX);
1168 }
1169
1170 #[test]
1171 fn budgeted_reader_rejects_non_empty_zst_collections() {
1172 let data = [];
1173 let inner = SliceReader::new(&data);
1174 let mut reader = BudgetedReader::new(inner, 64);
1175
1176 assert!(reader.read_many_iter::<()>(0).is_ok());
1177 assert!(reader.read_many_iter::<()>(1).is_err());
1178 }
1179
1180 #[test]
1181 fn slice_reader_rejects_overflowing_read_lengths() {
1182 let data = [1u8];
1183 let mut reader = SliceReader::new(&data);
1184
1185 assert_eq!(reader.read_u8().unwrap(), 1);
1186 assert_eq!(reader.read_slice(usize::MAX), Err(DeserializationError::UnexpectedEOF));
1187 assert_eq!(reader.check_eor(usize::MAX), Err(DeserializationError::UnexpectedEOF));
1188 }
1189
1190 // ============================================================================================
1191 // The following tests document the threat model and defense layers.
1192 // ============================================================================================
1193
1194 /// SliceReader alone does NOT reject fake length prefixes.
1195 ///
1196 /// A malicious input claiming 1000 elements will be accepted by read_many_iter
1197 /// because SliceReader.max_alloc() returns usize::MAX. The deserialization will
1198 /// eventually fail with UnexpectedEOF, but only after attempting to iterate.
1199 #[test]
1200 fn slice_reader_accepts_fake_length_prefix() {
1201 let mut data = Vec::new();
1202 // Write length = 1000 (vint64 encoding: 0x07D0 << 2 | 0b10 = 0x1F42)
1203 // For simplicity, use the 9-byte form
1204 data.push(0); // 9-byte marker
1205 data.extend_from_slice(&1000u64.to_le_bytes());
1206 // Only 8 bytes of actual u64 data (1 element, not 1000)
1207 data.extend_from_slice(&42u64.to_le_bytes());
1208
1209 let mut reader = SliceReader::new(&data);
1210 let _len = reader.read_usize().unwrap();
1211 let iter_result = reader.read_many_iter::<u64>(1000);
1212
1213 assert!(iter_result.is_ok());
1214
1215 let collect_result: Result<Vec<u64>, _> = iter_result.unwrap().collect();
1216 assert!(collect_result.is_err());
1217 assert!(matches!(collect_result.unwrap_err(), DeserializationError::UnexpectedEOF));
1218 }
1219
1220 /// BudgetedReader rejects fake length prefixes BEFORE iteration begins.
1221 ///
1222 /// With a 64-byte budget, max_alloc(8) = 8, so a claim of 1000 elements
1223 /// is rejected immediately by read_many_iter.
1224 #[test]
1225 fn budgeted_reader_rejects_fake_length_upfront() {
1226 let mut data = Vec::new();
1227 data.push(0); // 9-byte vint64 marker
1228 data.extend_from_slice(&1000u64.to_le_bytes());
1229 data.extend_from_slice(&42u64.to_le_bytes());
1230
1231 let inner = SliceReader::new(&data);
1232 let mut reader = BudgetedReader::new(inner, 64);
1233
1234 let _len = reader.read_usize().unwrap(); // consumes 9 bytes, 55 remaining
1235 // 55 / 8 = 6 elements max
1236 let iter_result = reader.read_many_iter::<u64>(1000);
1237
1238 // Rejected immediately: 1000 > 6
1239 match iter_result {
1240 Err(DeserializationError::InvalidValue(_)) => {}, // expected
1241 other => panic!("expected InvalidValue error, got {:?}", other.map(|_| "Ok")),
1242 }
1243 }
1244
1245 #[test]
1246 fn read_many_iter_advertises_exact_remaining_count() {
1247 let data = [0u8; 8];
1248 let mut reader = SliceReader::new(&data);
1249 let iter = reader.read_many_iter::<u64>(1000).unwrap();
1250
1251 // size_hint is exact: the iterator knows precisely how many items it will yield
1252 // (each call to next() returns Some until remaining hits 0, whether the item is a
1253 // deserialization Ok or Err). This satisfies ExactSizeIterator.
1254 assert_eq!(iter.size_hint(), (1000, Some(1000)));
1255 assert_eq!(iter.len(), 1000);
1256 }
1257
1258 /// Best practice: budget = input length provides both protections.
1259 ///
1260 /// 1. Fake length prefixes are bounded by max_alloc (remaining_bytes / element_size)
1261 /// 2. Total consumption is bounded by the budget
1262 #[test]
1263 fn budget_equals_input_length_is_safe() {
1264 // Valid input: 2 u64s
1265 let original = vec![100u64, 200];
1266 let mut data = Vec::new();
1267 crate::Serializable::write_into(&original, &mut data);
1268
1269 // Budget = exact input size
1270 let result = Vec::<u64>::read_from_bytes_with_budget(&data, data.len());
1271 assert_eq!(result.unwrap(), vec![100, 200]);
1272
1273 // Malicious input claiming 1000 elements (same serialized prefix manipulation)
1274 let mut evil_data = Vec::new();
1275 evil_data.push(0); // 9-byte vint64
1276 evil_data.extend_from_slice(&1000u64.to_le_bytes());
1277 evil_data.extend_from_slice(&42u64.to_le_bytes()); // only 1 actual element
1278
1279 // Budget = input length (17 bytes). After reading length (9 bytes), 8 remain.
1280 // max_alloc(8) = 8/8 = 1, so 1000 > 1 fails.
1281 let result = Vec::<u64>::read_from_bytes_with_budget(&evil_data, evil_data.len());
1282 assert!(result.is_err());
1283 }
1284
1285 // ============================================================================================
1286 // Tests documenting min_serialized_size()-based allocation bounds (defaults to size_of)
1287 // ============================================================================================
1288
1289 /// The max_alloc check uses D::min_serialized_size() to bound memory allocation.
1290 /// By default, min_serialized_size() returns size_of::<D>().
1291 ///
1292 /// For flat collections like Vec<u64>, this works well: we check that
1293 /// budget / min_serialized_size() >= requested_count before allocating.
1294 #[test]
1295 fn min_serialized_size_bounds_flat_collections() {
1296 let mut data = Vec::new();
1297 data.push(0); // 9-byte vint64 marker
1298 data.extend_from_slice(&1000u64.to_le_bytes()); // claim 1000 u64s
1299 data.extend_from_slice(&[0u8; 16]); // only 2 u64s of actual data
1300
1301 let inner = SliceReader::new(&data);
1302 // Budget of 80 bytes: after reading 9-byte length, 71 remain.
1303 // max_alloc(u64::min_serialized_size()) = 71 / 8 = 8 elements max
1304 let mut reader = BudgetedReader::new(inner, 80);
1305
1306 let _len = reader.read_usize().unwrap();
1307 let result = reader.read_many_iter::<u64>(1000);
1308
1309 // Rejected: 1000 > 8
1310 assert!(result.is_err());
1311 }
1312
1313 /// For nested collections like Vec<Vec<u64>>, min_serialized_size() returns 1 (the minimum
1314 /// vint length prefix), not size_of. This is more permissive but accurate: a
1315 /// serialized Vec can be as small as 1 byte (empty vec).
1316 ///
1317 /// The early-abort check uses this minimum, and budget enforcement during actual
1318 /// reads provides the real protection against malicious input.
1319 #[test]
1320 fn min_serialized_size_override_for_nested_collections() {
1321 // Vec<u64>::min_serialized_size() returns 1 (minimum vint prefix), not size_of
1322 assert_eq!(<Vec<u64>>::min_serialized_size(), 1);
1323
1324 let mut data = Vec::new();
1325 data.push(0); // 9-byte vint64 marker
1326 data.extend_from_slice(&100u64.to_le_bytes()); // claim 100 inner Vecs
1327 // Only provide enough data for 1 empty inner Vec
1328 data.push(0b10); // vint64 for 0 (empty inner vec)
1329
1330 let inner = SliceReader::new(&data);
1331 // With min_serialized_size() = 1, we need budget >= 100 to pass the early check.
1332 // After reading 9-byte length, 101 - 9 = 92 remaining, 92 / 1 = 92 < 100.
1333 // So with budget = 110, we get 110 - 9 = 101 remaining, 101 >= 100.
1334 let mut reader = BudgetedReader::new(inner, 110);
1335
1336 let _len = reader.read_usize().unwrap();
1337 let result = reader.read_many_iter::<Vec<u64>>(100);
1338
1339 // The early check passes (100 <= 101)
1340 assert!(result.is_ok());
1341
1342 // But deserialization fails when we try to read 100 inner Vecs with only 1
1343 let collect_result: Result<Vec<Vec<u64>>, _> = result.unwrap().collect();
1344 assert!(collect_result.is_err());
1345 }
1346
1347 /// Demonstrates that min_serialized_size() approach still provides security for nested
1348 /// collections, just with later detection. The budget is enforced during reads.
1349 #[test]
1350 fn nested_collections_still_protected_by_budget() {
1351 // With Vec::min_serialized_size() = 1, the early check is permissive.
1352 // Security comes from budget enforcement during actual reads.
1353 let mut data = Vec::new();
1354 data.push(0); // 9-byte vint64 marker
1355 data.extend_from_slice(&10u64.to_le_bytes()); // claim 10 inner Vecs
1356 // Each inner vec claims 1000 u64s but provides none
1357 for _ in 0..10 {
1358 data.push(0); // 9-byte vint64 marker
1359 data.extend_from_slice(&1000u64.to_le_bytes());
1360 }
1361
1362 let inner = SliceReader::new(&data);
1363 // Small budget: will run out during inner deserialization
1364 let mut reader = BudgetedReader::new(inner, 100);
1365
1366 // Outer length read succeeds (consumes 9 bytes, 91 remaining)
1367 let _len = reader.read_usize().unwrap();
1368
1369 // With Vec::min_serialized_size() = 1, early check passes: 91 / 1 = 91 >= 10
1370 let result = reader.read_many_iter::<Vec<u64>>(10);
1371 assert!(result.is_ok());
1372
1373 // But collecting fails because the inner vecs claim 1000 u64s each,
1374 // exhausting the budget during inner deserialization
1375 let collect_result: Result<Vec<Vec<u64>>, _> = result.unwrap().collect();
1376 assert!(collect_result.is_err());
1377 }
1378
1379 /// Tuples should use sum of element min_serialized_size, not size_of (which includes padding).
1380 ///
1381 /// This test verifies that (u8, u64) has min_serialized_size = 9 (1 + 8) not 16 (in-memory size
1382 /// with 7 bytes of alignment padding).
1383 #[test]
1384 fn tuple_min_serialized_size_excludes_padding() {
1385 // Serialized: 1 byte for u8 + 8 bytes for u64 = 9 bytes
1386 // In-memory: 8 bytes for u8 (with 7 bytes padding) + 8 bytes for u64 = 16 bytes
1387 assert_eq!(<(u8, u64)>::min_serialized_size(), 9);
1388 assert_eq!(size_of::<(u8, u64)>(), 16);
1389
1390 // Verify budget calculation uses 9, not 16
1391 let mut data = Vec::new();
1392 data.push(0); // 9-byte vint64 marker
1393 data.extend_from_slice(&4u64.to_le_bytes()); // claim 4 tuples
1394 // Provide exactly 4 tuples worth of data: 4 * 9 = 36 bytes
1395 for i in 0u8..4 {
1396 data.push(i); // u8
1397 data.extend_from_slice(&(i as u64).to_le_bytes()); // u64
1398 }
1399
1400 let inner = SliceReader::new(&data);
1401 // Budget: 9 (length prefix) + 36 (data) = 45 bytes
1402 let mut reader = BudgetedReader::new(inner, 45);
1403
1404 let _len = reader.read_usize().unwrap();
1405 // With min_serialized_size = 9: remaining = 45 - 9 = 36, max_elements = 36 / 9 = 4
1406 // This should succeed (4 <= 4)
1407 let result = reader.read_many_iter::<(u8, u64)>(4);
1408 assert!(result.is_ok());
1409
1410 // With min_serialized_size = 16 (wrong): max_elements = 36 / 16 = 2
1411 // This would fail (4 > 2)
1412 let collect_result: Result<Vec<(u8, u64)>, _> = result.unwrap().collect();
1413 assert!(collect_result.is_ok());
1414 assert_eq!(collect_result.unwrap().len(), 4);
1415 }
1416}