Skip to main content

ntex_net/connect/
message.rs

1use std::collections::{VecDeque, vec_deque};
2use std::{fmt, iter::FusedIterator, net::SocketAddr};
3
4use ntex_bytes::ByteString;
5use ntex_util::future::Either;
6
7/// Address information required by [`Connect`].
8pub trait Address: Unpin + 'static {
9    /// Returns the host name.
10    fn host(&self) -> &str;
11
12    /// Returns an explicitly configured port, if available.
13    fn port(&self) -> Option<u16>;
14
15    /// Returns a pre-resolved socket address, if available.
16    fn addr(&self) -> Option<SocketAddr> {
17        None
18    }
19}
20
21impl Address for String {
22    fn host(&self) -> &str {
23        self
24    }
25
26    fn port(&self) -> Option<u16> {
27        None
28    }
29}
30
31impl Address for ByteString {
32    fn host(&self) -> &str {
33        self
34    }
35
36    fn port(&self) -> Option<u16> {
37        None
38    }
39}
40
41impl Address for &'static str {
42    fn host(&self) -> &str {
43        self
44    }
45
46    fn port(&self) -> Option<u16> {
47        None
48    }
49}
50
51impl Address for SocketAddr {
52    fn host(&self) -> &'static str {
53        ""
54    }
55
56    fn port(&self) -> Option<u16> {
57        None
58    }
59
60    fn addr(&self) -> Option<SocketAddr> {
61        Some(*self)
62    }
63}
64
65/// Request to resolve and connect to a remote address.
66#[derive(Eq, PartialEq, Debug, Hash)]
67pub struct Connect<T> {
68    pub(super) req: T,
69    pub(super) port: u16,
70    pub(super) addr: Option<Either<SocketAddr, VecDeque<SocketAddr>>>,
71}
72
73impl<T: Address> Connect<T> {
74    /// Creates a connection request and derives a port from the host when present.
75    #[must_use]
76    pub fn new(req: T) -> Connect<T> {
77        let (_, port) = parse(req.host());
78        Connect {
79            req,
80            port: port.unwrap_or(0),
81            addr: None,
82        }
83    }
84
85    /// Creates a request with a pre-resolved socket address.
86    ///
87    /// The connector skips DNS resolution for this request.
88    #[must_use]
89    pub fn with(req: T, addr: SocketAddr) -> Connect<T> {
90        Connect {
91            req,
92            port: 0,
93            addr: Some(Either::Left(addr)),
94        }
95    }
96
97    /// Sets the port used when [`Address::port()`] does not provide one.
98    ///
99    /// This replaces the port parsed from a `host:port` host by
100    /// [`new()`](Self::new), which is zero if the host has none.
101    #[must_use]
102    pub fn set_port(mut self, port: u16) -> Self {
103        self.port = port;
104        self
105    }
106
107    /// Sets one pre-resolved socket address.
108    #[must_use]
109    pub fn set_addr(mut self, addr: Option<SocketAddr>) -> Self {
110        if let Some(addr) = addr {
111            self.addr = Some(Either::Left(addr));
112        }
113        self
114    }
115
116    /// Sets multiple pre-resolved socket addresses.
117    #[must_use]
118    pub fn set_addrs<I>(mut self, addrs: I) -> Self
119    where
120        I: IntoIterator<Item = SocketAddr>,
121    {
122        let mut addrs = VecDeque::from_iter(addrs);
123        self.addr = if addrs.len() < 2 {
124            addrs.pop_front().map(Either::Left)
125        } else {
126            Some(Either::Right(addrs))
127        };
128        self
129    }
130
131    /// Returns the request host name.
132    pub fn host(&self) -> &str {
133        self.req.host()
134    }
135
136    /// Returns the port from [`Address::port()`], or the one parsed from the
137    /// host or set by [`set_port()`](Self::set_port).
138    pub fn port(&self) -> u16 {
139        self.req.port().unwrap_or(self.port)
140    }
141
142    /// Iterates over the request's pre-resolved addresses.
143    pub fn addrs(&self) -> ConnectAddrsIter<'_> {
144        if let Some(addr) = self.req.addr() {
145            ConnectAddrsIter {
146                inner: Either::Left(Some(addr)),
147            }
148        } else {
149            let inner = match self.addr {
150                None => Either::Left(None),
151                Some(Either::Left(addr)) => Either::Left(Some(addr)),
152                Some(Either::Right(ref addrs)) => Either::Right(addrs.iter()),
153            };
154
155            ConnectAddrsIter { inner }
156        }
157    }
158
159    /// Removes and returns the request's pre-resolved addresses.
160    pub fn take_addrs(&mut self) -> ConnectTakeAddrsIter {
161        if let Some(addr) = self.req.addr() {
162            ConnectTakeAddrsIter {
163                inner: Either::Left(Some(addr)),
164            }
165        } else {
166            let inner = match self.addr.take() {
167                None => Either::Left(None),
168                Some(Either::Left(addr)) => Either::Left(Some(addr)),
169                Some(Either::Right(addrs)) => Either::Right(addrs.into_iter()),
170            };
171
172            ConnectTakeAddrsIter { inner }
173        }
174    }
175
176    /// Returns the original address value.
177    pub fn get_ref(&self) -> &T {
178        &self.req
179    }
180
181    /// Maps the original address value while preserving port and resolved addresses.
182    pub fn map_addr<F, R>(self, f: F) -> Connect<R>
183    where
184        F: FnOnce(T) -> R,
185    {
186        let req = f(self.req);
187
188        Connect {
189            req,
190            port: self.port,
191            addr: self.addr,
192        }
193    }
194}
195
196impl<T: Clone> Clone for Connect<T> {
197    fn clone(&self) -> Self {
198        Connect {
199            req: self.req.clone(),
200            port: self.port,
201            addr: self.addr.clone(),
202        }
203    }
204}
205
206impl<T: Address> From<T> for Connect<T> {
207    fn from(addr: T) -> Self {
208        Connect::new(addr)
209    }
210}
211
212impl<T: Address> fmt::Display for Connect<T> {
213    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
214        let (host, _) = parse(self.host());
215        if host.contains(':') {
216            write!(f, "[{host}]:{}", self.port())
217        } else {
218            write!(f, "{host}:{}", self.port())
219        }
220    }
221}
222
223/// Iterator over addresses in a [`Connect`] request.
224#[derive(Clone)]
225pub struct ConnectAddrsIter<'a> {
226    inner: Either<Option<SocketAddr>, vec_deque::Iter<'a, SocketAddr>>,
227}
228
229impl Iterator for ConnectAddrsIter<'_> {
230    type Item = SocketAddr;
231
232    fn next(&mut self) -> Option<Self::Item> {
233        match self.inner {
234            Either::Left(ref mut opt) => opt.take(),
235            Either::Right(ref mut iter) => iter.next().copied(),
236        }
237    }
238
239    fn size_hint(&self) -> (usize, Option<usize>) {
240        match self.inner {
241            Either::Left(Some(_)) => (1, Some(1)),
242            Either::Left(None) => (0, Some(0)),
243            Either::Right(ref iter) => iter.size_hint(),
244        }
245    }
246}
247
248impl fmt::Debug for ConnectAddrsIter<'_> {
249    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
250        f.debug_list().entries(self.clone()).finish()
251    }
252}
253
254impl ExactSizeIterator for ConnectAddrsIter<'_> {}
255
256impl FusedIterator for ConnectAddrsIter<'_> {}
257
258/// Owning iterator over addresses removed from a [`Connect`] request.
259#[derive(Debug)]
260pub struct ConnectTakeAddrsIter {
261    inner: Either<Option<SocketAddr>, vec_deque::IntoIter<SocketAddr>>,
262}
263
264impl Iterator for ConnectTakeAddrsIter {
265    type Item = SocketAddr;
266
267    fn next(&mut self) -> Option<Self::Item> {
268        match self.inner {
269            Either::Left(ref mut opt) => opt.take(),
270            Either::Right(ref mut iter) => iter.next(),
271        }
272    }
273
274    fn size_hint(&self) -> (usize, Option<usize>) {
275        match self.inner {
276            Either::Left(Some(_)) => (1, Some(1)),
277            Either::Left(None) => (0, Some(0)),
278            Either::Right(ref iter) => iter.size_hint(),
279        }
280    }
281}
282
283impl ExactSizeIterator for ConnectTakeAddrsIter {}
284
285impl FusedIterator for ConnectTakeAddrsIter {}
286
287/// Splits `host`, `host:port`, `[v6]`, `[v6]:port` or bare `v6` into host and port.
288///
289/// Brackets are stripped from IPv6 hosts.
290pub(super) fn parse(host: &str) -> (&str, Option<u16>) {
291    let (name, port) = if let Some(rest) = host.strip_prefix('[') {
292        match rest.split_once(']') {
293            Some((ip, tail)) => (ip, tail.strip_prefix(':')),
294            None => (host, None),
295        }
296    } else {
297        match host.split_once(':') {
298            Some((name, port)) if !port.contains(':') => (name, Some(port)),
299            _ => (host, None),
300        }
301    };
302    (name, port.and_then(|p| p.parse::<u16>().ok()))
303}
304
305#[cfg(test)]
306mod tests {
307    use super::*;
308
309    #[test]
310    fn address() {
311        assert_eq!("test".host(), "test");
312        assert_eq!("test".port(), None);
313
314        let s = "test".to_string();
315        assert_eq!(s.host(), "test");
316        assert_eq!(s.port(), None);
317
318        let s = ByteString::from("test");
319        assert_eq!(s.host(), "test");
320        assert_eq!(s.port(), None);
321    }
322
323    #[test]
324    fn parse_host() {
325        assert_eq!(parse("example.com"), ("example.com", None));
326        assert_eq!(parse("example.com:443"), ("example.com", Some(443)));
327        assert_eq!(parse("example.com:bad"), ("example.com", None));
328        assert_eq!(parse("127.0.0.1:8080"), ("127.0.0.1", Some(8080)));
329        assert_eq!(parse("[::1]"), ("::1", None));
330        assert_eq!(parse("[::1]:443"), ("::1", Some(443)));
331        assert_eq!(parse("[::1]443"), ("::1", None));
332        assert_eq!(parse("::1"), ("::1", None));
333        assert_eq!(parse("2001:db8::1"), ("2001:db8::1", None));
334        assert_eq!(parse("[::1"), ("[::1", None));
335        assert_eq!(parse(""), ("", None));
336
337        assert_eq!(Connect::new("[::1]:8080").port(), 8080);
338        assert_eq!(Connect::new("::1").set_port(80).port(), 80);
339    }
340
341    #[test]
342    fn display() {
343        assert_eq!(
344            Connect::new("example.com:443").to_string(),
345            "example.com:443"
346        );
347        assert_eq!(
348            Connect::new("example.com").set_port(80).to_string(),
349            "example.com:80"
350        );
351        assert_eq!(Connect::new("[::1]:8080").to_string(), "[::1]:8080");
352        assert_eq!(Connect::new("::1").set_port(80).to_string(), "[::1]:80");
353        assert_eq!(
354            Connect::new("fe80::1%3").set_port(80).to_string(),
355            "[fe80::1%3]:80"
356        );
357    }
358
359    #[test]
360    #[allow(clippy::similar_names)]
361    fn connect() {
362        let mut connect = Connect::new("www.rust-lang.org");
363        assert_eq!(connect.host(), "www.rust-lang.org");
364        assert_eq!(connect.port(), 0);
365        assert_eq!(*connect.get_ref(), "www.rust-lang.org");
366        connect = connect.set_port(80);
367        assert_eq!(connect.port(), 80);
368        let addrs = connect.addrs().clone();
369        assert_eq!(format!("{addrs:?}"), "[]");
370        assert!(connect.addrs().next().is_none());
371        assert!(format!("{:?}", connect.clone()).contains("Connect"));
372
373        let c = connect.clone().map_addr(|_| "www.rust-lang.org:80");
374        assert_eq!(c.host(), "www.rust-lang.org:80");
375        assert_eq!(c.port(), 80);
376        let addrs = c.addrs().clone();
377        assert_eq!(format!("{addrs:?}"), "[]");
378        assert!(c.addrs().next().is_none());
379
380        let addr: SocketAddr = "127.0.0.1:8080".parse().unwrap();
381        connect = connect.set_addrs(vec![addr]);
382        let addrs = connect.addrs().clone();
383        assert_eq!(format!("{addrs:?}"), "[127.0.0.1:8080]");
384        let addrs: Vec<_> = connect.take_addrs().collect();
385        assert_eq!(addrs.len(), 1);
386        assert!(addrs.contains(&addr));
387
388        let addr2: SocketAddr = "127.0.0.1:8081".parse().unwrap();
389        connect = connect.set_addrs(vec![addr, addr2]);
390        let addrs: Vec<_> = connect.addrs().collect();
391        assert_eq!(addrs.len(), 2);
392        assert!(addrs.contains(&addr));
393        assert!(addrs.contains(&addr2));
394
395        let addrs: Vec<_> = connect.take_addrs().collect();
396        assert_eq!(addrs.len(), 2);
397        assert!(addrs.contains(&addr));
398        assert!(addrs.contains(&addr2));
399        assert!(connect.addrs().next().is_none());
400
401        connect = connect.set_addrs(vec![addr]);
402        assert_eq!(format!("{connect}"), "www.rust-lang.org:80");
403
404        let addr: SocketAddr = "127.0.0.1:8080".parse().unwrap();
405        let mut connect = Connect::new(addr);
406        assert_eq!(connect.host(), "");
407        assert_eq!(connect.port(), 0);
408        let addrs: Vec<_> = connect.addrs().collect();
409        assert_eq!(addrs.len(), 1);
410        assert!(addrs.contains(&addr));
411        let addrs: Vec<_> = connect.take_addrs().collect();
412        assert_eq!(addrs.len(), 1);
413        assert!(addrs.contains(&addr));
414    }
415
416    #[test]
417    fn connect_with_addr() {
418        let addr: SocketAddr = "127.0.0.1:8080".parse().unwrap();
419        let connect = Connect::with("example.com", addr);
420        assert_eq!(connect.port(), 0);
421        assert_eq!(connect.addrs().len(), 1);
422        assert_eq!(connect.addrs().next(), Some(addr));
423
424        let connect: Connect<&str> = "example.com:80".into();
425        assert_eq!(connect.port(), 80);
426        assert_eq!(connect.addrs().len(), 0);
427
428        // `None` keeps the current addresses
429        let connect = connect.set_addr(None);
430        assert_eq!(connect.addrs().len(), 0);
431        let mut connect = connect.set_addr(Some(addr)).set_addr(None);
432        assert_eq!(connect.addrs().collect::<Vec<_>>(), vec![addr]);
433
434        let mut it = connect.take_addrs();
435        assert_eq!(it.len(), 1);
436        assert_eq!(it.next(), Some(addr));
437        assert_eq!(it.len(), 0);
438        assert_eq!(connect.take_addrs().len(), 0);
439
440        let mut connect = Connect::new(addr);
441        assert_eq!(connect.addrs().len(), 1);
442        assert_eq!(connect.take_addrs().len(), 1);
443
444        let addr2: SocketAddr = "127.0.0.1:8081".parse().unwrap();
445        let mut connect = Connect::new("example.com").set_addrs([addr, addr2]);
446        assert_eq!(connect.addrs().len(), 2);
447        let mut it = connect.take_addrs();
448        assert_eq!(it.len(), 2);
449        it.next();
450        assert_eq!(it.len(), 1);
451        assert!(format!("{it:?}").contains("8081"));
452    }
453}