Skip to main content

rama_http_headers/
map_ext.rs

1#![expect(
2    clippy::unreachable,
3    reason = "vendored from upstream `headers`: `State::Tmp` is a transient placeholder used inside `mem::replace`, never observed by the next iteration"
4)]
5
6use rama_http_types::{HeaderValue, header, header::AsHeaderName};
7
8use crate::{HeaderDecode, HeaderEncode};
9
10use super::Error;
11
12/// An extension trait adding "typed" methods to `http::HeaderMap`.
13pub trait HeaderMapExt: self::sealed::Sealed {
14    /// Inserts the typed header into this `HeaderMap`.
15    fn typed_insert<H>(&mut self, header: H)
16    where
17        H: HeaderEncode;
18
19    /// Tries to find the header by name, and then decode it into `H`.
20    fn typed_get<H>(&self) -> Option<H>
21    where
22        H: HeaderDecode;
23
24    /// Tries to find the header by name, and then decode it into `H`.
25    fn typed_try_get<H>(&self) -> Result<Option<H>, Error>
26    where
27        H: HeaderDecode;
28
29    /// Remove every value associated with a header name and return how many
30    /// were removed.
31    ///
32    /// Note: `HeaderMap::remove` already removes the whole entry (all
33    /// values) for a header — it just *returns* only the first value as
34    /// `Option<HeaderValue>`. This method gives callers an accurate count
35    /// without needing to call `HeaderMap::get_all` manually beforehand.
36    /// Use [`remove_all_values`](Self::remove_all_values) when the values
37    /// themselves are needed.
38    fn remove_all<K>(&mut self, name: K) -> usize
39    where
40        K: AsHeaderName + Clone;
41
42    /// Remove every value associated with a header name and return them in
43    /// iteration order.
44    ///
45    /// Useful when you need to inspect, log, or relocate the removed
46    /// values — `HeaderMap::remove` only surfaces the first.
47    fn remove_all_values<K>(&mut self, name: K) -> Vec<HeaderValue>
48    where
49        K: AsHeaderName + Clone;
50}
51
52impl HeaderMapExt for rama_http_types::HeaderMap {
53    fn typed_insert<H>(&mut self, header: H)
54    where
55        H: HeaderEncode,
56    {
57        let entry = self.entry(H::name());
58        let mut values = ToValues {
59            state: State::First(entry),
60        };
61        header.encode(&mut values);
62    }
63
64    fn typed_get<H>(&self) -> Option<H>
65    where
66        H: HeaderDecode,
67    {
68        HeaderMapExt::typed_try_get(self).unwrap_or(None)
69    }
70
71    fn typed_try_get<H>(&self) -> Result<Option<H>, Error>
72    where
73        H: HeaderDecode,
74    {
75        let mut values = self.get_all(H::name()).iter();
76        if values.size_hint() == (0, Some(0)) {
77            Ok(None)
78        } else {
79            H::decode(&mut values).map(Some)
80        }
81    }
82
83    fn remove_all<K>(&mut self, name: K) -> usize
84    where
85        K: AsHeaderName + Clone,
86    {
87        // `HeaderMap::remove` removes the whole entry; count separately
88        // because it only surfaces the first value.
89        let count = self.get_all(name.clone()).iter().count();
90        self.remove(name);
91        count
92    }
93
94    fn remove_all_values<K>(&mut self, name: K) -> Vec<HeaderValue>
95    where
96        K: AsHeaderName + Clone,
97    {
98        // Collect every value first, then drop the entry in one call.
99        let values: Vec<HeaderValue> = self.get_all(name.clone()).iter().cloned().collect();
100        self.remove(name);
101        values
102    }
103}
104
105struct ToValues<'a> {
106    state: State<'a>,
107}
108
109#[derive(Debug)]
110enum State<'a> {
111    First(header::Entry<'a, HeaderValue>),
112    Latter(header::OccupiedEntry<'a, HeaderValue>),
113    Tmp,
114}
115
116impl Extend<HeaderValue> for ToValues<'_> {
117    fn extend<T: IntoIterator<Item = HeaderValue>>(&mut self, iter: T) {
118        for value in iter {
119            let entry = match ::std::mem::replace(&mut self.state, State::Tmp) {
120                State::First(header::Entry::Occupied(mut e)) => {
121                    e.insert(value);
122                    e
123                }
124                State::First(header::Entry::Vacant(e)) => e.insert_entry(value),
125                State::Latter(mut e) => {
126                    e.append(value);
127                    e
128                }
129                State::Tmp => unreachable!("ToValues State::Tmp"),
130            };
131            self.state = State::Latter(entry);
132        }
133    }
134}
135
136mod sealed {
137    pub trait Sealed {}
138    impl Sealed for ::rama_http_types::HeaderMap {}
139}
140
141#[cfg(test)]
142mod test {
143    use super::*;
144    use rama_http_types::HeaderMap;
145
146    #[test]
147    fn test_remove_all_drops_every_value() {
148        let mut map = HeaderMap::new();
149        map.append(header::CONTENT_LENGTH, HeaderValue::from(42u64));
150        map.append(header::CONTENT_LENGTH, HeaderValue::from(99u64));
151        map.append(header::CONTENT_TYPE, HeaderValue::from_static("text/plain"));
152
153        let removed = map.remove_all(&header::CONTENT_LENGTH);
154        assert_eq!(removed, 2);
155        assert!(!map.contains_key(header::CONTENT_LENGTH));
156        assert_eq!(
157            map.get(header::CONTENT_TYPE).unwrap().as_bytes(),
158            b"text/plain"
159        );
160    }
161
162    #[test]
163    fn test_remove_all_returns_zero_when_absent() {
164        let mut map = HeaderMap::new();
165        map.insert(header::CONTENT_TYPE, HeaderValue::from_static("text/plain"));
166        assert_eq!(map.remove_all(&header::CONTENT_LENGTH), 0);
167        assert_eq!(map.len(), 1);
168    }
169
170    #[test]
171    fn test_remove_all_values_collects_in_order() {
172        let mut map = HeaderMap::new();
173        map.append(header::CONTENT_LENGTH, HeaderValue::from(1u64));
174        map.append(header::CONTENT_LENGTH, HeaderValue::from(2u64));
175        map.append(header::CONTENT_LENGTH, HeaderValue::from(3u64));
176
177        let values = map.remove_all_values(&header::CONTENT_LENGTH);
178        assert_eq!(values.len(), 3);
179        assert_eq!(values[0].to_str().unwrap(), "1");
180        assert_eq!(values[1].to_str().unwrap(), "2");
181        assert_eq!(values[2].to_str().unwrap(), "3");
182        assert!(!map.contains_key(header::CONTENT_LENGTH));
183    }
184}