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