1use serde::de::{Error, SeqAccess, Visitor};
2use serde::{Deserialize, Deserializer, Serialize, Serializer};
3
4use crate::types::*;
5
6macro_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
30macro_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 #[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#[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 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 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#[cfg(any(
191 all(feature = "protected", any(unix, windows)),
192 all(doc, not(doctest), feature = "std")
193))]
194const MAX_PREALLOCATION: usize = 4096;
195
196#[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 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 fn from_bytes<'de, T: Deserialize<'de>>(bytes: &'de [u8]) -> Result<T, ValueError> {
280 T::deserialize(BytesDeserializer::<ValueError>::new(bytes))
281 }
282
283 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 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 #[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 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 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 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 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}