rama_http_headers/
map_ext.rs1#![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
12pub trait HeaderMapExt: self::sealed::Sealed {
14 fn typed_insert<H>(&mut self, header: H)
16 where
17 H: HeaderEncode;
18
19 fn typed_get<H>(&self) -> Option<H>
21 where
22 H: HeaderDecode;
23
24 fn typed_try_get<H>(&self) -> Result<Option<H>, Error>
26 where
27 H: HeaderDecode;
28
29 fn remove_all<K>(&mut self, name: K) -> usize
39 where
40 K: AsHeaderName + Clone;
41
42 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 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 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}