hyper_util/client/legacy/connect/
dns.rs1use std::error::Error;
24use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6, ToSocketAddrs};
25use std::pin::Pin;
26use std::str::FromStr;
27use std::task::{self, Poll};
28use std::{fmt, io, vec};
29
30use tokio::task::JoinHandle;
31use tower_service::Service;
32
33pub(super) use self::sealed::Resolve;
34
35#[derive(Clone, Hash, Eq, PartialEq)]
37pub struct Name {
38 host: Box<str>,
39}
40
41#[derive(Clone)]
43pub struct GaiResolver {
44 _priv: (),
45}
46
47pub struct GaiAddrs {
49 inner: SocketAddrs,
50}
51
52pub struct GaiFuture {
54 inner: JoinHandle<Result<SocketAddrs, io::Error>>,
55}
56
57impl Name {
58 pub(super) fn new(host: Box<str>) -> Name {
59 Name { host }
60 }
61
62 pub fn as_str(&self) -> &str {
64 &self.host
65 }
66}
67
68impl fmt::Debug for Name {
69 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
70 fmt::Debug::fmt(&self.host, f)
71 }
72}
73
74impl fmt::Display for Name {
75 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
76 fmt::Display::fmt(&self.host, f)
77 }
78}
79
80impl FromStr for Name {
81 type Err = InvalidNameError;
82
83 fn from_str(host: &str) -> Result<Self, Self::Err> {
84 Ok(Name::new(host.into()))
86 }
87}
88
89#[derive(Debug)]
91pub struct InvalidNameError(());
92
93impl fmt::Display for InvalidNameError {
94 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
95 f.write_str("Not a valid domain name")
96 }
97}
98
99impl Error for InvalidNameError {}
100
101impl GaiResolver {
102 pub fn new() -> Self {
104 GaiResolver { _priv: () }
105 }
106}
107
108impl Service<Name> for GaiResolver {
109 type Response = GaiAddrs;
110 type Error = io::Error;
111 type Future = GaiFuture;
112
113 fn poll_ready(&mut self, _cx: &mut task::Context<'_>) -> Poll<Result<(), io::Error>> {
114 Poll::Ready(Ok(()))
115 }
116
117 fn call(&mut self, name: Name) -> Self::Future {
118 let blocking = tokio::task::spawn_blocking(move || {
119 (&*name.host, 0)
120 .to_socket_addrs()
121 .map(|i| SocketAddrs { iter: i })
122 });
123
124 GaiFuture { inner: blocking }
125 }
126}
127
128impl fmt::Debug for GaiResolver {
129 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
130 f.pad("GaiResolver")
131 }
132}
133
134impl Future for GaiFuture {
135 type Output = Result<GaiAddrs, io::Error>;
136
137 fn poll(mut self: Pin<&mut Self>, cx: &mut task::Context<'_>) -> Poll<Self::Output> {
138 Pin::new(&mut self.inner).poll(cx).map(|res| match res {
139 Ok(Ok(addrs)) => Ok(GaiAddrs { inner: addrs }),
140 Ok(Err(err)) => Err(err),
141 Err(join_err) => {
142 if join_err.is_cancelled() {
143 Err(io::Error::new(io::ErrorKind::Interrupted, join_err))
144 } else {
145 panic!("gai background task failed: {join_err:?}")
146 }
147 }
148 })
149 }
150}
151
152impl fmt::Debug for GaiFuture {
153 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
154 f.pad("GaiFuture")
155 }
156}
157
158impl Drop for GaiFuture {
159 fn drop(&mut self) {
160 self.inner.abort();
161 }
162}
163
164impl Iterator for GaiAddrs {
165 type Item = SocketAddr;
166
167 fn next(&mut self) -> Option<Self::Item> {
168 self.inner.next()
169 }
170}
171
172impl fmt::Debug for GaiAddrs {
173 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
174 f.pad("GaiAddrs")
175 }
176}
177
178pub(super) struct SocketAddrs {
179 iter: vec::IntoIter<SocketAddr>,
180}
181
182impl SocketAddrs {
183 pub(super) fn new(addrs: Vec<SocketAddr>) -> Self {
184 SocketAddrs {
185 iter: addrs.into_iter(),
186 }
187 }
188
189 pub(super) fn try_parse(host: &str, port: u16) -> Option<SocketAddrs> {
190 if let Ok(addr) = host.parse::<Ipv4Addr>() {
191 let addr = SocketAddrV4::new(addr, port);
192 return Some(SocketAddrs {
193 iter: vec![SocketAddr::V4(addr)].into_iter(),
194 });
195 }
196 if let Ok(addr) = host.parse::<Ipv6Addr>() {
197 let addr = SocketAddrV6::new(addr, port, 0, 0);
198 return Some(SocketAddrs {
199 iter: vec![SocketAddr::V6(addr)].into_iter(),
200 });
201 }
202 None
203 }
204
205 #[inline]
206 fn filter(self, predicate: impl FnMut(&SocketAddr) -> bool) -> SocketAddrs {
207 SocketAddrs::new(self.iter.filter(predicate).collect())
208 }
209
210 pub(super) fn split_by_preference(
211 self,
212 local_addr_ipv4: Option<Ipv4Addr>,
213 local_addr_ipv6: Option<Ipv6Addr>,
214 ) -> (SocketAddrs, SocketAddrs) {
215 match (local_addr_ipv4, local_addr_ipv6) {
216 (Some(_), None) => (self.filter(SocketAddr::is_ipv4), SocketAddrs::new(vec![])),
217 (None, Some(_)) => (self.filter(SocketAddr::is_ipv6), SocketAddrs::new(vec![])),
218 _ => {
219 let preferring_v6 = self
220 .iter
221 .as_slice()
222 .first()
223 .map(SocketAddr::is_ipv6)
224 .unwrap_or(false);
225
226 let (preferred, fallback) = self
227 .iter
228 .partition::<Vec<_>, _>(|addr| addr.is_ipv6() == preferring_v6);
229
230 (SocketAddrs::new(preferred), SocketAddrs::new(fallback))
231 }
232 }
233 }
234
235 pub(super) fn is_empty(&self) -> bool {
236 self.iter.as_slice().is_empty()
237 }
238
239 pub(super) fn len(&self) -> usize {
240 self.iter.as_slice().len()
241 }
242}
243
244impl Iterator for SocketAddrs {
245 type Item = SocketAddr;
246 #[inline]
247 fn next(&mut self) -> Option<SocketAddr> {
248 self.iter.next()
249 }
250}
251
252mod sealed {
253 use std::task::{self, Poll};
254
255 use super::{Name, SocketAddr};
256 use tower_service::Service;
257
258 pub trait Resolve {
260 type Addrs: Iterator<Item = SocketAddr>;
261 type Error: Into<Box<dyn std::error::Error + Send + Sync>>;
262 type Future: Future<Output = Result<Self::Addrs, Self::Error>>;
263
264 fn poll_ready(&mut self, cx: &mut task::Context<'_>) -> Poll<Result<(), Self::Error>>;
265 fn resolve(&mut self, name: Name) -> Self::Future;
266 }
267
268 impl<S> Resolve for S
269 where
270 S: Service<Name>,
271 S::Response: Iterator<Item = SocketAddr>,
272 S::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
273 {
274 type Addrs = S::Response;
275 type Error = S::Error;
276 type Future = S::Future;
277
278 fn poll_ready(&mut self, cx: &mut task::Context<'_>) -> Poll<Result<(), Self::Error>> {
279 Service::poll_ready(self, cx)
280 }
281
282 fn resolve(&mut self, name: Name) -> Self::Future {
283 Service::call(self, name)
284 }
285 }
286}
287
288pub(super) async fn resolve<R>(resolver: &mut R, name: Name) -> Result<R::Addrs, R::Error>
289where
290 R: Resolve,
291{
292 std::future::poll_fn(|cx| resolver.poll_ready(cx)).await?;
293 resolver.resolve(name).await
294}
295
296#[cfg(test)]
297mod tests {
298 use super::*;
299
300 #[test]
301 fn test_ip_addrs_split_by_preference() {
302 let ip_v4 = Ipv4Addr::new(127, 0, 0, 1);
303 let ip_v6 = Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 1);
304 let v4_addr = (ip_v4, 80).into();
305 let v6_addr = (ip_v6, 80).into();
306
307 let (mut preferred, mut fallback) = SocketAddrs {
308 iter: vec![v4_addr, v6_addr].into_iter(),
309 }
310 .split_by_preference(None, None);
311 assert!(preferred.next().unwrap().is_ipv4());
312 assert!(fallback.next().unwrap().is_ipv6());
313
314 let (mut preferred, mut fallback) = SocketAddrs {
315 iter: vec![v6_addr, v4_addr].into_iter(),
316 }
317 .split_by_preference(None, None);
318 assert!(preferred.next().unwrap().is_ipv6());
319 assert!(fallback.next().unwrap().is_ipv4());
320
321 let (mut preferred, mut fallback) = SocketAddrs {
322 iter: vec![v4_addr, v6_addr].into_iter(),
323 }
324 .split_by_preference(Some(ip_v4), Some(ip_v6));
325 assert!(preferred.next().unwrap().is_ipv4());
326 assert!(fallback.next().unwrap().is_ipv6());
327
328 let (mut preferred, mut fallback) = SocketAddrs {
329 iter: vec![v6_addr, v4_addr].into_iter(),
330 }
331 .split_by_preference(Some(ip_v4), Some(ip_v6));
332 assert!(preferred.next().unwrap().is_ipv6());
333 assert!(fallback.next().unwrap().is_ipv4());
334
335 let (mut preferred, fallback) = SocketAddrs {
336 iter: vec![v4_addr, v6_addr].into_iter(),
337 }
338 .split_by_preference(Some(ip_v4), None);
339 assert!(preferred.next().unwrap().is_ipv4());
340 assert!(fallback.is_empty());
341
342 let (mut preferred, fallback) = SocketAddrs {
343 iter: vec![v4_addr, v6_addr].into_iter(),
344 }
345 .split_by_preference(None, Some(ip_v6));
346 assert!(preferred.next().unwrap().is_ipv6());
347 assert!(fallback.is_empty());
348 }
349
350 #[test]
351 fn test_name_from_str() {
352 const DOMAIN: &str = "test.example.com";
353 let name = Name::from_str(DOMAIN).expect("Should be a valid domain");
354 assert_eq!(name.as_str(), DOMAIN);
355 assert_eq!(name.to_string(), DOMAIN);
356 }
357}