Skip to main content

rama_net/forwarded/
node.rs

1use core::{
2    fmt,
3    net::{IpAddr, Ipv6Addr, SocketAddr},
4};
5
6use crate::std::borrow::ToOwned;
7use crate::std::{self as std, string::String, vec::Vec};
8
9use super::{ObfNode, ObfPort};
10use crate::address::{Domain, Host, HostWithOptPort, HostWithPort, SocketAddress};
11
12use rama_core::error::BoxErrorExt as _;
13use rama_core::error::{BoxError, ErrorContext};
14use rama_utils::str::smol_str::SmolStr;
15
16#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
17/// Node Identifier
18///
19/// The node identifier is one of the following:
20///
21/// - The client's IP address, with an optional port number
22/// - A token indicating that the IP address of the client is not known
23///   to the proxy server (unknown)
24/// - A generated token, allowing for tracing and debugging, while
25///   allowing the internal structure or sensitive information to be
26///   hidden
27///
28/// As specified in proposal section:
29/// <https://datatracker.ietf.org/doc/html/rfc7239#section-6>
30pub struct NodeId {
31    name: NodeName,
32    port: Option<NodePort>,
33}
34
35#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
36enum NodeName {
37    Unknown,
38    Ip(IpAddr),
39    Obf(ObfNode),
40}
41
42#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
43enum NodePort {
44    Num(u16),
45    Obf(ObfPort),
46}
47
48impl NodeId {
49    /// Try to convert a vector of bytes to a [`NodeId`].
50    pub fn try_from_bytes(vec: Vec<u8>) -> Result<Self, BoxError> {
51        vec.try_into()
52    }
53
54    /// Try to convert a string slice to a [`NodeId`].
55    pub fn try_from_str(s: &str) -> Result<Self, BoxError> {
56        s.to_owned().try_into()
57    }
58
59    #[inline]
60    /// Converts a vector of bytes to a [`NodeId`], converting invalid characters to underscore.
61    #[must_use]
62    pub fn from_bytes_lossy(vec: &[u8]) -> Self {
63        let s = String::from_utf8_lossy(vec);
64        Self::from_str_lossy(&s)
65    }
66
67    /// Converts a string slice to a [`NodeId`], converting invalid characters to underscore.
68    pub fn from_str_lossy(s: &str) -> Self {
69        let s_original = s;
70
71        if s.eq_ignore_ascii_case(UNKNOWN_STR) {
72            return Self {
73                name: NodeName::Unknown,
74                port: None,
75            };
76        }
77
78        if let Ok(ip) = try_to_parse_str_to_ip(s) {
79            // early return to prevent stuff like `::1` to
80            // be interpreted as node { name = obf(:), port = num(1) }
81            return Self {
82                name: NodeName::Ip(ip),
83                port: None,
84            };
85        }
86
87        let (s, port) = try_to_split_node_port_lossy_from_str(s);
88        let name = try_to_parse_str_to_ip(s)
89            .map(NodeName::Ip)
90            .unwrap_or_else(|_| NodeName::Obf(ObfNode::from_str_lossy(s)));
91
92        match name {
93            NodeName::Ip(IpAddr::V6(_)) if port.is_some() && !s.starts_with('[') => Self {
94                name: NodeName::Obf(ObfNode::from_str_lossy(s_original)),
95                port: None,
96            },
97            _ => Self { name, port },
98        }
99    }
100
101    /// Return the [`IpAddr`] if one was defined for this [`NodeId`].
102    #[must_use]
103    pub fn ip(&self) -> Option<IpAddr> {
104        match &self.name {
105            NodeName::Ip(addr) => Some(*addr),
106            NodeName::Unknown | NodeName::Obf(_) => None,
107        }
108    }
109
110    /// Return true if this [`NodeId`] has a any kind of port defined,
111    /// even if obfuscated.
112    #[must_use]
113    pub fn has_any_port(&self) -> bool {
114        self.port.is_some()
115    }
116
117    /// Return the numeric port if one was defined for this [`NodeId`].
118    #[must_use]
119    pub fn port(&self) -> Option<u16> {
120        if let Some(NodePort::Num(n)) = self.port {
121            Some(n)
122        } else {
123            None
124        }
125    }
126
127    /// Return the [`HostWithPort`] if this [`NodeId`] has either
128    /// an [`IpAddr`] or [`Domain`] defined, as well as a numeric port.
129    #[must_use]
130    pub fn authority(&self) -> Option<HostWithPort> {
131        match (&self.name, self.port()) {
132            (NodeName::Ip(ip), Some(port)) => Some((*ip, port).into()),
133            // every domain is a valid node name, but not every valid node name is a valid domain!!
134            (NodeName::Obf(s), Some(port)) => s
135                .as_str()
136                .parse::<Domain>()
137                .ok()
138                .map(|domain| (domain, port).into()),
139            _ => None,
140        }
141    }
142}
143
144impl NodePort {
145    /// Converts a string slice to a [`NodePort`], converting invalid characters to underscore.
146    fn from_str_lossy(s: &str) -> Self {
147        s.parse::<u16>()
148            .map(NodePort::Num)
149            .unwrap_or_else(|_| Self::Obf(ObfPort::from_str_lossy(s)))
150    }
151}
152
153impl From<IpAddr> for NodeId {
154    #[inline]
155    fn from(ip: IpAddr) -> Self {
156        (ip, None).into()
157    }
158}
159
160impl From<(IpAddr, u16)> for NodeId {
161    #[inline]
162    fn from((ip, port): (IpAddr, u16)) -> Self {
163        (ip, Some(port)).into()
164    }
165}
166
167impl From<(IpAddr, Option<u16>)> for NodeId {
168    fn from((ip, port): (IpAddr, Option<u16>)) -> Self {
169        Self {
170            name: NodeName::Ip(ip),
171            port: port.map(NodePort::Num),
172        }
173    }
174}
175
176impl From<Domain> for NodeId {
177    #[inline]
178    fn from(domain: Domain) -> Self {
179        (domain, None).into()
180    }
181}
182
183impl From<(Domain, u16)> for NodeId {
184    #[inline]
185    fn from((domain, port): (Domain, u16)) -> Self {
186        (domain, Some(port)).into()
187    }
188}
189
190impl From<(Domain, Option<u16>)> for NodeId {
191    fn from((domain, port): (Domain, Option<u16>)) -> Self {
192        Self {
193            // NOTE: this assumes all domains are valid obf nodes,
194            // which should be ok given the validation rules for domains are more strict!
195            name: NodeName::Obf(ObfNode::from_inner(SmolStr::from(domain.as_str()))),
196            port: port.map(NodePort::Num),
197        }
198    }
199}
200
201impl From<HostWithOptPort> for NodeId {
202    fn from(value: HostWithOptPort) -> Self {
203        let HostWithOptPort { host, port } = value;
204        // RFC 7239 has no representation for "explicit empty port" — fold
205        // `OptPort::Empty` into `Unset` at the typed-identifier layer.
206        let port = port.as_u16();
207        node_id_from_host_port(&host, port)
208    }
209}
210
211impl From<HostWithPort> for NodeId {
212    fn from(value: HostWithPort) -> Self {
213        let HostWithPort { host, port } = value;
214        node_id_from_host_port(&host, Some(port))
215    }
216}
217
218/// Promote `host` to the most specific RFC 7239 node identifier:
219/// IP first (pct-encoded IP literals inside `Uninterpreted` bridge for
220/// free), then Domain (pct-encoded reg-names that decode to a domain
221/// land here too), and only genuinely non-typed shapes (sub-delim
222/// reg-name, IPvFuture) fall through to lossy `obfnode`.
223fn node_id_from_host_port(host: &Host, port: Option<u16>) -> NodeId {
224    if let Ok(ip) = host.try_as_ip() {
225        return (ip, port).into();
226    }
227    match host.try_as_domain() {
228        Ok(domain) => (domain.into_owned(), port).into(),
229        Err(_) => NodeId {
230            // Sub-delim reg-name / IPvFuture / non-promotable bytes —
231            // route through `from_str_lossy` so any non-obfnode byte
232            // becomes `_`. Lossy by design.
233            name: NodeName::Obf(ObfNode::from_str_lossy(&host.to_str())),
234            port: port.map(NodePort::Num),
235        },
236    }
237}
238
239impl From<SocketAddr> for NodeId {
240    fn from(addr: SocketAddr) -> Self {
241        Self {
242            name: NodeName::Ip(addr.ip()),
243            port: Some(NodePort::Num(addr.port())),
244        }
245    }
246}
247
248impl From<&SocketAddr> for NodeId {
249    fn from(addr: &SocketAddr) -> Self {
250        Self {
251            name: NodeName::Ip(addr.ip()),
252            port: Some(NodePort::Num(addr.port())),
253        }
254    }
255}
256
257impl From<SocketAddress> for NodeId {
258    fn from(addr: SocketAddress) -> Self {
259        Self {
260            name: NodeName::Ip(addr.ip_addr),
261            port: Some(NodePort::Num(addr.port)),
262        }
263    }
264}
265
266impl From<&SocketAddress> for NodeId {
267    fn from(addr: &SocketAddress) -> Self {
268        Self {
269            name: NodeName::Ip(addr.ip_addr),
270            port: Some(NodePort::Num(addr.port)),
271        }
272    }
273}
274
275const UNKNOWN_STR: &str = "unknown";
276
277impl fmt::Display for NodeId {
278    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> core::fmt::Result {
279        match &self.name {
280            NodeName::Unknown => UNKNOWN_STR.fmt(f),
281            NodeName::Ip(ip) => match &self.port {
282                None => ip.fmt(f),
283                Some(port) => match ip {
284                    core::net::IpAddr::V4(ip) => write!(f, "{ip}:{port}"),
285                    core::net::IpAddr::V6(ip) => write!(f, "[{ip}]:{port}"),
286                },
287            },
288            NodeName::Obf(s) => match &self.port {
289                None => s.fmt(f),
290                Some(port) => write!(f, "{s}:{port}"),
291            },
292        }
293    }
294}
295
296impl fmt::Display for NodePort {
297    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
298        match self {
299            Self::Num(num) => num.fmt(f),
300            Self::Obf(s) => s.fmt(f),
301        }
302    }
303}
304
305impl core::str::FromStr for NodeId {
306    type Err = BoxError;
307
308    fn from_str(s: &str) -> Result<Self, Self::Err> {
309        Self::try_from(s)
310    }
311}
312
313impl TryFrom<String> for NodeId {
314    type Error = BoxError;
315
316    fn try_from(s: String) -> Result<Self, Self::Error> {
317        s.as_str().try_into()
318    }
319}
320
321impl TryFrom<&str> for NodeId {
322    type Error = BoxError;
323
324    fn try_from(s: &str) -> Result<Self, Self::Error> {
325        if s.eq_ignore_ascii_case(UNKNOWN_STR) {
326            return Ok(Self {
327                name: NodeName::Unknown,
328                port: None,
329            });
330        }
331
332        if let Ok(ip) = try_to_parse_str_to_ip(s) {
333            // early return to prevent stuff like `::1` to
334            // be interpreted as node { name = obf(:), port = num(1) }
335            return Ok(Self {
336                name: NodeName::Ip(ip),
337                port: None,
338            });
339        }
340
341        let (s, port) = try_to_split_node_port_from_str(s);
342        let name = try_to_parse_str_to_ip(s)
343            .map(NodeName::Ip)
344            .or_else(|_| s.parse::<ObfNode>().map(NodeName::Obf))
345            .context("parse str as Node")?;
346
347        match name {
348            NodeName::Ip(IpAddr::V6(_)) if port.is_some() && !s.starts_with('[') => Err(
349                BoxError::from_static_str("missing brackets for node IPv6 address with port"),
350            ),
351            _ => Ok(Self { name, port }),
352        }
353    }
354}
355
356fn try_to_parse_str_to_ip(value: &str) -> Result<IpAddr, BoxError> {
357    if value.starts_with('[') || value.ends_with(']') {
358        let value = value
359            .strip_prefix('[')
360            .and_then(|value| value.strip_suffix(']'))
361            .context("strip brackets from ipv6 str")?;
362        Ok(IpAddr::V6(
363            value.parse::<Ipv6Addr>().context("parse str as ipv6")?,
364        ))
365    } else {
366        value.parse::<IpAddr>().context("parse ipv4/6 str")
367    }
368}
369
370impl TryFrom<Vec<u8>> for NodeId {
371    type Error = BoxError;
372
373    fn try_from(bytes: Vec<u8>) -> Result<Self, Self::Error> {
374        let s = String::from_utf8(bytes).context("parse node from bytes")?;
375        s.try_into()
376    }
377}
378
379impl TryFrom<&[u8]> for NodeId {
380    type Error = BoxError;
381
382    fn try_from(bytes: &[u8]) -> Result<Self, Self::Error> {
383        let s = core::str::from_utf8(bytes).context("parse node from bytes")?;
384        s.try_into()
385    }
386}
387
388fn try_to_split_node_port_from_str(s: &str) -> (&str, Option<NodePort>) {
389    if let Some(colon) = s.as_bytes().iter().rposition(|c| *c == b':') {
390        match s[colon + 1..].parse() {
391            Ok(port) => (&s[..colon], Some(port)),
392            Err(_) => (s, None),
393        }
394    } else {
395        (s, None)
396    }
397}
398
399fn try_to_split_node_port_lossy_from_str(s: &str) -> (&str, Option<NodePort>) {
400    if let Some(colon) = s.as_bytes().iter().rposition(|c| *c == b':') {
401        let port = NodePort::from_str_lossy(&s[colon + 1..]);
402        let s = &s[..colon];
403        (s, Some(port))
404    } else {
405        (s, None)
406    }
407}
408
409impl core::str::FromStr for NodePort {
410    type Err = BoxError;
411
412    fn from_str(s: &str) -> Result<Self, Self::Err> {
413        s.parse::<u16>()
414            .map(NodePort::Num)
415            .or_else(|_| s.parse::<ObfPort>().map(NodePort::Obf))
416            .context("parse str as NodePort")
417    }
418}
419
420use rama_utils::macros::serde_str::impl_serde_str;
421
422impl_serde_str!(display NodeId);
423
424#[cfg(test)]
425mod tests {
426    use super::*;
427
428    #[test]
429    fn test_parse_node_id_valid() {
430        for (s, expected) in [
431            (
432                "unknown",
433                NodeId {
434                    name: NodeName::Unknown,
435                    port: None,
436                },
437            ),
438            (
439                "::1",
440                NodeId {
441                    name: NodeName::Ip("::1".parse().unwrap()),
442                    port: None,
443                },
444            ),
445            (
446                "127.0.0.1",
447                NodeId {
448                    name: NodeName::Ip("127.0.0.1".parse().unwrap()),
449                    port: None,
450                },
451            ),
452            (
453                "192.0.2.43:47011",
454                NodeId {
455                    name: NodeName::Ip("192.0.2.43".parse().unwrap()),
456                    port: Some(NodePort::Num(47011)),
457                },
458            ),
459            (
460                "[2001:db8:cafe::17]:47011",
461                NodeId {
462                    name: NodeName::Ip("2001:db8:cafe::17".parse().unwrap()),
463                    port: Some(NodePort::Num(47011)),
464                },
465            ),
466            (
467                "192.0.2.43:_foo",
468                NodeId {
469                    name: NodeName::Ip("192.0.2.43".parse().unwrap()),
470                    port: Some(NodePort::Obf(ObfPort::from_static("_foo"))),
471                },
472            ),
473            (
474                "[2001:db8:cafe::17]:_bar",
475                NodeId {
476                    name: NodeName::Ip("2001:db8:cafe::17".parse().unwrap()),
477                    port: Some(NodePort::Obf(ObfPort::from_static("_bar"))),
478                },
479            ),
480            (
481                "foo",
482                NodeId {
483                    name: NodeName::Obf(ObfNode::from_static("foo")),
484                    port: None,
485                },
486            ),
487            (
488                "_foo",
489                NodeId {
490                    name: NodeName::Obf(ObfNode::from_static("_foo")),
491                    port: None,
492                },
493            ),
494            (
495                "foo:_bar",
496                NodeId {
497                    name: NodeName::Obf(ObfNode::from_static("foo")),
498                    port: Some(NodePort::Obf(ObfPort::from_static("_bar"))),
499                },
500            ),
501            (
502                "foo:42",
503                NodeId {
504                    name: NodeName::Obf(ObfNode::from_static("foo")),
505                    port: Some(NodePort::Num(42)),
506                },
507            ),
508        ] {
509            match s.parse::<NodeId>() {
510                Err(err) => panic!("failed to parse '{s}': {err}"),
511                Ok(node_id) => assert_eq!(node_id, expected, "parse: {s}"),
512            }
513        }
514    }
515
516    #[test]
517    fn test_parse_node_id_invalid() {
518        for s in [
519            "",
520            "@",
521            "2001:db8:3333:4444:5555:6666:7777:8888:80",
522            "foo:bar",
523            "foo:_b+r",
524            "😀",
525            "abcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyz",
526        ] {
527            let node_result = s.parse::<NodeId>();
528            assert!(
529                node_result.is_err(),
530                "parse invalid: {s}; parsed: {node_result:?}",
531            );
532        }
533    }
534
535    #[test]
536    fn test_parse_node_id_lossy() {
537        for (s, expected) in [
538            (
539                "",
540                NodeId {
541                    name: NodeName::Obf(ObfNode::from_static("_")),
542                    port: None,
543                },
544            ),
545            (
546                "@",
547                NodeId {
548                    name: NodeName::Obf(ObfNode::from_static("_")),
549                    port: None,
550                },
551            ),
552            (
553                "2001:db8:3333:4444:5555:6666:7777:8888:80",
554                NodeId {
555                    name: NodeName::Obf(ObfNode::from_static(
556                        "2001_db8_3333_4444_5555_6666_7777_8888_80",
557                    )),
558                    port: None,
559                },
560            ),
561            (
562                "foo:bar",
563                NodeId {
564                    name: NodeName::Obf(ObfNode::from_static("foo")),
565                    port: Some(NodePort::Obf(ObfPort::from_static("_bar"))),
566                },
567            ),
568            (
569                "foo:_b+r",
570                NodeId {
571                    name: NodeName::Obf(ObfNode::from_static("foo")),
572                    port: Some(NodePort::Obf(ObfPort::from_static("_b_r"))),
573                },
574            ),
575            (
576                "😀",
577                NodeId {
578                    name: NodeName::Obf(ObfNode::from_static("____")),
579                    port: None,
580                },
581            ),
582            (
583                "abcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyz",
584                NodeId {
585                    name: NodeName::Obf(ObfNode::from_static(
586                        "abcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuv",
587                    )),
588                    port: None,
589                },
590            ),
591        ] {
592            let node_id = NodeId::from_str_lossy(s);
593            assert_eq!(node_id, expected, "parse str: {s}");
594
595            let node_id = NodeId::from_bytes_lossy(s.as_bytes());
596            assert_eq!(node_id, expected, "parse bytes: {s}");
597        }
598    }
599}