Skip to main content

reqwest/dns/
resolve.rs

1use hyper_util::client::legacy::connect::dns::Name as HyperName;
2use tower_service::Service;
3
4use std::collections::HashMap;
5use std::future::Future;
6use std::net::SocketAddr;
7use std::pin::Pin;
8use std::str::FromStr;
9use std::sync::Arc;
10use std::task::{Context, Poll};
11
12use crate::error::BoxError;
13
14/// Alias for an `Iterator` trait object over `SocketAddr`.
15pub type Addrs = Box<dyn Iterator<Item = SocketAddr> + Send>;
16
17/// Alias for the `Future` type returned by a DNS resolver.
18pub type Resolving = Pin<Box<dyn Future<Output = Result<Addrs, BoxError>> + Send>>;
19
20/// Trait for customizing DNS resolution in reqwest.
21pub trait Resolve: Send + Sync {
22    /// Performs DNS resolution on a `Name`.
23    /// The return type is a future containing an iterator of `SocketAddr`.
24    ///
25    /// It differs from `tower_service::Service<Name>` in several ways:
26    ///  * It is assumed that `resolve` will always be ready to poll.
27    ///  * It does not need a mutable reference to `self`.
28    ///  * Since trait objects cannot make use of associated types, it requires
29    ///    wrapping the returned `Future` and its contained `Iterator` with `Box`.
30    ///
31    /// Explicitly specified port in the URL will override any port in the resolved `SocketAddr`s.
32    /// Otherwise, port `0` will be replaced by the conventional port for the given scheme (e.g. 80 for http).
33    fn resolve(&self, name: Name) -> Resolving;
34}
35
36/// A name that must be resolved to addresses.
37#[derive(Debug)]
38pub struct Name(pub(super) HyperName);
39
40/// A more general trait implemented for types implementing `Resolve`.
41///
42/// Unnameable, only exported to aid seeing what implements this.
43pub trait IntoResolve {
44    #[doc(hidden)]
45    fn into_resolve(self) -> Arc<dyn Resolve>;
46}
47
48impl Name {
49    /// View the name as a string.
50    pub fn as_str(&self) -> &str {
51        self.0.as_str()
52    }
53}
54
55impl FromStr for Name {
56    type Err = sealed::InvalidNameError;
57
58    fn from_str(host: &str) -> Result<Self, Self::Err> {
59        HyperName::from_str(host)
60            .map(Name)
61            .map_err(|_| sealed::InvalidNameError { _ext: () })
62    }
63}
64
65#[derive(Clone)]
66pub(crate) struct DynResolver {
67    resolver: Arc<dyn Resolve>,
68}
69
70impl DynResolver {
71    pub(crate) fn new(resolver: Arc<dyn Resolve>) -> Self {
72        Self { resolver }
73    }
74
75    #[cfg(feature = "socks")]
76    pub(crate) fn gai() -> Self {
77        Self::new(Arc::new(super::gai::GaiResolver::new()))
78    }
79
80    /// Resolve an HTTP host and port, not just a domain name.
81    ///
82    /// This does the same thing that hyper-util's HttpConnector does, before
83    /// calling out to its underlying DNS resolver.
84    #[cfg(feature = "socks")]
85    pub(crate) async fn http_resolve(
86        &self,
87        target: &http::Uri,
88    ) -> Result<impl Iterator<Item = std::net::SocketAddr>, BoxError> {
89        let host = target.host().ok_or("missing host")?;
90        let port = target
91            .port_u16()
92            .unwrap_or_else(|| match target.scheme_str() {
93                Some("https") => 443,
94                Some("socks4") | Some("socks4a") | Some("socks5") | Some("socks5h") => 1080,
95                _ => 80,
96            });
97
98        let explicit_port = target.port().is_some();
99
100        let addrs = self
101            .resolver
102            .resolve(host.parse()?)
103            .await
104            .map_err(crate::error::dns)?;
105
106        Ok(addrs.map(move |mut addr| {
107            if explicit_port || addr.port() == 0 {
108                addr.set_port(port);
109            }
110            addr
111        }))
112    }
113}
114
115impl Service<HyperName> for DynResolver {
116    type Response = Addrs;
117    type Error = BoxError;
118    type Future = Resolving;
119
120    fn poll_ready(&mut self, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
121        Poll::Ready(Ok(()))
122    }
123
124    fn call(&mut self, name: HyperName) -> Self::Future {
125        let resolving = self.resolver.resolve(Name(name));
126        // Tag resolution failures so `Error::is_dns` can recognize them once
127        Box::pin(async move { resolving.await.map_err(crate::error::dns) })
128    }
129}
130
131pub(crate) struct DnsResolverWithOverrides {
132    dns_resolver: Arc<dyn Resolve>,
133    overrides: Arc<HashMap<String, Vec<SocketAddr>>>,
134}
135
136impl DnsResolverWithOverrides {
137    pub(crate) fn new(
138        dns_resolver: Arc<dyn Resolve>,
139        overrides: HashMap<String, Vec<SocketAddr>>,
140    ) -> Self {
141        DnsResolverWithOverrides {
142            dns_resolver,
143            overrides: Arc::new(overrides),
144        }
145    }
146}
147
148impl Resolve for DnsResolverWithOverrides {
149    fn resolve(&self, name: Name) -> Resolving {
150        match self.overrides.get(name.as_str()) {
151            Some(dest) => {
152                let addrs: Addrs = Box::new(dest.clone().into_iter());
153                Box::pin(std::future::ready(Ok(addrs)))
154            }
155            None => self.dns_resolver.resolve(name),
156        }
157    }
158}
159
160impl IntoResolve for Arc<dyn Resolve> {
161    fn into_resolve(self) -> Arc<dyn Resolve> {
162        self
163    }
164}
165
166impl<R> IntoResolve for Arc<R>
167where
168    R: Resolve + 'static,
169{
170    fn into_resolve(self) -> Arc<dyn Resolve> {
171        self
172    }
173}
174
175impl<R> IntoResolve for R
176where
177    R: Resolve + 'static,
178{
179    fn into_resolve(self) -> Arc<dyn Resolve> {
180        Arc::new(self)
181    }
182}
183
184mod sealed {
185    use std::fmt;
186
187    #[derive(Debug)]
188    pub struct InvalidNameError {
189        pub(super) _ext: (),
190    }
191
192    impl fmt::Display for InvalidNameError {
193        fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
194            f.write_str("invalid DNS name")
195        }
196    }
197
198    impl std::error::Error for InvalidNameError {}
199}