Skip to main content

dryoc/
bytes_serde.rs

1use serde::de::{Error, SeqAccess, Visitor};
2use serde::{Deserialize, Deserializer, Serialize, Serializer};
3
4use crate::types::*;
5
6/// Serializes a byte container with [`Serializer::serialize_bytes`].
7macro_rules! impl_serialize_bytes {
8    ([$($generics:tt)*] $ty:ty) => {
9        impl<$($generics)*> Serialize for $ty {
10            fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
11            where
12                S: Serializer,
13            {
14                serializer.serialize_bytes(self.as_slice())
15            }
16        }
17    };
18    ($ty:ty) => {
19        impl Serialize for $ty {
20            fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
21            where
22                S: Serializer,
23            {
24                serializer.serialize_bytes(self.as_slice())
25            }
26        }
27    };
28}
29
30/// Implements [`Deserialize`] for a fixed-size byte container, accepting a
31/// byte string or a sequence of exactly `LENGTH` bytes.
32///
33/// * `$ty`: the container type, mentioning `LENGTH`.
34/// * `$new`: builds an empty `$ty` in `visit_seq`; may fail with `A::Error` for
35///   locked allocation.
36/// * `$from_slice`: converts the length-checked `v: &[u8]` into `$ty` in
37///   `visit_bytes`; may fail with `E`.
38macro_rules! impl_deserialize_fixed {
39    ($ty:ty, $new:expr, $from_slice:expr) => {
40        impl<'de, const LENGTH: usize> Deserialize<'de> for $ty {
41            fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
42            where
43                D: Deserializer<'de>,
44            {
45                struct ByteArrayVisitor<const LENGTH: usize>;
46
47                impl<'de, const LENGTH: usize> Visitor<'de> for ByteArrayVisitor<LENGTH> {
48                    type Value = $ty;
49
50                    fn expecting(&self, formatter: &mut core::fmt::Formatter) -> core::fmt::Result {
51                        write!(formatter, "exactly {LENGTH} bytes")
52                    }
53
54                    fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
55                    where
56                        A: SeqAccess<'de>,
57                    {
58                        let mut arr = $new;
59                        let mut idx: usize = 0;
60
61                        while let Some(elem) = seq.next_element()? {
62                            if idx >= LENGTH {
63                                return Err(Error::invalid_length(idx + 1, &self));
64                            }
65                            arr[idx] = elem;
66                            idx += 1;
67                        }
68
69                        if idx != LENGTH {
70                            return Err(Error::invalid_length(idx, &self));
71                        }
72
73                        Ok(arr)
74                    }
75
76                    fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
77                    where
78                        E: Error,
79                    {
80                        if v.len() != LENGTH {
81                            return Err(Error::invalid_length(v.len(), &self));
82                        }
83                        $from_slice(v)
84                    }
85
86                    /// Wipes the owned buffer that the deserializer handed
87                    /// over, which would otherwise be freed with its bytes.
88                    #[cfg(feature = "alloc")]
89                    fn visit_byte_buf<E>(self, v: alloc::vec::Vec<u8>) -> Result<Self::Value, E>
90                    where
91                        E: Error,
92                    {
93                        let v = zeroize::Zeroizing::new(v);
94                        self.visit_bytes(&v)
95                    }
96                }
97
98                deserializer.deserialize_bytes(ByteArrayVisitor::<LENGTH>)
99            }
100        }
101    };
102}
103
104/// Implements [`Deserialize`] for a variable-length byte container, accepting
105/// a byte string or a sequence of bytes.
106///
107/// * `$ty`: the container type, with a fallible `try_resize(new_len, value)`.
108/// * `$new`: builds an empty `$ty` as a `Result<$ty, crate::error::Error>`.
109///
110/// Allocation and locking failures are returned as deserialization errors.
111// Only the `protected` module below uses this macro.
112#[cfg(any(
113    all(feature = "protected", any(unix, windows)),
114    all(doc, not(doctest), feature = "std")
115))]
116macro_rules! impl_deserialize_bytes {
117    ($ty:ty, $new:expr) => {
118        impl<'de> Deserialize<'de> for $ty {
119            fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
120            where
121                D: Deserializer<'de>,
122            {
123                struct BytesVisitor;
124
125                impl<'de> Visitor<'de> for BytesVisitor {
126                    type Value = $ty;
127
128                    fn expecting(&self, formatter: &mut core::fmt::Formatter) -> core::fmt::Result {
129                        write!(formatter, "bytes")
130                    }
131
132                    fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
133                    where
134                        A: SeqAccess<'de>,
135                    {
136                        let mut arr = $new.map_err(A::Error::custom)?;
137                        // The size hint comes from the input, so it only sizes
138                        // a bounded first allocation; the buffer then doubles
139                        // as elements arrive and is trimmed once at the end.
140                        let initial = seq.size_hint().unwrap_or(0).min(MAX_PREALLOCATION);
141                        arr.try_resize(initial, 0).map_err(A::Error::custom)?;
142                        let mut len: usize = 0;
143
144                        while let Some(elem) = seq.next_element()? {
145                            if len == arr.len() {
146                                let grown = len
147                                    .checked_mul(2)
148                                    .ok_or_else(|| A::Error::custom("byte sequence is too long"))?
149                                    .max(MIN_GROWTH);
150                                arr.try_resize(grown, 0).map_err(A::Error::custom)?;
151                            }
152                            arr[len] = elem;
153                            len += 1;
154                        }
155
156                        arr.try_resize(len, 0).map_err(A::Error::custom)?;
157
158                        Ok(arr)
159                    }
160
161                    fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
162                    where
163                        E: Error,
164                    {
165                        let mut arr = $new.map_err(E::custom)?;
166                        arr.try_resize(v.len(), 0).map_err(E::custom)?;
167                        arr.copy_from_slice(v);
168                        Ok(arr)
169                    }
170
171                    /// Wipes the owned buffer that the deserializer handed
172                    /// over, which would otherwise be freed with its bytes.
173                    fn visit_byte_buf<E>(self, v: alloc::vec::Vec<u8>) -> Result<Self::Value, E>
174                    where
175                        E: Error,
176                    {
177                        let v = zeroize::Zeroizing::new(v);
178                        self.visit_bytes(&v)
179                    }
180                }
181
182                deserializer.deserialize_bytes(BytesVisitor)
183            }
184        }
185    };
186}
187
188/// The largest allocation a sequence's (untrusted) size hint can request
189/// before any of its elements have been read.
190#[cfg(any(
191    all(feature = "protected", any(unix, windows)),
192    all(doc, not(doctest), feature = "std")
193))]
194const MAX_PREALLOCATION: usize = 4096;
195
196/// The smallest capacity a growing sequence buffer is resized to.
197#[cfg(any(
198    all(feature = "protected", any(unix, windows)),
199    all(doc, not(doctest), feature = "std")
200))]
201const MIN_GROWTH: usize = 64;
202
203impl_serialize_bytes!([const LENGTH: usize] StackByteArray<LENGTH>);
204
205impl_deserialize_fixed!(
206    StackByteArray<LENGTH>,
207    StackByteArray::<LENGTH>::default(),
208    |v| {
209        let mut arr = StackByteArray::<LENGTH>::default();
210        arr.copy_from_slice(v);
211        Ok(arr)
212    }
213);
214
215#[cfg(any(
216    all(feature = "protected", any(unix, windows)),
217    all(doc, not(doctest), feature = "std")
218))]
219mod protected {
220    use super::*;
221    use crate::protected::*;
222
223    impl_serialize_bytes!([const LENGTH: usize] HeapByteArray<LENGTH>);
224
225    impl_serialize_bytes!([const LENGTH: usize] Locked<HeapByteArray<LENGTH>>);
226
227    impl_deserialize_fixed!(
228        HeapByteArray<LENGTH>,
229        HeapByteArray::<LENGTH>::default(),
230        |v| HeapByteArray::<LENGTH>::try_from(v).map_err(E::custom)
231    );
232
233    impl_serialize_bytes!(HeapBytes);
234
235    impl_serialize_bytes!(LockedBytes);
236
237    impl_serialize_bytes!(LockedRO<HeapBytes>);
238
239    impl_deserialize_bytes!(
240        HeapBytes,
241        Ok::<_, crate::error::Error>(HeapBytes::default())
242    );
243
244    impl_deserialize_bytes!(LockedBytes, HeapBytes::new_locked());
245
246    impl_deserialize_fixed!(
247        Locked<HeapByteArray<LENGTH>>,
248        HeapByteArray::<LENGTH>::new_locked().map_err(A::Error::custom)?,
249        |v| HeapByteArray::<LENGTH>::from_slice_into_locked(v).map_err(E::custom)
250    );
251}
252
253#[cfg(test)]
254mod tests {
255    use serde::de::value::{BytesDeserializer, Error as ValueError, SeqDeserializer};
256
257    use super::*;
258
259    /// Iterator whose `size_hint` is wrong, so `visit_seq` cannot rely on it
260    /// for the final length.
261    struct LyingHint<I> {
262        iter: I,
263        hint: usize,
264    }
265
266    impl<I: Iterator> Iterator for LyingHint<I> {
267        type Item = I::Item;
268
269        fn next(&mut self) -> Option<I::Item> {
270            self.iter.next()
271        }
272
273        fn size_hint(&self) -> (usize, Option<usize>) {
274            (self.hint, Some(self.hint))
275        }
276    }
277
278    /// Drives the `visit_bytes` path with a byte string.
279    fn from_bytes<'de, T: Deserialize<'de>>(bytes: &'de [u8]) -> Result<T, ValueError> {
280        T::deserialize(BytesDeserializer::<ValueError>::new(bytes))
281    }
282
283    /// Drives the `visit_seq` path with a sequence claiming `hint` elements.
284    fn from_seq<T: for<'de> Deserialize<'de>>(bytes: &[u8], hint: usize) -> Result<T, ValueError> {
285        let iter = LyingHint {
286            iter: bytes.iter().copied(),
287            hint,
288        };
289        T::deserialize(SeqDeserializer::<_, ValueError>::new(iter))
290    }
291
292    /// A fixed-size container accepts exactly `LENGTH` bytes from either
293    /// input form and rejects one fewer or one more.
294    fn check_fixed<T: for<'de> Deserialize<'de> + Bytes>() {
295        let data = [7u8, 8, 9];
296
297        assert_eq!(
298            from_bytes::<T>(&data).expect("exact bytes").as_slice(),
299            &data
300        );
301        assert!(from_bytes::<T>(&data[..2]).is_err());
302        assert!(from_bytes::<T>(&[7, 8, 9, 10]).is_err());
303        assert!(from_bytes::<T>(&[]).is_err());
304
305        for hint in [0, 3, 100] {
306            assert_eq!(
307                from_seq::<T>(&data, hint).expect("exact seq").as_slice(),
308                &data,
309                "hint {hint}"
310            );
311            assert!(from_seq::<T>(&data[..2], hint).is_err(), "hint {hint}");
312            assert!(from_seq::<T>(&[7, 8, 9, 10], hint).is_err(), "hint {hint}");
313        }
314    }
315
316    /// A variable-size container accepts any length from either input form,
317    /// regardless of what the sequence claims about its length. A claimed
318    /// length of `usize::MAX` must not be allocated up front, and sequences
319    /// longer than the first allocation grow to fit.
320    #[cfg(all(feature = "protected", any(unix, windows)))]
321    fn check_variable<T: for<'de> Deserialize<'de> + Bytes>() {
322        for len in [0usize, 1, 5, 17, 200] {
323            let data: alloc::vec::Vec<u8> = (1..=len as u8).collect();
324            assert_eq!(from_bytes::<T>(&data).expect("bytes").as_slice(), &data);
325            for hint in [0, 1, len, 100, usize::MAX] {
326                assert_eq!(
327                    from_seq::<T>(&data, hint).expect("seq").as_slice(),
328                    &data,
329                    "len {len} hint {hint}"
330                );
331            }
332        }
333    }
334
335    #[test]
336    fn stack_byte_array_deserializes_only_exact_length() {
337        check_fixed::<StackByteArray<3>>();
338    }
339
340    #[test]
341    fn stack_byte_array_json_uses_byte_array_form() {
342        let array = StackByteArray::from([1u8, 2, 3]);
343        let json = serde_json::to_string(&array).expect("serialize");
344        assert_eq!(json, "[1,2,3]");
345
346        let decoded: StackByteArray<3> = serde_json::from_str(&json).expect("deserialize");
347        assert_eq!(decoded, array);
348
349        // A JSON string reaches the visitor as a byte string.
350        let from_string: StackByteArray<3> = serde_json::from_str("\"abc\"").expect("string");
351        assert_eq!(from_string.as_slice(), b"abc");
352        assert!(serde_json::from_str::<StackByteArray<3>>("\"ab\"").is_err());
353        assert!(serde_json::from_str::<StackByteArray<3>>("[1,2]").is_err());
354        assert!(serde_json::from_str::<StackByteArray<3>>("[1,2,3,4]").is_err());
355        assert!(serde_json::from_str::<StackByteArray<3>>("null").is_err());
356    }
357
358    #[cfg(all(feature = "protected", any(unix, windows)))]
359    mod protected {
360        use super::*;
361        use crate::protected::test_util::can_lock_pages;
362        use crate::protected::*;
363
364        #[test]
365        fn fixed_protected_containers_deserialize_only_exact_length() {
366            check_fixed::<HeapByteArray<3>>();
367            if can_lock_pages(1) {
368                check_fixed::<Locked<HeapByteArray<3>>>();
369            }
370        }
371
372        #[test]
373        fn variable_protected_containers_deserialize_any_length() {
374            check_variable::<HeapBytes>();
375            // A locked sequence grows by locked `resize`, which holds the old
376            // and the new page at once.
377            if can_lock_pages(2) {
378                check_variable::<LockedBytes>();
379            }
380        }
381
382        #[test]
383        fn locked_deserialization_yields_locked_values() {
384            if !can_lock_pages(1) {
385                return;
386            }
387            let locked: LockedBytes = from_bytes(&[1, 2, 3]).expect("locked bytes");
388            let unlocked = locked.munlock().expect("munlock");
389            assert_eq!(unlocked.as_slice(), &[1, 2, 3]);
390
391            let locked: Locked<HeapByteArray<3>> = from_seq(&[4, 5, 6], 0).expect("locked array");
392            let unlocked = locked.munlock().expect("munlock");
393            assert_eq!(unlocked.as_slice(), &[4, 5, 6]);
394        }
395
396        #[test]
397        fn protected_containers_serialize_as_their_bytes() {
398            let data = [1u8, 2, 3];
399            let expected = serde_json::to_string(&data).expect("serialize array");
400
401            let heap = HeapBytes::from(&data[..]);
402            assert_eq!(serde_json::to_string(&heap).expect("heap"), expected);
403
404            let array = HeapByteArray::<3>::from(&data);
405            assert_eq!(serde_json::to_string(&array).expect("array"), expected);
406
407            assert_eq!(
408                serde_json::to_string(&HeapBytes::default()).expect("empty"),
409                "[]"
410            );
411
412            // Each locked form is released before the next is created, so
413            // one lockable page suffices.
414            if !can_lock_pages(1) {
415                return;
416            }
417            {
418                let locked = HeapBytes::from_slice_into_locked(&data).expect("locked");
419                assert_eq!(serde_json::to_string(&locked).expect("locked"), expected);
420            }
421            {
422                // `LockedRO<HeapBytes>` is serialize-only: there is no
423                // `Deserialize` impl for it, so it must serialize exactly like
424                // the unlocked, read-write form it was created from.
425                let readonly = HeapBytes::from_slice_into_readonly_locked(&data).expect("readonly");
426                assert_eq!(
427                    serde_json::to_string(&readonly).expect("readonly"),
428                    expected
429                );
430            }
431            let locked_array = HeapByteArray::<3>::from_slice_into_locked(&data).expect("locked");
432            assert_eq!(
433                serde_json::to_string(&locked_array).expect("locked array"),
434                expected
435            );
436        }
437    }
438}