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#[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#[derive(Debug, Copy, Clone, PartialEq, Eq, PartialOrd, Ord)]
26pub enum Protocol {
27 PossiblyPlain,
36 TLS,
42}
43
44#[derive(Debug, Copy, Clone, PartialEq, Eq, PartialOrd, Ord)]
46pub enum Transport {
47 TCP,
49 Websocket,
51}
52
53#[derive(Debug, Clone, PartialEq, Eq)]
55pub enum Host {
56 Ip(IpAddr),
58 Dns(ByteString),
60}
61
62impl ServerAddr {
63 pub fn protocol(&self) -> Protocol {
65 self.protocol
66 }
67
68 pub fn transport(&self) -> Transport {
70 self.transport
71 }
72
73 pub fn host(&self) -> &Host {
75 &self.host
76 }
77
78 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 pub fn username(&self) -> Option<&str> {
89 if self.username.is_empty() {
90 None
91 } else {
92 Some(&self.username)
93 }
94 }
95
96 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 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#[derive(Debug, thiserror::Error)]
245pub enum ServerAddrError {
246 #[error("invalid Url")]
248 InvalidUrl(#[source] url::ParseError),
249 #[error("invalid Url scheme")]
251 InvalidScheme,
252 #[error("missing host")]
254 MissingHost,
255 #[error("username is not utf-8")]
257 UsernameInvalidUtf8,
258 #[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}