1use std::collections::{VecDeque, vec_deque};
2use std::{fmt, iter::FusedIterator, net::SocketAddr};
3
4use ntex_bytes::ByteString;
5use ntex_util::future::Either;
6
7pub trait Address: Unpin + 'static {
9 fn host(&self) -> &str;
11
12 fn port(&self) -> Option<u16>;
14
15 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#[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 #[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 #[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 #[must_use]
102 pub fn set_port(mut self, port: u16) -> Self {
103 self.port = port;
104 self
105 }
106
107 #[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 #[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 pub fn host(&self) -> &str {
133 self.req.host()
134 }
135
136 pub fn port(&self) -> u16 {
139 self.req.port().unwrap_or(self.port)
140 }
141
142 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 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 pub fn get_ref(&self) -> &T {
178 &self.req
179 }
180
181 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#[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#[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
287pub(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 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}