Skip to main content

juicebox_sdk_marshalling/
bytes.rs

1//! Serde helpers for serializing byte arrays and vectors.
2extern crate alloc;
3use alloc::vec::Vec;
4use core::fmt;
5
6pub fn serialize<Ser, B>(bytes: &B, serializer: Ser) -> Result<Ser::Ok, Ser::Error>
7where
8    Ser: serde::ser::Serializer,
9    B: Bytes,
10{
11    bytes.serialize(serializer)
12}
13
14pub fn deserialize<'de, De, B>(deserializer: De) -> Result<B, De::Error>
15where
16    De: serde::de::Deserializer<'de>,
17    B: Bytes,
18{
19    B::deserialize(deserializer)
20}
21
22pub trait Bytes: Sized {
23    fn serialize<Ser>(&self, serializer: Ser) -> Result<Ser::Ok, Ser::Error>
24    where
25        Ser: serde::ser::Serializer;
26
27    fn deserialize<'de, De>(deserializer: De) -> Result<Self, De::Error>
28    where
29        De: serde::de::Deserializer<'de>;
30}
31
32impl<const N: usize> Bytes for [u8; N] {
33    fn serialize<Ser>(&self, serializer: Ser) -> Result<Ser::Ok, Ser::Error>
34    where
35        Ser: serde::ser::Serializer,
36    {
37        serializer.serialize_bytes(self)
38    }
39
40    fn deserialize<'de, De>(deserializer: De) -> Result<Self, De::Error>
41    where
42        De: serde::de::Deserializer<'de>,
43    {
44        struct Visitor<const N: usize>;
45
46        impl<'de, const N: usize> serde::de::Visitor<'de> for Visitor<N> {
47            type Value = [u8; N];
48
49            fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
50                formatter.write_fmt(format_args!("byte array of length {}", N))
51            }
52
53            fn visit_bytes<E>(self, slice: &[u8]) -> Result<Self::Value, E>
54            where
55                E: serde::de::Error,
56            {
57                Self::Value::try_from(slice)
58                    .map_err(|_| serde::de::Error::invalid_length(slice.len(), &self))
59            }
60
61            fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
62            where
63                A: serde::de::SeqAccess<'de>,
64            {
65                let mut buf: Vec<u8> = Vec::with_capacity(N);
66                while let Some(x) = seq.next_element()? {
67                    buf.push(x);
68                }
69                self.visit_bytes(&buf)
70            }
71        }
72
73        deserializer.deserialize_any(Visitor)
74    }
75}
76
77impl Bytes for Vec<u8> {
78    fn serialize<Ser>(&self, serializer: Ser) -> Result<Ser::Ok, Ser::Error>
79    where
80        Ser: serde::ser::Serializer,
81    {
82        serializer.serialize_bytes(self)
83    }
84
85    fn deserialize<'de, De>(deserializer: De) -> Result<Self, De::Error>
86    where
87        De: serde::de::Deserializer<'de>,
88    {
89        struct Visitor;
90
91        impl<'de> serde::de::Visitor<'de> for Visitor {
92            type Value = Vec<u8>;
93
94            fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
95                formatter.write_str("Vec<u8>")
96            }
97
98            fn visit_bytes<E>(self, slice: &[u8]) -> Result<Self::Value, E>
99            where
100                E: serde::de::Error,
101            {
102                Ok(slice.to_vec())
103            }
104
105            fn visit_byte_buf<E>(self, value: Vec<u8>) -> Result<Self::Value, E>
106            where
107                E: serde::de::Error,
108            {
109                Ok(value)
110            }
111
112            fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
113            where
114                A: serde::de::SeqAccess<'de>,
115            {
116                let mut buf: Vec<u8> = match seq.size_hint() {
117                    Some(hint) => Vec::with_capacity(hint),
118                    None => Vec::new(),
119                };
120                while let Some(x) = seq.next_element()? {
121                    buf.push(x);
122                }
123                Ok(buf)
124            }
125        }
126
127        deserializer.deserialize_any(Visitor)
128    }
129}
130
131#[cfg(test)]
132mod tests {
133    use crate::{bytes, from_slice, to_vec};
134    use serde::{Deserialize, Serialize};
135
136    // A `[u8; N]` wrapper type that uses `serde(with = "bytes")`.
137    #[derive(Debug, Deserialize, Eq, PartialEq, Serialize)]
138    struct BytesArray<const N: usize>(#[serde(with = "bytes")] [u8; N]);
139
140    // A `Vec<u8>` wrapper type that uses `serde(with = "bytes")`.
141    #[derive(Debug, Deserialize, Eq, PartialEq, Serialize)]
142    struct BytesVec(#[serde(with = "bytes")] Vec<u8>);
143
144    fn expected_serialized_bytes(input: &[u8]) -> Vec<u8> {
145        // cbor bytes are tagged with 0x40 & length (a simplification, this
146        // gets more complicated for larger length values)
147        let len = u8::try_from(input.len()).unwrap();
148        if len > 0x17 {
149            unimplemented!("bigger integer encoding");
150        }
151        let bytes_marker = 0x40;
152        let mut buf = vec![bytes_marker + len];
153        buf.extend_from_slice(input);
154        buf
155    }
156
157    #[test]
158    fn test_array_bytes() {
159        let input = BytesArray([0xff; 16]);
160        let serialized = to_vec(&input).unwrap();
161        assert_eq!(serialized, expected_serialized_bytes(&input.0));
162        let output: BytesArray<16> = from_slice(&serialized).unwrap();
163        assert_eq!(input, output);
164    }
165
166    #[test]
167    fn test_array_bytes_error() {
168        let input = BytesArray([0xff; 16]);
169        let serialized = to_vec(&input).unwrap();
170        assert_eq!(serialized, expected_serialized_bytes(&input.0));
171        let result = from_slice::<BytesArray<10>>(&serialized);
172        assert!(format!("{:?}", result.unwrap_err())
173            .contains("invalid length 16, expected byte array of length 10"));
174    }
175
176    #[test]
177    fn test_array_bytes_from_inefficient() {
178        let input = [0xff; 16];
179        let serialized = to_vec(&input).unwrap();
180        assert_eq!(33, serialized.len());
181        let output: BytesArray<16> = from_slice(&serialized).unwrap();
182        assert_eq!(input, output.0);
183    }
184
185    #[test]
186    fn test_vec_bytes() {
187        let input = BytesVec(vec![15; 16]);
188        let serialized = to_vec(&input).unwrap();
189        assert_eq!(serialized, expected_serialized_bytes(&input.0));
190        let output: BytesVec = from_slice(&serialized).unwrap();
191        assert_eq!(input, output);
192    }
193
194    #[test]
195    fn test_vec_bytes_from_inefficient() {
196        let input = vec![0xff; 16];
197        let serialized = to_vec(&input).unwrap();
198        assert_eq!(33, serialized.len());
199        let output: BytesVec = from_slice(&serialized).unwrap();
200        assert_eq!(input, output.0);
201    }
202}