Skip to main content

pb_mapper_core/
addr.rs

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        // This uses a helper method because clippy doesn't like the `to_vec()`
91        // call here (it will allocate, whereas `self.iter().copied()` would
92        // not), but it's actually necessary in order to ensure that the
93        // returned iterator is valid for the `'static` lifetime, which the
94        // borrowed `slice::Iter` iterator would not be.
95        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/// copy from tokio::net::addr
113#[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        // First check if the input parses as a socket address
158        let res: Result<SocketAddr, _> = self.parse();
159        if let Ok(addr) = res {
160            return MaybeReady(State::Ready(Some(addr)));
161        }
162
163        // Run DNS lookup on the blocking pool
164        let s = self.to_owned();
165
166        MaybeReady(State::Blocking(tokio::task::spawn_blocking(move || {
167            // Customized dns resolvers are preferred, if a custom resolver does not exist then the
168            // standard library's
169            get_socket_addrs(&s).map(|v| v.into_iter())
170        })))
171    }
172}
173
174/// Implement this trait for &T of type !Sized(such as str), since &T of type Sized all implement it
175/// by default.
176impl<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        // try to parse the host as a regular IP address first
196        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
227// ===== impl String =====
228
229impl 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
238/// Customized dns resolution server
239const DEFAULT_DNS_SERVER_GROUP: &[IpAddr] = &[
240    IpAddr::V4(Ipv4Addr::new(223, 5, 5, 5)), // alibaba
241    IpAddr::V4(Ipv4Addr::new(223, 6, 6, 6)),
242    IpAddr::V4(Ipv4Addr::new(119, 29, 29, 29)), // tencent
243    IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8)),      // google
244    IpAddr::V6(Ipv6Addr::new(0x2001, 0x4860, 0x4860, 0, 0, 0, 0, 0x8888)), // google
245];
246
247#[inline]
248fn custom_resolver_config() -> ResolverConfig {
249    // `udp_and_tcp` uses the standard DNS port and trusts negative responses,
250    // matching what this passed explicitly before.
251    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/// The custom resolver, or `None` if it could not be built.
263///
264/// Built once: `build()` can fail, and retrying per lookup would repeat the
265/// same failure. Callers fall back to the system resolver.
266#[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/// Blocking DNS lookup, via the system resolver.
301///
302/// The custom DNS servers are only reachable from the async helpers below.
303/// hickory has no blocking resolver, and the sync path never reached the
304/// custom one in practice anyway: it was skipped inside a Tokio runtime, and
305/// outside one this fell back to `std` whenever it was unavailable.
306#[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/// Blocking DNS lookup. Avoid calling this from inside a Tokio runtime thread.
312#[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
319/// Async DNS lookup using the custom resolver, safe to call inside Tokio runtimes.
320pub 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/// Async DNS lookup for host + port, falls back to system resolver on failure.
330#[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        // Resolve dns properly with the system resolver
338        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/// Async DNS lookup for `domain:port` forms, such as `bilibili.com:1080`.
346#[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}