1use std::{collections::hash_map::Entry, fmt};
2
3use serde::de::{self, Deserialize, Deserializer, MapAccess, Unexpected, Visitor};
4use serde::ser::{self, Serialize, SerializeMap, Serializer};
5
6use super::{HeaderMap, HeaderName, HeaderValue, Value};
7
8impl Serialize for HeaderMap {
9 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
10 where
11 S: Serializer,
12 {
13 let mut map = serializer.serialize_map(Some(self.len()))?;
14 for (name, value) in &self.inner {
15 map.serialize_entry(name.as_str(), value)?;
16 }
17 map.end()
18 }
19}
20
21impl<'de> Deserialize<'de> for HeaderMap {
22 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
23 where
24 D: Deserializer<'de>,
25 {
26 deserializer.deserialize_map(HeaderMapVisitor)
27 }
28}
29
30struct HeaderMapVisitor;
31
32impl<'de> Visitor<'de> for HeaderMapVisitor {
33 type Value = HeaderMap;
34
35 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
36 formatter.write_str("a header map")
37 }
38
39 fn visit_map<M>(self, mut map: M) -> Result<Self::Value, M::Error>
40 where
41 M: MapAccess<'de>,
42 {
43 let mut headers = HeaderMap::with_capacity(map.size_hint().unwrap_or(0));
44 while let Some((NameKey(name), value)) = map.next_entry::<NameKey, Value>()? {
45 match headers.inner.entry(name) {
47 Entry::Occupied(mut entry) => entry.get_mut().extend(value),
48 Entry::Vacant(entry) => {
49 entry.insert(value);
50 }
51 }
52 }
53 Ok(headers)
54 }
55}
56
57#[derive(Debug)]
59struct NameKey(HeaderName);
60
61impl<'de> Deserialize<'de> for NameKey {
62 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
63 where
64 D: Deserializer<'de>,
65 {
66 deserializer.deserialize_str(NameKeyVisitor)
67 }
68}
69
70struct NameKeyVisitor;
71
72impl Visitor<'_> for NameKeyVisitor {
73 type Value = NameKey;
74
75 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
76 formatter.write_str("a header name")
77 }
78
79 fn visit_str<E>(self, v: &str) -> Result<Self::Value, E>
80 where
81 E: de::Error,
82 {
83 HeaderName::from_bytes(v.as_bytes())
84 .map(NameKey)
85 .map_err(|_| de::Error::invalid_value(Unexpected::Str(v), &"a valid header name"))
86 }
87
88 fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
89 where
90 E: de::Error,
91 {
92 HeaderName::from_bytes(v)
93 .map(NameKey)
94 .map_err(|_| de::Error::invalid_value(Unexpected::Bytes(v), &"a valid header name"))
95 }
96}
97
98impl Serialize for Value {
99 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
100 where
101 S: Serializer,
102 {
103 match self {
104 Value::One(val) if serializer.is_human_readable() => val.serialize(serializer),
105 Value::One(val) => [val].as_slice().serialize(serializer),
107 Value::Multi(vec) => vec.serialize(serializer),
108 }
109 }
110}
111
112impl<'de> Deserialize<'de> for Value {
113 fn deserialize<D>(deserializer: D) -> Result<Value, D::Error>
114 where
115 D: Deserializer<'de>,
116 {
117 if deserializer.is_human_readable() {
118 return deserializer.deserialize_any(ValueVisitor);
119 }
120 deserializer.deserialize_seq(ValueVisitor)
121 }
122}
123
124struct ValueVisitor;
125
126impl<'de> Visitor<'de> for ValueVisitor {
127 type Value = Value;
128
129 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
130 formatter.write_str("a single header value or sequence of values")
131 }
132
133 fn visit_str<E>(self, v: &str) -> Result<Self::Value, E>
134 where
135 E: de::Error,
136 {
137 Ok(Value::One(HeaderValueVisitor.visit_str(v)?))
138 }
139
140 fn visit_string<E>(self, v: String) -> Result<Self::Value, E>
141 where
142 E: de::Error,
143 {
144 Ok(Value::One(HeaderValueVisitor.visit_string(v)?))
145 }
146
147 fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
148 where
149 E: de::Error,
150 {
151 Ok(Value::One(HeaderValueVisitor.visit_bytes(v)?))
152 }
153
154 fn visit_byte_buf<E>(self, v: Vec<u8>) -> Result<Self::Value, E>
155 where
156 E: de::Error,
157 {
158 Ok(Value::One(HeaderValueVisitor.visit_byte_buf(v)?))
159 }
160
161 fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
162 where
163 A: de::SeqAccess<'de>,
164 {
165 let mut value: Option<Value> = None;
166 while let Some(next_val) = seq.next_element()? {
167 match value.as_mut() {
168 Some(value) => value.append(next_val),
169 None => value = Some(Value::One(next_val)),
170 }
171 }
172 value.ok_or_else(|| de::Error::invalid_length(0, &"non-empty value"))
173 }
174}
175
176impl Serialize for HeaderValue {
177 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
178 where
179 S: Serializer,
180 {
181 if serializer.is_human_readable() {
182 return serializer.serialize_str(
183 self.to_str()
184 .map_err(|err| ser::Error::custom(err.to_string()))?,
185 );
186 }
187 serializer.serialize_bytes(self.as_bytes())
188 }
189}
190
191impl<'de> Deserialize<'de> for HeaderValue {
192 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
193 where
194 D: Deserializer<'de>,
195 {
196 if deserializer.is_human_readable() {
197 return deserializer.deserialize_string(HeaderValueVisitor);
198 }
199 deserializer.deserialize_byte_buf(HeaderValueVisitor)
200 }
201}
202
203struct HeaderValueVisitor;
204
205impl Visitor<'_> for HeaderValueVisitor {
206 type Value = HeaderValue;
207
208 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
209 formatter.write_str("a header value")
210 }
211
212 fn visit_str<E>(self, v: &str) -> Result<Self::Value, E>
213 where
214 E: de::Error,
215 {
216 HeaderValue::from_str(v)
217 .map_err(|_| de::Error::invalid_value(Unexpected::Str(v), &"a valid header value"))
218 }
219
220 fn visit_string<E>(self, v: String) -> Result<Self::Value, E>
221 where
222 E: de::Error,
223 {
224 HeaderValue::from_shared(v).map_err(|err| de::Error::custom(err.to_string()))
225 }
226
227 fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
228 where
229 E: de::Error,
230 {
231 HeaderValue::from_bytes(v)
232 .map_err(|_| de::Error::invalid_value(Unexpected::Bytes(v), &"a valid header value"))
233 }
234
235 fn visit_byte_buf<E>(self, v: Vec<u8>) -> Result<Self::Value, E>
236 where
237 E: de::Error,
238 {
239 HeaderValue::from_shared(v).map_err(|err| de::Error::custom(err.to_string()))
240 }
241}
242
243#[cfg(test)]
244mod tests {
245 use crate::header::*;
246
247 #[test]
248 fn test_serde_json() {
249 let mut map = HeaderMap::new();
250 map.insert(USER_AGENT, HeaderValue::from_static("hello"));
251 map.append(USER_AGENT, HeaderValue::from_static("world"));
252 assert_eq!(
253 serde_json::to_string(&map).unwrap(),
254 r#"{"user-agent":["hello","world"]}"#
255 );
256
257 map.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
259 map.insert(CONTENT_LENGTH, 0.into());
260 let map_json = serde_json::to_string(&map).unwrap();
261 let map2 = serde_json::from_str::<HeaderMap>(&map_json).unwrap();
262 assert_eq!(map, map2);
263
264 let map_uc = serde_json::from_str::<HeaderMap>(r#"{"X-Foo":"BAR"}"#).unwrap();
266 assert_eq!(map_uc.get("x-foo").unwrap(), "BAR");
267 assert_eq!(
268 serde_json::to_string(&map_uc).unwrap(),
269 r#"{"x-foo":"BAR"}"#
270 );
271
272 let map_empty = serde_json::from_str::<HeaderMap>(r#"{"user-agent":[]}"#);
274 assert!(map_empty.is_err());
275 assert!(
276 map_empty
277 .unwrap_err()
278 .to_string()
279 .contains("invalid length 0, expected non-empty value")
280 );
281 }
282
283 #[test]
284 fn test_serde_bincode() {
285 let mut map = HeaderMap::new();
286 map.insert(USER_AGENT, HeaderValue::from_static("hello"));
287 map.append(USER_AGENT, HeaderValue::from_static("world"));
288 map.insert(HeaderName::from_static("x-foo"), "bar".parse().unwrap());
289 let map_bin = bincode::serialize(&map).unwrap();
290 let map2 = bincode::deserialize::<HeaderMap>(&map_bin).unwrap();
291 assert_eq!(map, map2);
292 }
293
294 #[test]
295 fn test_serde_owned_keys() {
296 let v = serde_json::json!({"x-foo": "bar"});
298 let map = serde_json::from_value::<HeaderMap>(v).unwrap();
299 assert_eq!(map.get("x-foo").unwrap(), "bar");
300
301 let map = serde_json::from_reader::<_, HeaderMap>(&br#"{"x-foo":"bar"}"#[..]).unwrap();
302 assert_eq!(map.get("x-foo").unwrap(), "bar");
303
304 let map = serde_json::from_str::<HeaderMap>(r#"{"x-f\u006fo":"bar"}"#).unwrap();
305 assert_eq!(map.get("x-foo").unwrap(), "bar");
306
307 let map = serde_json::from_str::<HeaderMap>(r#"{"x-foo":"a","X-Foo":["b","c"]}"#).unwrap();
309 assert_eq!(map.len(), 1);
310 assert_eq!(map.get_all("x-foo").collect::<Vec<_>>(), ["a", "b", "c"]);
311
312 let err = serde_json::from_str::<HeaderMap>(r#"{"x foo":"a"}"#).unwrap_err();
313 assert!(
314 err.to_string().contains("expected a valid header name"),
315 "{err}"
316 );
317 let err = serde_json::from_value::<HeaderMap>(serde_json::json!([1])).unwrap_err();
318 assert!(err.to_string().contains("expected a header map"), "{err}");
319 }
320
321 #[test]
322 fn test_serde_visitors() {
323 use serde::de::{Visitor, value::Error as DeError};
324
325 use super::{HeaderValueVisitor, NameKey, NameKeyVisitor, Value, ValueVisitor};
326
327 let NameKey(name) = NameKeyVisitor.visit_bytes::<DeError>(b"x-foo").unwrap();
328 assert_eq!(name, "x-foo");
329 assert!(NameKeyVisitor.visit_bytes::<DeError>(b"x foo").is_err());
330
331 let v = ValueVisitor.visit_bytes::<DeError>(b"a").unwrap();
332 assert_eq!(v, Value::One(HeaderValue::from_static("a")));
333 let v = ValueVisitor
334 .visit_byte_buf::<DeError>(b"b".to_vec())
335 .unwrap();
336 assert_eq!(v, Value::One(HeaderValue::from_static("b")));
337 let v = ValueVisitor
338 .visit_string::<DeError>("c".to_string())
339 .unwrap();
340 assert_eq!(v, Value::One(HeaderValue::from_static("c")));
341 assert!(ValueVisitor.visit_bytes::<DeError>(b"\n").is_err());
342 assert!(
343 HeaderValueVisitor
344 .visit_byte_buf::<DeError>(b"\n".to_vec())
345 .is_err()
346 );
347 assert!(
348 HeaderValueVisitor
349 .visit_string::<DeError>("\n".to_string())
350 .is_err()
351 );
352
353 let map = serde_json::from_str::<HeaderMap>(r#"{"1":"a"}"#).unwrap();
355 assert_eq!(map.get("1").unwrap(), "a");
356 let err = serde_json::from_str::<HeaderMap>(r#"{"x":1}"#).unwrap_err();
357 assert!(
358 err.to_string()
359 .contains("a single header value or sequence"),
360 "{err}"
361 );
362 let err = serde_json::from_str::<HeaderMap>(r#"{"x":[1]}"#).unwrap_err();
363 assert!(err.to_string().contains("a header value"), "{err}");
364 let err = serde_json::from_str::<HeaderMap>(r#"{"x":"\u0001"}"#).unwrap_err();
365 assert!(err.to_string().contains("a valid header value"), "{err}");
366 let err =
367 serde_json::from_value::<HeaderMap>(serde_json::json!({"x": ["\u{1}"]})).unwrap_err();
368 assert_ne!(err.to_string(), "");
369 let err = NameKeyVisitor.visit_u8::<DeError>(1).unwrap_err();
370 assert!(err.to_string().contains("a header name"), "{err}");
371
372 let mut map = HeaderMap::new();
374 map.insert(USER_AGENT, HeaderValue::from_bytes(b"caf\xe9").unwrap());
375 assert!(serde_json::to_string(&map).is_err());
376 let map2 = bincode::deserialize::<HeaderMap>(&bincode::serialize(&map).unwrap()).unwrap();
377 assert_eq!(map, map2);
378
379 let mut map = HeaderMap::new();
380 map.insert(USER_AGENT, HeaderValue::from_bytes(b"caf\xc3\xa9").unwrap());
381 let json = serde_json::to_string(&map).unwrap();
382 assert_eq!(json, "{\"user-agent\":\"caf\u{e9}\"}");
383 assert_eq!(serde_json::from_str::<HeaderMap>(&json).unwrap(), map);
384 }
385}