Skip to main content

watermelon_proto/
server_addr.rs

1use alloc::{str::FromStr, string::String};
2use core::{
3    fmt::{self, Debug, Display, Write},
4    net::IpAddr,
5    ops::Deref,
6};
7
8use bytestring::ByteString;
9use percent_encoding::{NON_ALPHANUMERIC, percent_decode_str, percent_encode};
10use serde::{Deserialize, Deserializer, Serialize, Serializer, de};
11use url::Url;
12
13/// Address of a NATS server
14#[derive(Clone, PartialEq, Eq)]
15pub struct ServerAddr {
16    protocol: Protocol,
17    transport: Transport,
18    host: Host,
19    port: u16,
20    username: ByteString,
21    password: ByteString,
22}
23
24/// The protocol of the NATS server
25#[derive(Debug, Copy, Clone, PartialEq, Eq, PartialOrd, Ord)]
26pub enum Protocol {
27    /// Plaintext with the option to later upgrade to TLS
28    ///
29    /// This option should only be used when esplicity wanting to
30    /// connect using a plaintext connection. Using this option
31    /// over the public internet or other untrusted networks
32    /// leaves the client open to MITM attacks.
33    ///
34    /// Corresponds to the `nats` scheme.
35    PossiblyPlain,
36    /// TLS connection
37    ///
38    /// Requires the TCP connection to successfully upgrade to TLS.
39    ///
40    /// Corresponds to the `tls` scheme.
41    TLS,
42}
43
44/// The transport protocol of the NATS server
45#[derive(Debug, Copy, Clone, PartialEq, Eq, PartialOrd, Ord)]
46pub enum Transport {
47    /// Transmit data over a TCP stream
48    TCP,
49    /// Transmit data over WebSocket frames
50    Websocket,
51}
52
53/// The hostname of the NATS server
54#[derive(Debug, Clone, PartialEq, Eq)]
55pub enum Host {
56    /// An IPv4 or IPv6 address
57    Ip(IpAddr),
58    /// A DNS hostname
59    Dns(ByteString),
60}
61
62impl ServerAddr {
63    /// Get the connection protocol
64    pub fn protocol(&self) -> Protocol {
65        self.protocol
66    }
67
68    /// Get the transport protocol
69    pub fn transport(&self) -> Transport {
70        self.transport
71    }
72
73    /// Get the hostname
74    pub fn host(&self) -> &Host {
75        &self.host
76    }
77
78    /// Get the port
79    pub fn port(&self) -> u16 {
80        self.port
81    }
82
83    fn is_default_port(&self) -> bool {
84        self.port == protocol_transport_to_port(self.protocol, self.transport)
85    }
86
87    /// Get the username
88    pub fn username(&self) -> Option<&str> {
89        if self.username.is_empty() {
90            None
91        } else {
92            Some(&self.username)
93        }
94    }
95
96    /// Get the password
97    pub fn password(&self) -> Option<&str> {
98        if self.password.is_empty() {
99            None
100        } else {
101            Some(&self.password)
102        }
103    }
104}
105
106impl FromStr for ServerAddr {
107    type Err = ServerAddrError;
108
109    fn from_str(value: &str) -> Result<Self, Self::Err> {
110        let url = value.parse::<Url>().map_err(ServerAddrError::InvalidUrl)?;
111
112        let (protocol, transport) = match url.scheme() {
113            "nats" => (Protocol::PossiblyPlain, Transport::TCP),
114            "tls" => (Protocol::TLS, Transport::TCP),
115            "ws" => (Protocol::PossiblyPlain, Transport::Websocket),
116            "wss" => (Protocol::TLS, Transport::Websocket),
117            _ => return Err(ServerAddrError::InvalidScheme),
118        };
119
120        let host = match url.host() {
121            Some(url::Host::Ipv4(addr)) => Host::Ip(IpAddr::V4(addr)),
122            Some(url::Host::Ipv6(addr)) => Host::Ip(IpAddr::V6(addr)),
123            Some(url::Host::Domain(host)) => {
124                // TODO: this shouldn't be necessary
125                let host = host
126                    .strip_prefix('[')
127                    .and_then(|host| host.strip_suffix(']'))
128                    .unwrap_or(host);
129                match host.parse::<IpAddr>() {
130                    Ok(ip) => Host::Ip(ip),
131                    Err(_) => Host::Dns(host.into()),
132                }
133            }
134            None => return Err(ServerAddrError::MissingHost),
135        };
136
137        let port = if let Some(port) = url.port() {
138            port
139        } else {
140            protocol_transport_to_port(protocol, transport)
141        };
142
143        let username = percent_decode_str(url.username())
144            .decode_utf8()
145            .map_err(|_| ServerAddrError::UsernameInvalidUtf8)?
146            .deref()
147            .into();
148        let password = percent_decode_str(url.password().unwrap_or_default())
149            .decode_utf8()
150            .map_err(|_| ServerAddrError::PasswordInvalidUtf8)?
151            .deref()
152            .into();
153
154        Ok(Self {
155            protocol,
156            transport,
157            host,
158            port,
159            username,
160            password,
161        })
162    }
163}
164
165impl Debug for ServerAddr {
166    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
167        let username = if self.username.is_empty() {
168            "<none>"
169        } else {
170            "<redacted>"
171        };
172        let password = if self.password.is_empty() {
173            "<none>"
174        } else {
175            "<redacted>"
176        };
177        f.debug_struct("ServerAddr")
178            .field("protocol", &self.protocol)
179            .field("transport", &self.transport)
180            .field("host", &self.host)
181            .field("port", &self.port)
182            .field("username", &username)
183            .field("password", &password)
184            .finish()
185    }
186}
187
188impl Display for ServerAddr {
189    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
190        f.write_str(match (self.protocol, self.transport) {
191            (Protocol::PossiblyPlain, Transport::TCP) => "nats",
192            (Protocol::TLS, Transport::TCP) => "tls",
193            (Protocol::PossiblyPlain, Transport::Websocket) => "ws",
194            (Protocol::TLS, Transport::Websocket) => "wss",
195        })?;
196        f.write_str("://")?;
197
198        if let Some(username) = self.username() {
199            Display::fmt(&percent_encode(username.as_bytes(), NON_ALPHANUMERIC), f)?;
200
201            if let Some(password) = self.password() {
202                write!(
203                    f,
204                    ":{}",
205                    percent_encode(password.as_bytes(), NON_ALPHANUMERIC)
206                )?;
207            }
208            f.write_char('@')?;
209        }
210
211        match &self.host {
212            Host::Ip(IpAddr::V4(addr)) => Display::fmt(addr, f)?,
213            Host::Ip(IpAddr::V6(addr)) => write!(f, "[{addr}]")?,
214            Host::Dns(record) => Display::fmt(record, f)?,
215        }
216        if !self.is_default_port() {
217            write!(f, ":{}", self.port)?;
218        }
219
220        Ok(())
221    }
222}
223
224impl<'de> Deserialize<'de> for ServerAddr {
225    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
226    where
227        D: Deserializer<'de>,
228    {
229        let val = String::deserialize(deserializer)?;
230        val.parse().map_err(de::Error::custom)
231    }
232}
233
234impl Serialize for ServerAddr {
235    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
236    where
237        S: Serializer,
238    {
239        serializer.collect_str(self)
240    }
241}
242
243/// An error encountered while parsing [`ServerAddr`]
244#[derive(Debug, thiserror::Error)]
245pub enum ServerAddrError {
246    /// The Url could not be parsed
247    #[error("invalid Url")]
248    InvalidUrl(#[source] url::ParseError),
249    /// The Url has a bad scheme
250    #[error("invalid Url scheme")]
251    InvalidScheme,
252    /// The Url is missing the hostname
253    #[error("missing host")]
254    MissingHost,
255    /// The Url contains a non-utf8 username
256    #[error("username is not utf-8")]
257    UsernameInvalidUtf8,
258    /// The Url contains a non-utf8 password
259    #[error("password is not utf-8")]
260    PasswordInvalidUtf8,
261}
262
263fn protocol_transport_to_port(protocol: Protocol, transport: Transport) -> u16 {
264    match (protocol, transport) {
265        (Protocol::PossiblyPlain | Protocol::TLS, Transport::TCP) => 4222,
266        (Protocol::PossiblyPlain, Transport::Websocket) => 80,
267        (Protocol::TLS, Transport::Websocket) => 443,
268    }
269}
270
271#[cfg(test)]
272mod tests {
273    use alloc::string::ToString;
274    use core::net::{IpAddr, Ipv4Addr, Ipv6Addr};
275
276    use super::{Host, Protocol, ServerAddr, Transport};
277
278    #[test]
279    fn nats() {
280        let server_addr = "nats://127.0.0.1".parse::<ServerAddr>().unwrap();
281        assert_eq!(server_addr.transport(), Transport::TCP);
282        assert_eq!(server_addr.protocol(), Protocol::PossiblyPlain);
283        assert_eq!(
284            server_addr.host(),
285            &Host::Ip(IpAddr::V4(Ipv4Addr::LOCALHOST))
286        );
287        assert_eq!(server_addr.port(), 4222);
288        assert_eq!(server_addr.username(), None);
289        assert_eq!(server_addr.password(), None);
290        assert_eq!(server_addr.to_string(), "nats://127.0.0.1");
291    }
292
293    #[test]
294    fn nats_non_default_port() {
295        let server_addr = "nats://127.0.0.1:4321".parse::<ServerAddr>().unwrap();
296        assert_eq!(server_addr.transport(), Transport::TCP);
297        assert_eq!(server_addr.protocol(), Protocol::PossiblyPlain);
298        assert_eq!(
299            server_addr.host(),
300            &Host::Ip(IpAddr::V4(Ipv4Addr::LOCALHOST))
301        );
302        assert_eq!(server_addr.port(), 4321);
303        assert_eq!(server_addr.username(), None);
304        assert_eq!(server_addr.password(), None);
305        assert_eq!(server_addr.to_string(), "nats://127.0.0.1:4321");
306    }
307
308    #[test]
309    fn nats_ipv6() {
310        let server_addr = "nats://[::1]".parse::<ServerAddr>().unwrap();
311        assert_eq!(server_addr.transport(), Transport::TCP);
312        assert_eq!(server_addr.protocol(), Protocol::PossiblyPlain);
313        assert_eq!(
314            server_addr.host(),
315            &Host::Ip(IpAddr::V6(Ipv6Addr::LOCALHOST))
316        );
317        assert_eq!(server_addr.port(), 4222);
318        assert_eq!(server_addr.username(), None);
319        assert_eq!(server_addr.password(), None);
320        assert_eq!(server_addr.to_string(), "nats://[::1]");
321    }
322
323    #[test]
324    fn tls() {
325        let server_addr = "tls://127.0.0.1".parse::<ServerAddr>().unwrap();
326        assert_eq!(server_addr.transport(), Transport::TCP);
327        assert_eq!(server_addr.protocol(), Protocol::TLS);
328        assert_eq!(
329            server_addr.host(),
330            &Host::Ip(IpAddr::V4(Ipv4Addr::LOCALHOST))
331        );
332        assert_eq!(server_addr.port(), 4222);
333        assert_eq!(server_addr.username(), None);
334        assert_eq!(server_addr.password(), None);
335        assert_eq!(server_addr.to_string(), "tls://127.0.0.1");
336    }
337
338    #[test]
339    fn ws() {
340        let server_addr = "ws://127.0.0.1".parse::<ServerAddr>().unwrap();
341        assert_eq!(server_addr.transport(), Transport::Websocket);
342        assert_eq!(server_addr.protocol(), Protocol::PossiblyPlain);
343        assert_eq!(
344            server_addr.host(),
345            &Host::Ip(IpAddr::V4(Ipv4Addr::LOCALHOST))
346        );
347        assert_eq!(server_addr.port(), 80);
348        assert_eq!(server_addr.username(), None);
349        assert_eq!(server_addr.password(), None);
350        assert_eq!(server_addr.to_string(), "ws://127.0.0.1");
351    }
352
353    #[test]
354    fn wss() {
355        let server_addr = "wss://127.0.0.1".parse::<ServerAddr>().unwrap();
356        assert_eq!(server_addr.transport(), Transport::Websocket);
357        assert_eq!(server_addr.protocol(), Protocol::TLS);
358        assert_eq!(
359            server_addr.host(),
360            &Host::Ip(IpAddr::V4(Ipv4Addr::LOCALHOST))
361        );
362        assert_eq!(server_addr.port(), 443);
363        assert_eq!(server_addr.username(), None);
364        assert_eq!(server_addr.password(), None);
365        assert_eq!(server_addr.to_string(), "wss://127.0.0.1");
366    }
367}