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 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)]
139pub 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 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}