rama_http_headers/forwarded/
via.rs1use crate::{HeaderDecode, HeaderEncode, TypedHeader, util};
2use rama_core::error::BoxErrorExt as _;
3use rama_core::{
4 error::{BoxError, ErrorContext},
5 telemetry::tracing,
6};
7use rama_http_types::{HeaderName, HeaderValue, header};
8use rama_net::forwarded::{ForwardedElement, ForwardedProtocol, ForwardedVersion, NodeId};
9
10#[derive(Debug, Clone, PartialEq, Eq)]
35pub struct Via(Vec<ViaElement>);
36
37#[derive(Debug, Clone, PartialEq, Eq)]
38struct ViaElement {
39 protocol: Option<ForwardedProtocol>,
40 version: ForwardedVersion,
41 node_id: NodeId,
42}
43
44impl From<ViaElement> for ForwardedElement {
45 fn from(via: ViaElement) -> Self {
46 let mut el = Self::new_forwarded_by(via.node_id);
47 el.set_forwarded_version(via.version);
48 if let Some(protocol) = via.protocol {
49 el.set_forwarded_proto(protocol);
50 }
51 el
52 }
53}
54
55impl TypedHeader for Via {
56 fn name() -> &'static HeaderName {
57 &header::VIA
58 }
59}
60
61impl HeaderDecode for Via {
62 fn decode<'i, I: Iterator<Item = &'i HeaderValue>>(
63 values: &mut I,
64 ) -> Result<Self, crate::Error> {
65 util::csv::from_comma_delimited(values).map(Via)
66 }
67}
68
69impl HeaderEncode for Via {
70 fn encode<E: Extend<HeaderValue>>(&self, values: &mut E) {
71 let s = rama_utils::fmt::display_fn(|f: &mut std::fmt::Formatter<'_>| {
72 util::csv::fmt_comma_delimited(&mut *f, self.0.iter())
73 })
74 .to_string();
75 match HeaderValue::try_from(s) {
76 Ok(value) => values.extend(::std::iter::once(value)),
77 Err(err) => tracing::debug!("failed to encode via as header value: {err}"),
78 }
79 }
80}
81
82impl FromIterator<ViaElement> for Via {
83 fn from_iter<T>(iter: T) -> Self
84 where
85 T: IntoIterator<Item = ViaElement>,
86 {
87 Self(iter.into_iter().collect())
88 }
89}
90
91impl super::ForwardHeader for Via {
92 fn try_from_forwarded<'a, I>(input: I) -> Option<Self>
93 where
94 I: IntoIterator<Item = &'a ForwardedElement>,
95 {
96 let vec: Vec<_> = input
97 .into_iter()
98 .filter_map(|el| {
99 let node_id = el.forwarded_by()?.clone();
100 let version = el.forwarded_version()?;
101 let protocol = el.forwarded_proto();
102 Some(ViaElement {
103 protocol,
104 version,
105 node_id,
106 })
107 })
108 .collect();
109 if vec.is_empty() {
110 None
111 } else {
112 Some(Self(vec))
113 }
114 }
115}
116
117impl IntoIterator for Via {
118 type Item = ForwardedElement;
119 type IntoIter = ViaIterator;
120
121 fn into_iter(self) -> Self::IntoIter {
122 ViaIterator(self.0.into_iter())
123 }
124}
125
126#[derive(Debug, Clone)]
127pub struct ViaIterator(std::vec::IntoIter<ViaElement>);
129
130impl Iterator for ViaIterator {
131 type Item = ForwardedElement;
132
133 fn next(&mut self) -> Option<Self::Item> {
134 self.0.next().map(Into::into)
135 }
136}
137
138impl std::str::FromStr for ViaElement {
139 type Err = BoxError;
140
141 #[expect(
142 clippy::unreachable,
143 reason = "the `position` predicate above only matches `b'/'` or `b' '`, so the wildcard arm is unreachable"
144 )]
145 fn from_str(s: &str) -> Result<Self, Self::Err> {
146 let mut bytes = s.as_bytes();
147
148 bytes = trim_left(bytes);
149
150 let (protocol, version) = match bytes.iter().position(|b| *b == b'/' || *b == b' ') {
151 Some(index) => match bytes[index] {
152 b'/' => {
153 let protocol: ForwardedProtocol = std::str::from_utf8(&bytes[..index])
154 .context("parse via protocol as utf-8")?
155 .try_into()
156 .context("parse via utf-8 protocol as protocol")?;
157 bytes = &bytes[index + 1..];
158 let index = bytes.iter().position(|b| *b == b' ').ok_or_else(|| {
159 BoxError::from_static_str("via str: missing space after protocol separator")
160 })?;
161 let version =
162 ForwardedVersion::try_from(&bytes[..index]).context("parse via version")?;
163 bytes = &bytes[index + 1..];
164 (Some(protocol), version)
165 }
166 b' ' => {
167 let version =
168 ForwardedVersion::try_from(&bytes[..index]).context("parse via version")?;
169 bytes = &bytes[index + 1..];
170 (None, version)
171 }
172 _ => unreachable!(),
173 },
174 None => {
175 return Err(BoxError::from_static_str("via str: missing version"));
176 }
177 };
178
179 bytes = trim_right(trim_left(bytes));
180 let node_id = NodeId::from_bytes_lossy(bytes);
181
182 Ok(Self {
183 protocol,
184 version,
185 node_id,
186 })
187 }
188}
189
190impl std::fmt::Display for ViaElement {
191 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
192 if let Some(ref proto) = self.protocol {
193 write!(f, "{proto}/")?;
194 }
195 write!(f, "{} {}", self.version, self.node_id)
196 }
197}
198
199fn trim_left(b: &[u8]) -> &[u8] {
200 let mut offset = 0;
201 while offset < b.len() && b[offset] == b' ' {
202 offset += 1;
203 }
204 &b[offset..]
205}
206
207fn trim_right(b: &[u8]) -> &[u8] {
208 if b.is_empty() {
209 return b;
210 }
211
212 let mut offset = b.len();
213 while offset > 0 && b[offset - 1] == b' ' {
214 offset -= 1;
215 }
216 &b[..offset]
217}
218
219#[cfg(test)]
220mod tests {
221 use super::*;
222
223 use rama_http_types::HeaderValue;
224
225 macro_rules! test_header {
226 ($name: ident, $input: expr, $expected: expr) => {
227 #[test]
228 fn $name() {
229 assert_eq!(
230 Via::decode(
231 &mut $input
232 .into_iter()
233 .map(|s| HeaderValue::from_bytes(s.as_bytes()).unwrap())
234 .collect::<Vec<_>>()
235 .iter()
236 )
237 .ok(),
238 $expected,
239 );
240 }
241 };
242 }
243
244 test_header!(
246 test1,
247 vec!["1.1 vegur"],
248 Some(Via(vec![ViaElement {
249 protocol: None,
250 version: ForwardedVersion::HTTP_11,
251 node_id: NodeId::try_from_str("vegur").unwrap(),
252 }]))
253 );
254 test_header!(
255 test2,
256 vec!["1.1 vegur "],
257 Some(Via(vec![ViaElement {
258 protocol: None,
259 version: ForwardedVersion::HTTP_11,
260 node_id: NodeId::try_from_str("vegur").unwrap(),
261 }]))
262 );
263 test_header!(
264 test3,
265 vec!["1.0 fred, 1.1 p.example.net"],
266 Some(Via(vec![
267 ViaElement {
268 protocol: None,
269 version: ForwardedVersion::HTTP_10,
270 node_id: NodeId::try_from_str("fred").unwrap(),
271 },
272 ViaElement {
273 protocol: None,
274 version: ForwardedVersion::HTTP_11,
275 node_id: NodeId::try_from_str("p.example.net").unwrap(),
276 }
277 ]))
278 );
279 test_header!(
280 test4,
281 vec!["1.0 fred , 1.1 p.example.net "],
282 Some(Via(vec![
283 ViaElement {
284 protocol: None,
285 version: ForwardedVersion::HTTP_10,
286 node_id: NodeId::try_from_str("fred").unwrap(),
287 },
288 ViaElement {
289 protocol: None,
290 version: ForwardedVersion::HTTP_11,
291 node_id: NodeId::try_from_str("p.example.net").unwrap(),
292 }
293 ]))
294 );
295 test_header!(
296 test5,
297 vec!["1.0 fred", "1.1 p.example.net"],
298 Some(Via(vec![
299 ViaElement {
300 protocol: None,
301 version: ForwardedVersion::HTTP_10,
302 node_id: NodeId::try_from_str("fred").unwrap(),
303 },
304 ViaElement {
305 protocol: None,
306 version: ForwardedVersion::HTTP_11,
307 node_id: NodeId::try_from_str("p.example.net").unwrap(),
308 }
309 ]))
310 );
311 test_header!(
312 test6,
313 vec!["HTTP/1.1 proxy.example.re, 1.1 edge_1"],
314 Some(Via(vec![
315 ViaElement {
316 protocol: Some(ForwardedProtocol::HTTP),
317 version: ForwardedVersion::HTTP_11,
318 node_id: NodeId::try_from_str("proxy.example.re").unwrap(),
319 },
320 ViaElement {
321 protocol: None,
322 version: ForwardedVersion::HTTP_11,
323 node_id: NodeId::try_from_str("edge_1").unwrap(),
324 }
325 ]))
326 );
327 test_header!(
328 test7,
329 vec!["1.1 2e9b3ee4d534903f433e1ed8ea30e57a.cloudfront.net (CloudFront)"],
330 Some(Via(vec![ViaElement {
331 protocol: None,
332 version: ForwardedVersion::HTTP_11,
333 node_id: NodeId::try_from_str(
334 "2e9b3ee4d534903f433e1ed8ea30e57a.cloudfront.net__CloudFront_"
335 )
336 .unwrap(),
337 }]))
338 );
339
340 #[test]
341 fn test_via_symmetric_encoder() {
342 for via_input in [
343 Via(vec![
344 ViaElement {
345 protocol: None,
346 version: ForwardedVersion::HTTP_10,
347 node_id: NodeId::try_from_str("fred").unwrap(),
348 },
349 ViaElement {
350 protocol: None,
351 version: ForwardedVersion::HTTP_11,
352 node_id: NodeId::try_from_str("p.example.net").unwrap(),
353 },
354 ]),
355 Via(vec![
356 ViaElement {
357 protocol: Some(ForwardedProtocol::HTTP),
358 version: ForwardedVersion::HTTP_11,
359 node_id: NodeId::try_from_str("proxy.example.re").unwrap(),
360 },
361 ViaElement {
362 protocol: None,
363 version: ForwardedVersion::HTTP_11,
364 node_id: NodeId::try_from_str("edge_1").unwrap(),
365 },
366 ]),
367 Via(vec![ViaElement {
368 protocol: None,
369 version: ForwardedVersion::HTTP_11,
370 node_id: NodeId::try_from_str(
371 "2e9b3ee4d534903f433e1ed8ea30e57a.cloudfront.net__CloudFront_",
372 )
373 .unwrap(),
374 }]),
375 ] {
376 let mut values = Vec::new();
377 via_input.encode(&mut values);
378 let via_output = Via::decode(&mut values.iter()).unwrap();
379 assert_eq!(via_input, via_output);
380 }
381 }
382}