Skip to main content

ntex_http/
serde.rs

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            // names are case-insensitive, merge values of duplicate keys
46            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/// Header name map key, supports both borrowed and owned strings.
58#[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            // For non-human-readable formats, always serialize as a sequence
106            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        // Make roundtrip
258        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        // Try mixed case header names
265        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        // Try decode empty header value
273        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        // non-borrowed keys
297        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        // duplicate keys are merged
308        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        // `expecting` messages
354        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        // non-utf8 values can not be serialized to human readable formats
373        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}