diff --git a/src/context_deserialize.rs b/src/context_deserialize.rs index 870ff9f..a8b3956 100644 --- a/src/context_deserialize.rs +++ b/src/context_deserialize.rs @@ -17,9 +17,10 @@ where } } -impl<'de, C, T> ContextDeserialize<'de, C> for ProgressiveVariableList +impl<'de, C, T, N> ContextDeserialize<'de, C> for ProgressiveVariableList where T: ContextDeserialize<'de, C>, + N: Unsigned, C: Clone, { fn context_deserialize(deserializer: D, context: C) -> Result @@ -27,6 +28,34 @@ where D: Deserializer<'de>, { let vec = Vec::::context_deserialize(deserializer, context)?; - Ok(ProgressiveVariableList::new(vec)) + ProgressiveVariableList::new(vec).map_err(D::Error::custom) + } +} + +#[cfg(test)] +mod test { + use crate::ProgressiveVariableList; + use context_deserialize::ContextDeserialize; + use typenum::U4; + + fn context_deserialize_u64s(json: &str) -> Result, String> { + let mut deserializer = serde_json::Deserializer::from_str(json); + ProgressiveVariableList::context_deserialize(&mut deserializer, ()) + .map_err(|e| e.to_string()) + } + + #[test] + fn context_deserialize_accepts_list_at_limit() { + let list = context_deserialize_u64s("[1,2,3,4]").unwrap(); + assert_eq!(&list[..], &[1, 2, 3, 4]); + } + + #[test] + fn context_deserialize_rejects_list_past_limit() { + let err = context_deserialize_u64s("[1,2,3,4,5]").unwrap_err(); + assert!( + err.contains("Index out of bounds: index 5, length 4"), + "{err}" + ); } } diff --git a/src/progressive_variable_list.rs b/src/progressive_variable_list.rs index 853a2d0..7594a0b 100644 --- a/src/progressive_variable_list.rs +++ b/src/progressive_variable_list.rs @@ -1,22 +1,31 @@ use crate::tree_hash::progressive_vec_tree_hash_root; +use crate::variable_list::MAX_ELEMENTS_TO_PRE_ALLOCATE; +use crate::Error; use serde::Deserialize; use serde_derive::Serialize; use std::any::TypeId; +use std::marker::PhantomData; use std::ops::{Deref, DerefMut, Index, IndexMut}; use std::slice::SliceIndex; use tree_hash::Hash256; +use typenum::{Unsigned, U0}; /// Emulates a SSZ `ProgressiveList` (EIP-7916). /// -/// An ordered, heap-allocated, variable-length, homogeneous collection of `T` with **no** capacity -/// limit. This is the progressive analogue of [`VariableList`](crate::VariableList): the two are -/// identical except that +/// An ordered, heap-allocated, variable-length, homogeneous collection of `T`. This is the +/// progressive analogue of [`VariableList`](crate::VariableList). The two differ in two ways. /// -/// - there is no type-level maximum length (`N`), so construction and `push` are infallible, and -/// - merkleization uses the progressive scheme of EIP-7916 (a right-leaning spine of binary +/// - Merkleization uses the progressive scheme of EIP-7916 (a right-leaning spine of binary /// subtrees whose capacities grow by 4x), so the hash tree root is independent of any limit. +/// - The length limit `N` is optional. A non-zero limit is enforced on construction and mutation. /// -/// Like `VariableList`, it is backed by a Rust `Vec` and serialized identically to a plain list. +/// The type parameter `N` is a [`typenum`] unsigned integer. The default `U0` means no limit, so +/// the list can have any length. SSZ decoding rejects inputs with more than `N` elements for a +/// non-zero `N` before allocating space for them. The limit does not change the SSZ encoding or the +/// hash tree root. You can add or change it on a field without a consensus change. +/// +/// Like `VariableList`, the list is backed by a Rust `Vec` and serialized identically to a plain +/// list. /// /// Known spec divergence: encoding does not enforce the SSZ requirement that the total /// encoding be less than 2^32 bytes. Callers must ensure this limit is respected. @@ -26,49 +35,56 @@ use tree_hash::Hash256; /// /// ``` /// use ssz_types::ProgressiveVariableList; +/// use ssz_types::typenum::U8; /// /// let base: Vec = vec![1, 2, 3, 4]; /// -/// let mut list: ProgressiveVariableList = ProgressiveVariableList::new(base.clone()); +/// // No limit (default). +/// let mut list: ProgressiveVariableList = ProgressiveVariableList::new(base.clone()).unwrap(); /// assert_eq!(&list[..], &[1, 2, 3, 4]); /// -/// // Unlike `VariableList`, `push` cannot fail. -/// list.push(5); +/// // `push` succeeds as long as the optional limit is not exceeded. +/// list.push(5).unwrap(); /// assert_eq!(&list[..], &[1, 2, 3, 4, 5]); +/// +/// // A limited list rejects oversized input. +/// type Bounded = ProgressiveVariableList; +/// assert_eq!(Bounded::max_len(), Some(8)); +/// assert!(Bounded::new(vec![0; 9]).is_err()); /// ``` #[derive(Clone, Serialize)] #[serde(transparent)] -pub struct ProgressiveVariableList { +pub struct ProgressiveVariableList { vec: Vec, + #[serde(skip)] + _phantom: PhantomData, } -impl PartialEq for ProgressiveVariableList { +impl PartialEq for ProgressiveVariableList { fn eq(&self, other: &Self) -> bool { self.vec == other.vec } } -impl Eq for ProgressiveVariableList {} -impl std::hash::Hash for ProgressiveVariableList { +impl Eq for ProgressiveVariableList {} +impl std::hash::Hash for ProgressiveVariableList { fn hash(&self, state: &mut H) { self.vec.hash(state); } } -impl std::fmt::Debug for ProgressiveVariableList { +impl std::fmt::Debug for ProgressiveVariableList { fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { self.vec.fmt(f) } } -impl ProgressiveVariableList { - /// Create a list from a `Vec`. Infallible, as there is no capacity limit. - pub fn new(vec: Vec) -> Self { - Self { vec } - } - +impl ProgressiveVariableList { /// Create an empty list. pub fn empty() -> Self { - Self { vec: vec![] } + Self { + vec: vec![], + _phantom: PhantomData, + } } /// Returns the number of values presently in `self`. @@ -81,11 +97,6 @@ impl ProgressiveVariableList { self.vec.is_empty() } - /// Appends `value` to the back of `self`. Infallible, as there is no capacity limit. - pub fn push(&mut self, value: T) { - self.vec.push(value); - } - /// Returns the contents as a slice. pub fn as_slice(&self) -> &[T] { &self.vec @@ -97,33 +108,80 @@ impl ProgressiveVariableList { } } -impl From> for ProgressiveVariableList { - fn from(vec: Vec) -> Self { +impl ProgressiveVariableList { + /// Create a list from a `Vec`, returning an error if it exceeds the optional limit. + pub fn new(vec: Vec) -> Result { + if let Some(max) = Self::max_len() { + if vec.len() > max { + return Err(Error::OutOfBounds { + i: vec.len(), + len: max, + }); + } + } + Ok(Self { + vec, + _phantom: PhantomData, + }) + } + + /// Appends `value` to the back of `self`, returning an error if it would exceed the optional limit. + pub fn push(&mut self, value: T) -> Result<(), Error> { + if let Some(max) = Self::max_len() { + if self.vec.len() >= max { + return Err(Error::OutOfBounds { + i: self.vec.len() + 1, + len: max, + }); + } + } + self.vec.push(value); + Ok(()) + } + + /// Returns the optional length limit. + /// + /// `None` means no limit (`N = U0`). `Some(n)` rejects any input with more than `n` elements. + pub fn max_len() -> Option { + match N::to_usize() { + 0 => None, + n => Some(n), + } + } + + /// The size hint comes from untrusted input, so the limit and a constant cap the reservation. + fn capacity_to_reserve((lower, upper): (usize, Option)) -> usize { + let cap = Self::max_len().map_or(MAX_ELEMENTS_TO_PRE_ALLOCATE, |max| { + max.min(MAX_ELEMENTS_TO_PRE_ALLOCATE) + }); + upper.unwrap_or(lower).min(cap) + } +} + +impl TryFrom> for ProgressiveVariableList { + type Error = Error; + + fn try_from(vec: Vec) -> Result { Self::new(vec) } } -impl From> for Vec { - fn from(list: ProgressiveVariableList) -> Vec { +impl From> for Vec { + fn from(list: ProgressiveVariableList) -> Vec { list.vec } } -impl Default for ProgressiveVariableList { +impl Default for ProgressiveVariableList { fn default() -> Self { Self { vec: Vec::default(), + _phantom: PhantomData, } } } -impl FromIterator for ProgressiveVariableList { - fn from_iter>(iter: I) -> Self { - Self::new(iter.into_iter().collect()) - } -} - -impl> Index for ProgressiveVariableList { +impl> Index for ProgressiveVariableList { type Output = I::Output; #[inline] @@ -132,14 +190,14 @@ impl> Index for ProgressiveVariableList { } } -impl> IndexMut for ProgressiveVariableList { +impl> IndexMut for ProgressiveVariableList { #[inline] fn index_mut(&mut self, index: I) -> &mut Self::Output { IndexMut::index_mut(&mut self.vec, index) } } -impl Deref for ProgressiveVariableList { +impl Deref for ProgressiveVariableList { type Target = [T]; fn deref(&self) -> &[T] { @@ -147,19 +205,19 @@ impl Deref for ProgressiveVariableList { } } -impl DerefMut for ProgressiveVariableList { +impl DerefMut for ProgressiveVariableList { fn deref_mut(&mut self) -> &mut [T] { &mut self.vec[..] } } -impl AsRef<[T]> for ProgressiveVariableList { +impl AsRef<[T]> for ProgressiveVariableList { fn as_ref(&self) -> &[T] { &self.vec[..] } } -impl<'a, T> IntoIterator for &'a ProgressiveVariableList { +impl<'a, T, N> IntoIterator for &'a ProgressiveVariableList { type Item = &'a T; type IntoIter = std::slice::Iter<'a, T>; @@ -168,7 +226,7 @@ impl<'a, T> IntoIterator for &'a ProgressiveVariableList { } } -impl IntoIterator for ProgressiveVariableList { +impl IntoIterator for ProgressiveVariableList { type Item = T; type IntoIter = std::vec::IntoIter; @@ -177,7 +235,7 @@ impl IntoIterator for ProgressiveVariableList { } } -impl tree_hash::TreeHash for ProgressiveVariableList +impl tree_hash::TreeHash for ProgressiveVariableList where T: tree_hash::TreeHash, { @@ -200,7 +258,7 @@ where } } -impl ssz::Encode for ProgressiveVariableList +impl ssz::Encode for ProgressiveVariableList where T: ssz::Encode, { @@ -221,20 +279,27 @@ where } } -impl ssz::TryFromIter for ProgressiveVariableList { - type Error = std::convert::Infallible; +impl ssz::TryFromIter for ProgressiveVariableList { + type Error = Error; fn try_from_iter(value: I) -> Result where I: IntoIterator, { - Ok(Self::new(value.into_iter().collect())) + let iter = value.into_iter(); + let capacity = Self::capacity_to_reserve(iter.size_hint()); + let mut list = Self::new(Vec::with_capacity(capacity))?; + for item in iter { + list.push(item)?; + } + Ok(list) } } -impl ssz::Decode for ProgressiveVariableList +impl ssz::Decode for ProgressiveVariableList where T: ssz::Decode + 'static, + N: Unsigned, { fn is_ssz_fixed_len() -> bool { false @@ -251,19 +316,40 @@ where return Ok(Self::default()); } + let max_len = Self::max_len(); + if TypeId::of::() == TypeId::of::() { - return Ok(Self::new(crate::u8_bytes_to_vec(bytes))); + if let Some(max) = max_len { + if bytes.len() > max { + return Err(ssz::DecodeError::BytesInvalid(format!( + "ProgressiveVariableList of {} items exceeds maximum of {}", + bytes.len(), + max + ))); + } + } + return Self::new(crate::u8_bytes_to_vec(bytes)) + .map_err(|e| ssz::DecodeError::BytesInvalid(e.to_string())); } if T::is_ssz_fixed_len() { let item_len = T::ssz_fixed_len(); // A zero-length item is a distinct error, matching `VariableList::from_ssz_bytes`. // It also guards the `chunks_exact` below against a zero divisor. - bytes + let num_items = bytes .len() .checked_div(item_len) .ok_or(ssz::DecodeError::ZeroLengthItem)?; + if let Some(max) = max_len { + if num_items > max { + return Err(ssz::DecodeError::BytesInvalid(format!( + "ProgressiveVariableList of {} items exceeds maximum of {}", + num_items, max + ))); + } + } + if !bytes.len().is_multiple_of(item_len) { return Err(ssz::DecodeError::BytesInvalid(format!( "ProgressiveVariableList has {} bytes, not a multiple of item length {}", @@ -272,33 +358,41 @@ where ))); } - bytes - .chunks_exact(item_len) - .map(T::from_ssz_bytes) - .collect::, _>>() - .map(Self::new) + let mut vec = Vec::with_capacity(num_items); + for chunk in bytes.chunks_exact(item_len) { + vec.push(T::from_ssz_bytes(chunk)?); + } + Self::new(vec).map_err(|e| ssz::DecodeError::BytesInvalid(e.to_string())) } else { - ssz::decode_list_of_variable_length_items(bytes, None).map(Self::new) + ssz::decode_list_of_variable_length_items(bytes, max_len) } } } -impl<'de, T> Deserialize<'de> for ProgressiveVariableList +impl<'de, T, N> Deserialize<'de> for ProgressiveVariableList where T: Deserialize<'de>, + N: Unsigned, { fn deserialize(deserializer: D) -> Result where D: serde::Deserializer<'de>, { - Ok(Self::new(Vec::::deserialize(deserializer)?)) + let vec = Vec::::deserialize(deserializer)?; + Self::new(vec).map_err(serde::de::Error::custom) } } #[cfg(feature = "arbitrary")] -impl<'a, T: arbitrary::Arbitrary<'a>> arbitrary::Arbitrary<'a> for ProgressiveVariableList { +impl<'a, T: arbitrary::Arbitrary<'a>, N: Unsigned> arbitrary::Arbitrary<'a> + for ProgressiveVariableList +{ fn arbitrary(u: &mut arbitrary::Unstructured<'a>) -> arbitrary::Result { - Ok(Self::new(>::arbitrary(u)?)) + let mut vec = >::arbitrary(u)?; + if let Some(max) = Self::max_len() { + vec.truncate(max); + } + Self::new(vec).map_err(|_| arbitrary::Error::IncorrectFormat) } fn size_hint(depth: usize) -> (usize, Option) { @@ -309,18 +403,107 @@ impl<'a, T: arbitrary::Arbitrary<'a>> arbitrary::Arbitrary<'a> for ProgressiveVa #[cfg(test)] mod test { use super::*; - use ssz::{Decode, Encode}; + use ssz::{Decode, Encode, TryFromIter}; use tree_hash::TreeHash; + use typenum::{U256, U4}; #[test] - fn new_and_push_infallible() { - let mut list: ProgressiveVariableList = ProgressiveVariableList::new(vec![1, 2, 3]); + fn new_and_push_unbounded() { + let mut list: ProgressiveVariableList = + ProgressiveVariableList::new(vec![1, 2, 3]).unwrap(); assert_eq!(&list[..], &[1, 2, 3]); - list.push(4); + list.push(4).unwrap(); assert_eq!(&list[..], &[1, 2, 3, 4]); assert_eq!(list.len(), 4); } + #[test] + fn limit_bounds_construction() { + type Bounded = ProgressiveVariableList; + for len in [0, 3, 4] { + let values = vec![1; len]; + assert_eq!(Bounded::new(values.clone()).unwrap().as_slice(), values); + assert_eq!( + Bounded::try_from(values.clone()).unwrap().as_slice(), + values + ); + assert_eq!( + Bounded::try_from_iter(values.clone()).unwrap().as_slice(), + values + ); + } + let err = Error::OutOfBounds { i: 5, len: 4 }; + assert_eq!(Bounded::new(vec![1; 5]), Err(err.clone())); + assert_eq!(Bounded::try_from(vec![1; 5]), Err(err.clone())); + let iter = (0..5).chain(std::iter::once_with(|| panic!("iterated past the limit"))); + assert_eq!(Bounded::try_from_iter(iter), Err(err)); + assert_eq!( + ProgressiveVariableList::::try_from_iter(0..5) + .unwrap() + .len(), + 5 + ); + } + + #[test] + fn limit_bounds_push() { + let mut list = ProgressiveVariableList::::new(vec![1, 2, 3]).unwrap(); + list.push(4).unwrap(); + assert_eq!(list.push(5), Err(Error::OutOfBounds { i: 5, len: 4 })); + assert_eq!(&list[..], &[1, 2, 3, 4]); + } + + /// Yields the items of a `Vec` but claims a size hint of `usize::MAX`. + struct LyingIter(std::vec::IntoIter); + + impl Iterator for LyingIter { + type Item = u64; + + fn next(&mut self) -> Option { + self.0.next() + } + + fn size_hint(&self) -> (usize, Option) { + (usize::MAX, Some(usize::MAX)) + } + } + + #[test] + fn try_from_iter_reserves_the_exact_size_hint() { + let list = ProgressiveVariableList::::try_from_iter(0..100).unwrap(); + assert_eq!(list.vec.capacity(), 100); + } + + #[test] + fn try_from_iter_caps_a_lying_size_hint_at_the_limit() { + let iter = LyingIter(vec![1, 2, 3].into_iter()); + let list = ProgressiveVariableList::::try_from_iter(iter).unwrap(); + assert_eq!(&list[..], &[1, 2, 3]); + assert_eq!(list.vec.capacity(), 4); + } + + #[test] + fn try_from_iter_caps_a_lying_size_hint_without_a_limit() { + let iter = LyingIter(vec![1, 2, 3].into_iter()); + let list = ProgressiveVariableList::::try_from_iter(iter).unwrap(); + assert_eq!(&list[..], &[1, 2, 3]); + assert_eq!(list.vec.capacity(), MAX_ELEMENTS_TO_PRE_ALLOCATE); + } + + #[cfg(feature = "arbitrary")] + #[test] + fn limit_bounds_arbitrary() { + use arbitrary::{Arbitrary, Unstructured}; + + let data = [1; 32]; + let unbounded = + ProgressiveVariableList::::arbitrary(&mut Unstructured::new(&data)).unwrap(); + let bounded = + ProgressiveVariableList::::arbitrary(&mut Unstructured::new(&data)).unwrap(); + assert!(unbounded.len() > 4); + assert_eq!(bounded.as_slice(), &unbounded[..4]); + } + fn ssz_round_trip(item: T) { let encoded = &item.as_ssz_bytes(); assert_eq!(item.ssz_bytes_len(), encoded.len()); @@ -329,22 +512,89 @@ mod test { #[test] fn ssz_round_trip_bytes() { - ssz_round_trip::>(ProgressiveVariableList::new(vec![])); - ssz_round_trip::>(ProgressiveVariableList::new(vec![42; 100])); + ssz_round_trip::>( + ProgressiveVariableList::new(vec![]).unwrap(), + ); + ssz_round_trip::>( + ProgressiveVariableList::new(vec![42; 100]).unwrap(), + ); // Serializes identically to a plain byte-list. - let bytes = ProgressiveVariableList::::new(vec![1, 2, 3]); + let bytes = ProgressiveVariableList::::new(vec![1, 2, 3]).unwrap(); assert_eq!(bytes.as_ssz_bytes(), vec![1, 2, 3]); } #[test] fn ssz_round_trip_u64() { - ssz_round_trip::>(ProgressiveVariableList::new(vec![42; 9])); + ssz_round_trip::>( + ProgressiveVariableList::new(vec![42; 9]).unwrap(), + ); + } + + #[test] + fn max_len_reports_optional_limit() { + assert_eq!(ProgressiveVariableList::::max_len(), None); + assert_eq!(ProgressiveVariableList::::max_len(), Some(4)); + } + + #[test] + fn limit_bounds_fixed_len_ssz_decode() { + let ok = ProgressiveVariableList::::new(vec![1, 2, 3, 4]) + .unwrap() + .as_ssz_bytes(); + assert!(ProgressiveVariableList::::from_ssz_bytes(&ok).is_ok()); + + let too_many = ProgressiveVariableList::::new(vec![1, 2, 3, 4, 5]) + .unwrap() + .as_ssz_bytes(); + assert!(ProgressiveVariableList::::from_ssz_bytes(&too_many).is_err()); + + assert!(ProgressiveVariableList::::from_ssz_bytes(&too_many).is_ok()); + } + + #[test] + fn limit_bounds_variable_len_ssz_decode() { + type Inner = ProgressiveVariableList; + let items: Vec = (0..5) + .map(|_| ProgressiveVariableList::new(vec![1]).unwrap()) + .collect(); + let encoded = ProgressiveVariableList::::new(items) + .unwrap() + .as_ssz_bytes(); + + assert!(ProgressiveVariableList::::from_ssz_bytes(&encoded).is_err()); + assert!(ProgressiveVariableList::::from_ssz_bytes(&encoded).is_ok()); + } + + #[test] + fn limit_bounds_byte_list_ssz_decode() { + let encoded = ProgressiveVariableList::::new(vec![0; 5]) + .unwrap() + .as_ssz_bytes(); + assert!(ProgressiveVariableList::::from_ssz_bytes(&encoded).is_err()); + assert!(ProgressiveVariableList::::from_ssz_bytes(&encoded).is_ok()); + } + + #[test] + fn limit_bounds_json_deserialize() { + let json = "[1,2,3,4,5]"; + assert!(serde_json::from_str::>(json).is_err()); + assert!(serde_json::from_str::>(json).is_ok()); + } + + #[test] + fn limit_does_not_change_encoding_or_root() { + let values = vec![9u64, 8, 7]; + let unbounded = ProgressiveVariableList::::new(values.clone()).unwrap(); + let bounded = ProgressiveVariableList::::new(values).unwrap(); + assert_eq!(unbounded.as_ssz_bytes(), bounded.as_ssz_bytes()); + assert_eq!(unbounded.tree_hash_root(), bounded.tree_hash_root()); } #[test] fn serde_is_a_sequence() { // Matches `VariableList`: the default serde representation is a JSON sequence, not hex. - let list: ProgressiveVariableList = ProgressiveVariableList::new(vec![1, 2, 255]); + let list: ProgressiveVariableList = + ProgressiveVariableList::new(vec![1, 2, 255]).unwrap(); let json = serde_json::to_string(&list).unwrap(); assert_eq!(json, "[1,2,255]"); assert_eq!( @@ -357,7 +607,6 @@ mod test { fn tree_hash_byte_list() { use crate::VariableList; use tree_hash::mix_in_length; - use typenum::U256; // Empty list: mix_in_length(ZERO, 0). (Exact progressive-hasher math is covered by the // `tree_hash` crate's own tests and the EF spec vectors.) @@ -368,10 +617,14 @@ mod test { // Deterministic and non-zero for a non-empty list. let bytes: Vec = (0..40).collect(); - let root = ProgressiveVariableList::new(bytes.clone()).tree_hash_root(); + let root = ProgressiveVariableList::::new(bytes.clone()) + .unwrap() + .tree_hash_root(); assert_eq!( root, - ProgressiveVariableList::new(bytes.clone()).tree_hash_root() + ProgressiveVariableList::::new(bytes.clone()) + .unwrap() + .tree_hash_root() ); assert_ne!(root, mix_in_length(&Hash256::ZERO, bytes.len())); diff --git a/src/serde_utils/hex_prog_var_list.rs b/src/serde_utils/hex_prog_var_list.rs index 38c728e..3e64bab 100644 --- a/src/serde_utils/hex_prog_var_list.rs +++ b/src/serde_utils/hex_prog_var_list.rs @@ -2,22 +2,27 @@ //! //! The progressive (EIP-7688) counterpart of [`hex_var_list`](super::hex_var_list). use crate::ProgressiveVariableList; -use serde::{Deserializer, Serializer}; +use serde::{de::Error, Deserializer, Serializer}; use serde_utils::hex::{self, PrefixedHexVisitor}; +use typenum::Unsigned; -pub fn serialize(bytes: &ProgressiveVariableList, serializer: S) -> Result +pub fn serialize( + bytes: &ProgressiveVariableList, + serializer: S, +) -> Result where S: Serializer, { serializer.serialize_str(&hex::encode(&**bytes)) } -pub fn deserialize<'de, D>(deserializer: D) -> Result, D::Error> +pub fn deserialize<'de, D, N>(deserializer: D) -> Result, D::Error> where D: Deserializer<'de>, + N: Unsigned, { let bytes = deserializer.deserialize_str(PrefixedHexVisitor)?; - Ok(ProgressiveVariableList::new(bytes)) + ProgressiveVariableList::new(bytes).map_err(D::Error::custom) } #[cfg(test)] @@ -34,7 +39,7 @@ mod test { #[test] fn round_trip_hex() { let obj = Obj { - bytes: ProgressiveVariableList::new(vec![1, 2, 3, 255]), + bytes: ProgressiveVariableList::new(vec![1, 2, 3, 255]).unwrap(), }; let json = serde_json::to_string(&obj).unwrap(); assert_eq!(json, r#"{"bytes":"0x010203ff"}"#); @@ -50,4 +55,16 @@ mod test { assert_eq!(json, r#"{"bytes":"0x"}"#); assert_eq!(serde_json::from_str::(&json).unwrap(), obj); } + + #[derive(Debug, PartialEq, Deserialize)] + struct Bounded { + #[serde(with = "crate::serde_utils::hex_prog_var_list")] + bytes: ProgressiveVariableList, + } + + #[test] + fn limit_rejects_oversized_hex() { + assert!(serde_json::from_str::(r#"{"bytes":"0x01020304"}"#).is_ok()); + assert!(serde_json::from_str::(r#"{"bytes":"0x0102030405"}"#).is_err()); + } } diff --git a/src/serde_utils/prog_list_of_hex_fixed_vec.rs b/src/serde_utils/prog_list_of_hex_fixed_vec.rs index ab32a6a..172f6e4 100644 --- a/src/serde_utils/prog_list_of_hex_fixed_vec.rs +++ b/src/serde_utils/prog_list_of_hex_fixed_vec.rs @@ -2,15 +2,15 @@ //! //! The progressive (EIP-7688) counterpart of [`list_of_hex_fixed_vec`](super::list_of_hex_fixed_vec). use crate::{FixedVector, ProgressiveVariableList}; -use serde::{ser::SerializeSeq, Deserializer, Serializer}; +use serde::{de::Error, ser::SerializeSeq, Deserializer, Serializer}; use std::marker::PhantomData; use typenum::Unsigned; // The element wrappers are identical to the bounded `list_of_hex_fixed_vec`, so reuse them. pub use super::list_of_hex_fixed_vec::{WrappedListOwned, WrappedListRef}; -pub fn serialize( - list: &ProgressiveVariableList>, +pub fn serialize( + list: &ProgressiveVariableList, N>, serializer: S, ) -> Result where @@ -24,16 +24,26 @@ where seq.end() } -#[derive(Default)] -pub struct Visitor { +pub struct Visitor { _phantom_m: PhantomData, + _phantom_n: PhantomData, } -impl<'a, M> serde::de::Visitor<'a> for Visitor +impl Default for Visitor { + fn default() -> Self { + Self { + _phantom_m: PhantomData, + _phantom_n: PhantomData, + } + } +} + +impl<'a, M, N> serde::de::Visitor<'a> for Visitor where M: Unsigned, + N: Unsigned, { - type Value = ProgressiveVariableList>; + type Value = ProgressiveVariableList, N>; fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result { write!(formatter, "a list of 0x-prefixed hex bytes") @@ -44,21 +54,51 @@ where A: serde::de::SeqAccess<'a>, { let mut list = ProgressiveVariableList::empty(); - while let Some(val) = seq.next_element::>()? { - list.push(val.0); + list.push(val.0).map_err(A::Error::custom)?; } - Ok(list) } } -pub fn deserialize<'de, D, M>( +pub fn deserialize<'de, D, M, N>( deserializer: D, -) -> Result>, D::Error> +) -> Result, N>, D::Error> where D: Deserializer<'de>, M: Unsigned, + N: Unsigned, { deserializer.deserialize_seq(Visitor::default()) } + +#[cfg(test)] +mod test { + use crate::{FixedVector, ProgressiveVariableList}; + use serde_derive::Deserialize; + use typenum::{U1, U2}; + + #[derive(Debug, Deserialize)] + struct Bounded { + #[serde(with = "crate::serde_utils::prog_list_of_hex_fixed_vec")] + vecs: ProgressiveVariableList, U2>, + } + + #[test] + fn accepts_list_at_limit() { + let json = r#"{"vecs":["0x01","0x02"]}"#; + assert_eq!(serde_json::from_str::(json).unwrap().vecs.len(), 2); + } + + #[test] + fn fails_at_first_item_past_limit() { + // The item after the limit is not hex. Only an early check reports the limit. + let json = r#"{"vecs":["0x01","0x02","0x03","not hex"]}"#; + let err = serde_json::from_str::(json).unwrap_err(); + assert!( + err.to_string() + .contains("Index out of bounds: index 3, length 2"), + "{err}" + ); + } +} diff --git a/src/serde_utils/prog_list_of_hex_prog_var_list.rs b/src/serde_utils/prog_list_of_hex_prog_var_list.rs index 5806de5..4f8e19d 100644 --- a/src/serde_utils/prog_list_of_hex_prog_var_list.rs +++ b/src/serde_utils/prog_list_of_hex_prog_var_list.rs @@ -3,22 +3,38 @@ //! //! The progressive (EIP-7688) counterpart of [`list_of_hex_var_list`](super::list_of_hex_var_list). use crate::ProgressiveVariableList; -use serde::{ser::SerializeSeq, Deserialize, Deserializer, Serialize, Serializer}; - -#[derive(Deserialize)] -#[serde(transparent)] -pub struct WrappedListOwned( - #[serde(with = "crate::serde_utils::hex_prog_var_list")] ProgressiveVariableList, -); - -#[derive(Serialize)] -#[serde(transparent)] -pub struct WrappedListRef<'a>( - #[serde(with = "crate::serde_utils::hex_prog_var_list")] &'a ProgressiveVariableList, -); - -pub fn serialize( - list: &ProgressiveVariableList>, +use serde::{de::Error, ser::SerializeSeq, Deserialize, Deserializer, Serialize, Serializer}; +use std::marker::PhantomData; +use typenum::Unsigned; + +/// The inner byte list. `M` is its optional length limit. +pub struct WrappedListOwned(ProgressiveVariableList); + +impl<'de, M> Deserialize<'de> for WrappedListOwned +where + M: Unsigned, +{ + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + Ok(Self(super::hex_prog_var_list::deserialize(deserializer)?)) + } +} + +pub struct WrappedListRef<'a, M>(&'a ProgressiveVariableList); + +impl Serialize for WrappedListRef<'_, M> { + fn serialize(&self, serializer: S) -> Result + where + S: Serializer, + { + super::hex_prog_var_list::serialize(self.0, serializer) + } +} + +pub fn serialize( + list: &ProgressiveVariableList, N>, serializer: S, ) -> Result where @@ -31,10 +47,26 @@ where seq.end() } -struct Visitor; +pub struct Visitor { + _phantom_m: PhantomData, + _phantom_n: PhantomData, +} -impl<'a> serde::de::Visitor<'a> for Visitor { - type Value = ProgressiveVariableList>; +impl Default for Visitor { + fn default() -> Self { + Self { + _phantom_m: PhantomData, + _phantom_n: PhantomData, + } + } +} + +impl<'a, M, N> serde::de::Visitor<'a> for Visitor +where + M: Unsigned, + N: Unsigned, +{ + type Value = ProgressiveVariableList, N>; fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result { write!(formatter, "a list of 0x-prefixed hex strings") @@ -44,27 +76,30 @@ impl<'a> serde::de::Visitor<'a> for Visitor { where A: serde::de::SeqAccess<'a>, { - let mut list = Vec::new(); - while let Some(val) = seq.next_element::()? { - list.push(val.0); + let mut list = ProgressiveVariableList::empty(); + while let Some(val) = seq.next_element::>()? { + list.push(val.0).map_err(A::Error::custom)?; } - Ok(ProgressiveVariableList::new(list)) + Ok(list) } } -pub fn deserialize<'de, D>( +pub fn deserialize<'de, D, M, N>( deserializer: D, -) -> Result>, D::Error> +) -> Result, N>, D::Error> where D: Deserializer<'de>, + M: Unsigned, + N: Unsigned, { - deserializer.deserialize_seq(Visitor) + deserializer.deserialize_seq(Visitor::default()) } #[cfg(test)] mod test { use crate::ProgressiveVariableList; use serde_derive::{Deserialize, Serialize}; + use typenum::U2; #[derive(Debug, PartialEq, Serialize, Deserialize)] struct Obj { @@ -76,9 +111,10 @@ mod test { fn round_trip_hex() { let obj = Obj { lists: ProgressiveVariableList::new(vec![ - ProgressiveVariableList::new(vec![1, 2, 3]), - ProgressiveVariableList::new(vec![255]), - ]), + ProgressiveVariableList::new(vec![1, 2, 3]).unwrap(), + ProgressiveVariableList::new(vec![255]).unwrap(), + ]) + .unwrap(), }; let json = serde_json::to_string(&obj).unwrap(); assert_eq!(json, r#"{"lists":["0x010203","0xff"]}"#); @@ -94,4 +130,31 @@ mod test { assert_eq!(json, r#"{"lists":[]}"#); assert_eq!(serde_json::from_str::(&json).unwrap(), obj); } + + #[derive(Debug, Deserialize)] + struct Bounded { + #[serde(with = "crate::serde_utils::prog_list_of_hex_prog_var_list")] + lists: ProgressiveVariableList, U2>, + } + + #[test] + fn accepts_list_at_limit() { + let json = r#"{"lists":["0x01","0x02"]}"#; + assert_eq!( + serde_json::from_str::(json).unwrap().lists.len(), + 2 + ); + } + + #[test] + fn fails_at_first_item_past_limit() { + // The item after the limit is not hex. Only an early check reports the limit. + let json = r#"{"lists":["0x01","0x02","0x03","not hex"]}"#; + let err = serde_json::from_str::(json).unwrap_err(); + assert!( + err.to_string() + .contains("Index out of bounds: index 3, length 2"), + "{err}" + ); + } } diff --git a/src/variable_list.rs b/src/variable_list.rs index 40b0c99..1bc8b15 100644 --- a/src/variable_list.rs +++ b/src/variable_list.rs @@ -80,7 +80,7 @@ impl std::fmt::Debug for VariableList { /// in memory. This value is set to 128K with the expectation that any list with a large maximum /// length (N) will contain at least a few thousand small values. i.e. we're targeting an /// allocation around the 1MiB to 10MiB mark. -const MAX_ELEMENTS_TO_PRE_ALLOCATE: usize = 128 * (1 << 10); +pub(crate) const MAX_ELEMENTS_TO_PRE_ALLOCATE: usize = 128 * (1 << 10); impl VariableList { /// Returns `Some` if the given `vec` equals the fixed length of `Self`. Otherwise returns