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
6impl<const LENGTH: usize> Serialize for StackByteArray<LENGTH> {
7    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
8    where
9        S: Serializer,
10    {
11        serializer.serialize_bytes(self.as_slice())
12    }
13}
14
15impl<'de, const LENGTH: usize> Deserialize<'de> for StackByteArray<LENGTH> {
16    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
17    where
18        D: Deserializer<'de>,
19    {
20        struct ByteArrayVisitor<const LENGTH: usize>;
21
22        impl<'de, const LENGTH: usize> Visitor<'de> for ByteArrayVisitor<LENGTH> {
23            type Value = StackByteArray<LENGTH>;
24
25            fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
26                write!(formatter, "exactly {LENGTH} bytes")
27            }
28
29            fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
30            where
31                A: SeqAccess<'de>,
32            {
33                let mut arr = StackByteArray::<LENGTH>::new();
34                let mut idx: usize = 0;
35
36                while let Some(elem) = seq.next_element()? {
37                    if idx >= LENGTH {
38                        return Err(Error::invalid_length(idx + 1, &self));
39                    }
40                    arr[idx] = elem;
41                    idx += 1;
42                }
43
44                if idx != LENGTH {
45                    return Err(Error::invalid_length(idx, &self));
46                }
47
48                Ok(arr)
49            }
50
51            fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
52            where
53                E: Error,
54            {
55                if v.len() != LENGTH {
56                    return Err(Error::invalid_length(v.len(), &self));
57                }
58                let mut arr = StackByteArray::<LENGTH>::new();
59                arr.copy_from_slice(v);
60                Ok(arr)
61            }
62        }
63
64        deserializer.deserialize_bytes(ByteArrayVisitor::<LENGTH>)
65    }
66}
67
68#[cfg(any(all(feature = "protected", any(unix, windows)), all(doc, not(doctest))))]
69mod protected {
70    use super::*;
71    use crate::protected::*;
72
73    impl<const LENGTH: usize> Serialize for HeapByteArray<LENGTH> {
74        fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
75        where
76            S: Serializer,
77        {
78            serializer.serialize_bytes(self.as_slice())
79        }
80    }
81
82    impl<const LENGTH: usize> Serialize for Locked<HeapByteArray<LENGTH>> {
83        fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
84        where
85            S: Serializer,
86        {
87            serializer.serialize_bytes(self.as_slice())
88        }
89    }
90
91    impl<'de, const LENGTH: usize> Deserialize<'de> for HeapByteArray<LENGTH> {
92        fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
93        where
94            D: Deserializer<'de>,
95        {
96            struct ByteArrayVisitor<const LENGTH: usize>;
97
98            impl<'de, const LENGTH: usize> Visitor<'de> for ByteArrayVisitor<LENGTH> {
99                type Value = HeapByteArray<LENGTH>;
100
101                fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
102                    write!(formatter, "exactly {LENGTH} bytes")
103                }
104
105                fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
106                where
107                    A: SeqAccess<'de>,
108                {
109                    let mut arr = HeapByteArray::<LENGTH>::default();
110                    let mut idx = 0;
111
112                    while let Some(elem) = seq.next_element()? {
113                        if idx >= LENGTH {
114                            return Err(Error::invalid_length(idx + 1, &self));
115                        }
116                        arr[idx] = elem;
117                        idx += 1;
118                    }
119
120                    if idx != LENGTH {
121                        return Err(Error::invalid_length(idx, &self));
122                    }
123
124                    Ok(arr)
125                }
126
127                fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
128                where
129                    E: Error,
130                {
131                    if v.len() != LENGTH {
132                        return Err(Error::invalid_length(v.len(), &self));
133                    }
134                    HeapByteArray::<LENGTH>::try_from(v).map_err(E::custom)
135                }
136            }
137
138            deserializer.deserialize_bytes(ByteArrayVisitor::<LENGTH>)
139        }
140    }
141
142    impl Serialize for HeapBytes {
143        fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
144        where
145            S: Serializer,
146        {
147            serializer.serialize_bytes(self.as_slice())
148        }
149    }
150
151    impl Serialize for LockedBytes {
152        fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
153        where
154            S: Serializer,
155        {
156            serializer.serialize_bytes(self.as_slice())
157        }
158    }
159
160    impl Serialize for LockedRO<HeapBytes> {
161        fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
162        where
163            S: Serializer,
164        {
165            serializer.serialize_bytes(self.as_slice())
166        }
167    }
168
169    impl<'de> Deserialize<'de> for HeapBytes {
170        fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
171        where
172            D: Deserializer<'de>,
173        {
174            struct BytesVisitor;
175
176            impl<'de> Visitor<'de> for BytesVisitor {
177                type Value = HeapBytes;
178
179                fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
180                    write!(formatter, "bytes")
181                }
182
183                fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
184                where
185                    A: SeqAccess<'de>,
186                {
187                    let mut arr = HeapBytes::default();
188                    let mut idx: usize = 0;
189                    let size_hint = seq.size_hint().unwrap_or(1);
190                    arr.resize(size_hint, 0);
191
192                    while let Some(elem) = seq.next_element()? {
193                        if idx >= arr.len() {
194                            arr.resize(idx + 1, 0);
195                        }
196                        arr[idx] = elem;
197                        idx += 1;
198                    }
199
200                    arr.resize(idx, 0);
201
202                    Ok(arr)
203                }
204
205                fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
206                where
207                    E: Error,
208                {
209                    Ok(HeapBytes::from(v))
210                }
211            }
212
213            deserializer.deserialize_bytes(BytesVisitor)
214        }
215    }
216
217    impl<'de> Deserialize<'de> for LockedBytes {
218        fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
219        where
220            D: Deserializer<'de>,
221        {
222            struct BytesVisitor;
223
224            impl<'de> Visitor<'de> for BytesVisitor {
225                type Value = LockedBytes;
226
227                fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
228                    write!(formatter, "bytes")
229                }
230
231                fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
232                where
233                    A: SeqAccess<'de>,
234                {
235                    let mut arr = HeapBytes::new_locked().map_err(A::Error::custom)?;
236                    let mut idx: usize = 0;
237                    let size_hint = seq.size_hint().unwrap_or(1);
238                    arr.resize(size_hint, 0);
239
240                    while let Some(elem) = seq.next_element()? {
241                        if idx >= arr.len() {
242                            arr.resize(idx + 1, 0);
243                        }
244                        arr[idx] = elem;
245                        idx += 1;
246                    }
247
248                    arr.resize(idx, 0);
249
250                    Ok(arr)
251                }
252
253                fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
254                where
255                    E: Error,
256                {
257                    HeapBytes::from_slice_into_locked(v).map_err(E::custom)
258                }
259            }
260
261            deserializer.deserialize_bytes(BytesVisitor)
262        }
263    }
264
265    impl<'de, const LENGTH: usize> Deserialize<'de> for Locked<HeapByteArray<LENGTH>> {
266        fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
267        where
268            D: Deserializer<'de>,
269        {
270            struct BytesVisitor<const LENGTH: usize>;
271
272            impl<'de, const LENGTH: usize> Visitor<'de> for BytesVisitor<LENGTH> {
273                type Value = Locked<HeapByteArray<LENGTH>>;
274
275                fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
276                    write!(formatter, "exactly {LENGTH} bytes")
277                }
278
279                fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
280                where
281                    A: SeqAccess<'de>,
282                {
283                    let mut arr =
284                        HeapByteArray::<LENGTH>::new_locked().map_err(A::Error::custom)?;
285                    let mut idx: usize = 0;
286                    while let Some(elem) = seq.next_element()? {
287                        if idx >= LENGTH {
288                            return Err(Error::invalid_length(idx + 1, &self));
289                        }
290                        arr[idx] = elem;
291                        idx += 1;
292                    }
293
294                    if idx != LENGTH {
295                        return Err(Error::invalid_length(idx, &self));
296                    }
297
298                    Ok(arr)
299                }
300
301                fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
302                where
303                    E: Error,
304                {
305                    if v.len() != LENGTH {
306                        Err(Error::invalid_length(v.len(), &self))
307                    } else {
308                        HeapByteArray::<LENGTH>::from_slice_into_locked(v).map_err(E::custom)
309                    }
310                }
311            }
312
313            deserializer.deserialize_bytes(BytesVisitor)
314        }
315    }
316}