Skip to main content

rama_http_headers/forwarded/
via.rs

1use 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/// The Via general header is added by proxies, both forward and reverse.
11///
12/// This header can appear in the request or response headers.
13/// It is used for tracking message forwards, avoiding request loops,
14/// and identifying the protocol capabilities of senders along the request/response chain.
15///
16/// It is recommended to use the [`Forwarded`](super::Forwarded) header instead if you can.
17///
18/// More info can be found at <https://developer.mozilla.org/en-US/docs/Web/HTTP/Headers/Via>.
19///
20/// # Syntax
21///
22/// ```text
23/// Via: [ <protocol-name> "/" ] <protocol-version> <host> [ ":" <port> ]
24/// Via: [ <protocol-name> "/" ] <protocol-version> <pseudonym>
25/// ```
26///
27/// # Example values
28///
29/// * `1.1 vegur`
30/// * `HTTP/1.1 GWA`
31/// * `1.0 fred, 1.1 p.example.net`
32/// * `HTTP/1.1 proxy.example.re, 1.1 edge_1`
33/// * `1.1 2e9b3ee4d534903f433e1ed8ea30e57a.cloudfront.net (CloudFront)`
34#[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)]
127/// An iterator over the `Via` header's elements.
128pub 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    // Tests from the Docs
245    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}