1use std::future::{self, Future};
2use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6};
3use std::pin::Pin;
4use std::sync::LazyLock;
5use std::task::{Context, Poll, ready};
6
7use hickory_resolver::config::{NameServerConfig, ResolverConfig, ResolverOpts};
8use hickory_resolver::net::runtime::TokioRuntimeProvider;
9use hickory_resolver::{Resolver, TokioResolver};
10use tokio::task::JoinHandle;
11
12type Result<T, E = std::io::Error> = std::result::Result<T, E>;
13type ReadyFuture<T> = future::Ready<Result<T>>;
14
15pub trait ToSocketAddrs {
16 type Iter: Iterator<Item = SocketAddr> + Send + 'static;
17 type Future: Future<Output = Result<Self::Iter>> + Send + 'static;
18
19 fn to_socket_addrs(&self) -> Self::Future;
20}
21
22impl ToSocketAddrs for SocketAddr {
23 type Future = ReadyFuture<Self::Iter>;
24 type Iter = std::option::IntoIter<SocketAddr>;
25
26 fn to_socket_addrs(&self) -> Self::Future {
27 let iter = Some(*self).into_iter();
28 future::ready(Ok(iter))
29 }
30}
31
32impl ToSocketAddrs for SocketAddrV4 {
33 type Future = ReadyFuture<Self::Iter>;
34 type Iter = std::option::IntoIter<SocketAddr>;
35
36 fn to_socket_addrs(&self) -> Self::Future {
37 SocketAddr::V4(*self).to_socket_addrs()
38 }
39}
40
41impl ToSocketAddrs for SocketAddrV6 {
42 type Future = ReadyFuture<Self::Iter>;
43 type Iter = std::option::IntoIter<SocketAddr>;
44
45 fn to_socket_addrs(&self) -> Self::Future {
46 SocketAddr::V6(*self).to_socket_addrs()
47 }
48}
49
50impl ToSocketAddrs for (IpAddr, u16) {
51 type Future = ReadyFuture<Self::Iter>;
52 type Iter = std::option::IntoIter<SocketAddr>;
53
54 fn to_socket_addrs(&self) -> Self::Future {
55 let iter = Some(SocketAddr::from(*self)).into_iter();
56 future::ready(Ok(iter))
57 }
58}
59
60impl ToSocketAddrs for (Ipv4Addr, u16) {
61 type Future = ReadyFuture<Self::Iter>;
62 type Iter = std::option::IntoIter<SocketAddr>;
63
64 fn to_socket_addrs(&self) -> Self::Future {
65 let (ip, port) = *self;
66 SocketAddrV4::new(ip, port).to_socket_addrs()
67 }
68}
69
70impl ToSocketAddrs for (Ipv6Addr, u16) {
71 type Future = ReadyFuture<Self::Iter>;
72 type Iter = std::option::IntoIter<SocketAddr>;
73
74 fn to_socket_addrs(&self) -> Self::Future {
75 let (ip, port) = *self;
76 SocketAddrV6::new(ip, port, 0, 0).to_socket_addrs()
77 }
78}
79
80impl ToSocketAddrs for &[SocketAddr] {
81 type Future = ReadyFuture<Self::Iter>;
82 type Iter = std::vec::IntoIter<SocketAddr>;
83
84 fn to_socket_addrs(&self) -> Self::Future {
85 #[inline]
86 fn slice_to_vec(addrs: &[SocketAddr]) -> Vec<SocketAddr> {
87 addrs.to_vec()
88 }
89
90 let iter = slice_to_vec(self).into_iter();
96 future::ready(Ok(iter))
97 }
98}
99
100#[derive(Debug)]
101pub enum OneOrMore {
102 One(std::option::IntoIter<SocketAddr>),
103 More(std::vec::IntoIter<SocketAddr>),
104}
105
106#[derive(Debug)]
107enum State {
108 Ready(Option<SocketAddr>),
109 Blocking(JoinHandle<Result<std::vec::IntoIter<SocketAddr>>>),
110}
111
112#[derive(Debug)]
114pub struct MaybeReady(State);
115
116impl Future for MaybeReady {
117 type Output = Result<OneOrMore>;
118
119 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
120 match self.0 {
121 State::Ready(ref mut i) => {
122 let iter = OneOrMore::One(i.take().into_iter());
123 Poll::Ready(Ok(iter))
124 }
125 State::Blocking(ref mut rx) => {
126 let res = ready!(Pin::new(rx).poll(cx))?.map(OneOrMore::More);
127
128 Poll::Ready(res)
129 }
130 }
131 }
132}
133
134impl Iterator for OneOrMore {
135 type Item = SocketAddr;
136
137 fn next(&mut self) -> Option<Self::Item> {
138 match self {
139 OneOrMore::One(i) => i.next(),
140 OneOrMore::More(i) => i.next(),
141 }
142 }
143
144 fn size_hint(&self) -> (usize, Option<usize>) {
145 match self {
146 OneOrMore::One(i) => i.size_hint(),
147 OneOrMore::More(i) => i.size_hint(),
148 }
149 }
150}
151
152impl ToSocketAddrs for str {
153 type Future = MaybeReady;
154 type Iter = OneOrMore;
155
156 fn to_socket_addrs(&self) -> Self::Future {
157 let res: Result<SocketAddr, _> = self.parse();
159 if let Ok(addr) = res {
160 return MaybeReady(State::Ready(Some(addr)));
161 }
162
163 let s = self.to_owned();
165
166 MaybeReady(State::Blocking(tokio::task::spawn_blocking(move || {
167 get_socket_addrs(&s).map(|v| v.into_iter())
170 })))
171 }
172}
173
174impl<T> ToSocketAddrs for &T
177where
178 T: ToSocketAddrs + ?Sized,
179{
180 type Future = T::Future;
181 type Iter = T::Iter;
182
183 fn to_socket_addrs(&self) -> Self::Future {
184 (**self).to_socket_addrs()
185 }
186}
187
188impl ToSocketAddrs for (&str, u16) {
189 type Future = MaybeReady;
190 type Iter = OneOrMore;
191
192 fn to_socket_addrs(&self) -> Self::Future {
193 let (host, port) = *self;
194
195 if let Ok(addr) = host.parse::<Ipv4Addr>() {
197 let addr = SocketAddrV4::new(addr, port);
198 let addr = SocketAddr::V4(addr);
199
200 return MaybeReady(State::Ready(Some(addr)));
201 }
202
203 if let Ok(addr) = host.parse::<Ipv6Addr>() {
204 let addr = SocketAddrV6::new(addr, port, 0, 0);
205 let addr = SocketAddr::V6(addr);
206
207 return MaybeReady(State::Ready(Some(addr)));
208 }
209
210 let host = host.to_owned();
211
212 MaybeReady(State::Blocking(tokio::task::spawn_blocking(move || {
213 get_socket_addrs_from_host_port(&host, port).map(|v| v.into_iter())
214 })))
215 }
216}
217
218impl ToSocketAddrs for (String, u16) {
219 type Future = MaybeReady;
220 type Iter = OneOrMore;
221
222 fn to_socket_addrs(&self) -> Self::Future {
223 (self.0.as_str(), self.1).to_socket_addrs()
224 }
225}
226
227impl ToSocketAddrs for String {
230 type Future = <str as ToSocketAddrs>::Future;
231 type Iter = <str as ToSocketAddrs>::Iter;
232
233 fn to_socket_addrs(&self) -> Self::Future {
234 self[..].to_socket_addrs()
235 }
236}
237
238const DEFAULT_DNS_SERVER_GROUP: &[IpAddr] = &[
240 IpAddr::V4(Ipv4Addr::new(223, 5, 5, 5)), IpAddr::V4(Ipv4Addr::new(223, 6, 6, 6)),
242 IpAddr::V4(Ipv4Addr::new(119, 29, 29, 29)), IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8)), IpAddr::V6(Ipv6Addr::new(0x2001, 0x4860, 0x4860, 0, 0, 0, 0, 0x8888)), ];
246
247#[inline]
248fn custom_resolver_config() -> ResolverConfig {
249 ResolverConfig::from_parts(
252 None,
253 vec![],
254 DEFAULT_DNS_SERVER_GROUP
255 .iter()
256 .copied()
257 .map(NameServerConfig::udp_and_tcp)
258 .collect::<Vec<_>>(),
259 )
260}
261
262#[inline]
267fn get_custom_async_resolver() -> Option<TokioResolver> {
268 static RESOLVER: LazyLock<Option<TokioResolver>> = LazyLock::new(|| {
269 let mut builder = Resolver::builder_with_config(
270 custom_resolver_config(),
271 TokioRuntimeProvider::default(),
272 );
273 *builder.options_mut() = ResolverOpts::default();
274 match builder.build() {
275 Ok(resolver) => Some(resolver),
276 Err(e) => {
277 tracing::error!("Create custom dns resolver error:{e},falling back to system dns");
278 None
279 }
280 }
281 });
282 RESOLVER.clone()
283}
284
285macro_rules! invalid_input {
286 ($msg:expr) => {
287 std::io::Error::new(std::io::ErrorKind::InvalidInput, $msg)
288 };
289}
290
291macro_rules! try_opt {
292 ($call:expr, $msg:expr) => {
293 match $call {
294 Some(v) => v,
295 None => Err(invalid_input!($msg))?,
296 }
297 };
298}
299
300#[inline]
307pub fn get_socket_addrs_from_host_port(host: &str, port: u16) -> Result<Vec<SocketAddr>> {
308 std::net::ToSocketAddrs::to_socket_addrs(&(host, port)).map(|v| v.collect())
309}
310
311#[inline]
313pub fn get_socket_addrs(s: &str) -> Result<Vec<SocketAddr>> {
314 let (host, port_str) = try_opt!(s.rsplit_once(':'), "invalid socket address");
315 let port: u16 = try_opt!(port_str.parse().ok(), "invalid port value");
316 get_socket_addrs_from_host_port(host, port)
317}
318
319pub async fn get_ip_addrs_async(s: &str) -> Result<Vec<IpAddr>> {
321 let resolver = try_opt!(get_custom_async_resolver(), "custom resolver not exist");
322 resolver
323 .lookup_ip(s)
324 .await
325 .map(|v| v.iter().collect())
326 .map_err(|e| invalid_input!(e))
327}
328
329#[inline]
331pub async fn get_socket_addrs_from_host_port_async(
332 host: &str,
333 port: u16,
334) -> Result<Vec<SocketAddr>> {
335 match get_ip_addrs_async(host).await {
336 Ok(r) => Ok(r.into_iter().map(|ip| SocketAddr::new(ip, port)).collect()),
337 Err(_) => tokio::net::lookup_host((host, port))
339 .await
340 .map(|v| v.collect())
341 .map_err(|e| invalid_input!(e)),
342 }
343}
344
345#[inline]
347pub async fn get_socket_addrs_async(s: &str) -> Result<Vec<SocketAddr>> {
348 let (host, port_str) = try_opt!(s.rsplit_once(':'), "invalid socket address");
349 let port: u16 = try_opt!(port_str.parse().ok(), "invalid port value");
350 get_socket_addrs_from_host_port_async(host, port).await
351}
352
353pub async fn each_addr<A: ToSocketAddrs, F, T, R>(addr: A, f: F) -> Result<T>
354where
355 F: Fn(SocketAddr) -> R,
356 R: std::future::Future<Output = Result<T>>,
357{
358 let addrs = match addr.to_socket_addrs().await {
359 Ok(addrs) => addrs,
360 Err(e) => return Err(e),
361 };
362 let mut last_err = None;
363 for addr in addrs {
364 match f(addr).await {
365 Ok(l) => return Ok(l),
366 Err(e) => last_err = Some(e),
367 }
368 }
369 Err(last_err.unwrap_or_else(|| {
370 std::io::Error::new(
371 std::io::ErrorKind::InvalidInput,
372 "could not resolve to any addresses",
373 )
374 }))
375}