juicebox_sdk_marshalling/
bytes.rs1extern 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 #[derive(Debug, Deserialize, Eq, PartialEq, Serialize)]
138 struct BytesArray<const N: usize>(#[serde(with = "bytes")] [u8; N]);
139
140 #[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 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}