1use base64::prelude::BASE64_STANDARD;
20use base64::Engine;
21use serde::de::{self, SeqAccess, Visitor};
22use serde::{Deserializer, Serializer};
23use std::fmt;
24
25pub fn serialize<S: Serializer, const N: usize>(
26 bytes: &[u8; N],
27 serializer: S,
28) -> Result<S::Ok, S::Error> {
29 if serializer.is_human_readable() {
30 serializer.serialize_str(&BASE64_STANDARD.encode(bytes))
31 } else {
32 serializer.serialize_bytes(bytes)
33 }
34}
35
36pub fn deserialize<'de, D: Deserializer<'de>, const N: usize>(
37 deserializer: D,
38) -> Result<[u8; N], D::Error> {
39 struct AnyShapeVisitor<const N: usize>;
49
50 impl<'de, const N: usize> Visitor<'de> for AnyShapeVisitor<N> {
51 type Value = [u8; N];
52
53 fn expecting(&self, f: &mut fmt::Formatter) -> fmt::Result {
54 write!(
55 f,
56 "{} bytes (as a byte buffer, sequence of u8, or base64-encoded string)",
57 N
58 )
59 }
60
61 fn visit_bytes<E: de::Error>(self, v: &[u8]) -> Result<Self::Value, E> {
62 v.try_into()
63 .map_err(|_| E::custom(format!("expected {} bytes, got {}", N, v.len())))
64 }
65
66 fn visit_byte_buf<E: de::Error>(self, v: Vec<u8>) -> Result<Self::Value, E> {
67 let len = v.len();
68 v.try_into()
69 .map_err(|_| E::custom(format!("expected {} bytes, got {}", N, len)))
70 }
71
72 fn visit_str<E: de::Error>(self, v: &str) -> Result<Self::Value, E> {
73 let vec = BASE64_STANDARD
74 .decode(v)
75 .map_err(|e| E::custom(format!("expected base64-encoded {} bytes: {}", N, e)))?;
76 self.visit_byte_buf(vec)
77 }
78
79 fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Self::Value, A::Error> {
80 let mut buf = Vec::with_capacity(N);
81 while let Some(b) = seq.next_element::<u8>()? {
82 buf.push(b);
83 }
84 let len = buf.len();
85 buf.try_into()
86 .map_err(|_| de::Error::custom(format!("expected {} bytes, got {}", N, len)))
87 }
88 }
89
90 if deserializer.is_human_readable() {
91 deserializer.deserialize_any(AnyShapeVisitor::<N>)
97 } else {
98 deserializer.deserialize_byte_buf(AnyShapeVisitor::<N>)
102 }
103}
104
105pub mod option {
112 use serde::de::{self, Visitor};
113 use serde::{Deserializer, Serializer};
114 use std::fmt;
115
116 pub fn serialize<S: Serializer, const N: usize>(
117 value: &Option<[u8; N]>,
118 serializer: S,
119 ) -> Result<S::Ok, S::Error> {
120 struct Inner<'a, const N: usize>(&'a [u8; N]);
125 impl<'a, const N: usize> serde::Serialize for Inner<'a, N> {
126 fn serialize<S: Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
127 super::serialize(self.0, s)
128 }
129 }
130 match value {
131 Some(bytes) => serializer.serialize_some(&Inner::<N>(bytes)),
132 None => serializer.serialize_none(),
133 }
134 }
135
136 pub fn deserialize<'de, D: Deserializer<'de>, const N: usize>(
137 deserializer: D,
138 ) -> Result<Option<[u8; N]>, D::Error> {
139 struct OptionVisitor<const N: usize>;
140
141 impl<'de, const N: usize> Visitor<'de> for OptionVisitor<N> {
142 type Value = Option<[u8; N]>;
143
144 fn expecting(&self, f: &mut fmt::Formatter) -> fmt::Result {
145 write!(f, "optional {} bytes", N)
146 }
147
148 fn visit_none<E: de::Error>(self) -> Result<Self::Value, E> {
149 Ok(None)
150 }
151
152 fn visit_unit<E: de::Error>(self) -> Result<Self::Value, E> {
153 Ok(None)
154 }
155
156 fn visit_some<D: Deserializer<'de>>(
157 self,
158 deserializer: D,
159 ) -> Result<Self::Value, D::Error> {
160 super::deserialize::<D, N>(deserializer).map(Some)
161 }
162 }
163
164 deserializer.deserialize_option(OptionVisitor::<N>)
165 }
166}
167
168#[cfg(test)]
169mod tests {
170 use base64::prelude::BASE64_STANDARD;
171 use base64::Engine;
172 use serde::{Deserialize, Serialize};
173
174 #[derive(Serialize, Deserialize, PartialEq, Debug)]
175 struct Wrap32(#[serde(with = "super")] [u8; 32]);
176
177 #[derive(Serialize, Deserialize, PartialEq, Debug)]
178 struct Wrap64(#[serde(with = "super")] [u8; 64]);
179
180 #[derive(Serialize, Deserialize, PartialEq, Debug)]
181 struct Wrap20(#[serde(with = "super")] [u8; 20]);
182
183 #[test]
184 fn json_round_trip_32_bytes_uses_base64_string() {
185 let original = Wrap32([0xab; 32]);
186 let value = serde_json::to_value(&original).expect("serialize");
187 assert_eq!(value, serde_json::json!(BASE64_STANDARD.encode([0xab; 32])));
188 let restored: Wrap32 = serde_json::from_value(value).expect("deserialize");
189 assert_eq!(original, restored);
190 }
191
192 #[test]
193 fn json_round_trip_64_bytes_uses_base64_string() {
194 let original = Wrap64([0xcd; 64]);
195 let value = serde_json::to_value(&original).expect("serialize");
196 assert_eq!(value, serde_json::json!(BASE64_STANDARD.encode([0xcd; 64])));
197 let restored: Wrap64 = serde_json::from_value(value).expect("deserialize");
198 assert_eq!(original, restored);
199 }
200
201 #[test]
202 fn json_round_trip_20_bytes_works_with_const_generic() {
203 let original = Wrap20([0x12; 20]);
204 let value = serde_json::to_value(&original).expect("serialize");
205 assert_eq!(value, serde_json::json!(BASE64_STANDARD.encode([0x12; 20])));
206 let restored: Wrap20 = serde_json::from_value(value).expect("deserialize");
207 assert_eq!(original, restored);
208 }
209
210 #[test]
211 fn rejects_wrong_length_base64() {
212 let result: Result<Wrap32, _> =
213 serde_json::from_value(serde_json::json!(BASE64_STANDARD.encode([0u8; 8])));
214 assert!(result.is_err());
215 }
216
217 #[test]
218 fn binary_round_trip_uses_raw_bytes() {
219 let original = Wrap32([0x55; 32]);
220 let bytes = bincode::serde::encode_to_vec(&original, bincode::config::standard())
221 .expect("bincode encode");
222 let (restored, _): (Wrap32, usize) =
223 bincode::serde::decode_from_slice(&bytes, bincode::config::standard())
224 .expect("bincode decode");
225 assert_eq!(original, restored);
226 }
227
228 #[derive(Serialize, Deserialize, PartialEq, Debug)]
231 struct OptWrap32(#[serde(with = "super::option")] Option<[u8; 32]>);
232
233 #[test]
234 fn option_some_json_round_trip() {
235 let original = OptWrap32(Some([0xab; 32]));
236 let value = serde_json::to_value(&original).expect("serialize");
237 assert_eq!(value, serde_json::json!(BASE64_STANDARD.encode([0xab; 32])));
238 let restored: OptWrap32 = serde_json::from_value(value).expect("deserialize");
239 assert_eq!(original, restored);
240 }
241
242 #[test]
243 fn option_none_json_round_trip() {
244 let original = OptWrap32(None);
245 let value = serde_json::to_value(&original).expect("serialize");
246 assert_eq!(value, serde_json::Value::Null);
247 let restored: OptWrap32 = serde_json::from_value(value).expect("deserialize");
248 assert_eq!(original, restored);
249 }
250
251 #[test]
252 fn option_some_binary_round_trip() {
253 let original = OptWrap32(Some([0x77; 32]));
254 let bytes = bincode::serde::encode_to_vec(&original, bincode::config::standard())
255 .expect("bincode encode");
256 let (restored, _): (OptWrap32, usize) =
257 bincode::serde::decode_from_slice(&bytes, bincode::config::standard())
258 .expect("bincode decode");
259 assert_eq!(original, restored);
260 }
261}