Skip to main content

ddns/
resolvers.rs

1#[cfg(feature = "resolvers")]
2use std::{
3    error::Error,
4    fmt::{self, Display},
5    sync::Arc,
6};
7
8#[cfg(feature = "resolvers")]
9use dquic::{
10    qbase::net::addr::EndpointAddr,
11    qresolve::{Resolve, ResolveFuture, Source},
12};
13#[cfg(feature = "resolvers")]
14use futures::{FutureExt, Stream, StreamExt, TryFutureExt, stream};
15#[cfg(feature = "resolvers")]
16use tokio::io;
17
18#[cfg(feature = "h3")]
19pub use crate::h3::H3Resolver;
20#[cfg(feature = "http")]
21pub use crate::http::HttpResolver;
22#[cfg(feature = "mdns")]
23pub use crate::mdns::MdnsResolver;
24#[cfg(all(feature = "mdns", feature = "dquic-network", feature = "resolvers"))]
25use crate::mdns::MdnsResolvers;
26
27/// Extract and validate the DNS host from `name`, which may include a `:port`
28/// suffix. Returns `Some(host)` if the host part is a valid RFC-compliant DNS
29/// name, or `None` for raw IP addresses, bracketed IPv6, or malformed input.
30#[cfg_attr(not(any(feature = "h3", feature = "http")), allow(dead_code))]
31pub(crate) fn resolvable_name(name: &str) -> Option<&str> {
32    let host = match name.rsplit_once(':') {
33        Some((h, port)) if !port.is_empty() && port.chars().all(|c| c.is_ascii_digit()) => h,
34        _ => name,
35    };
36    rustls::pki_types::DnsName::try_from(host).ok()?;
37    Some(host)
38}
39
40/// Default DNS-over-H3 server for DHTTP endpoints.
41pub const DHTTP_H3_DNS_SERVER: &str = crate::bootstrap::DHTTP_H3_DNS_SERVER;
42
43/// Default DNS-over-HTTP server for DHTTP endpoints.
44pub const DHTTP_HTTP_DNS_SERVER: &str = crate::bootstrap::DHTTP_HTTP_DNS_SERVER;
45
46/// mDNS service type used by DHTTP endpoints.
47pub const DHTTP_MDNS_SERVICE: &str = crate::bootstrap::DHTTP_MDNS_SERVICE;
48
49#[cfg(feature = "resolvers")]
50#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
51pub enum DnsScheme {
52    Mdns,
53    Http,
54    H3,
55    System,
56}
57
58#[cfg(feature = "resolvers")]
59impl Display for DnsScheme {
60    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
61        f.write_str(match self {
62            Self::Mdns => "mdns",
63            Self::Http => "http",
64            Self::H3 => "h3",
65            Self::System => "system",
66        })
67    }
68}
69
70#[cfg(feature = "resolvers")]
71#[derive(Debug, snafu::Snafu)]
72#[snafu(display("unsupported dns scheme {scheme}"))]
73pub struct ParseDnsSchemeError {
74    scheme: String,
75}
76
77#[cfg(feature = "resolvers")]
78impl std::str::FromStr for DnsScheme {
79    type Err = ParseDnsSchemeError;
80
81    fn from_str(s: &str) -> Result<Self, Self::Err> {
82        match s {
83            "mdns" => Ok(Self::Mdns),
84            "http" => Ok(Self::Http),
85            "h3" => Ok(Self::H3),
86            "system" => Ok(Self::System),
87            scheme => Err(ParseDnsSchemeError {
88                scheme: scheme.to_owned(),
89            }),
90        }
91    }
92}
93
94pub mod deferred;
95#[cfg(any(feature = "h3", feature = "mdns", test))]
96pub(crate) mod endpoint_group;
97pub mod weak;
98
99#[cfg(feature = "resolvers")]
100type ArcResolver = Arc<dyn Resolve + Send + Sync + 'static>;
101
102#[cfg(feature = "resolvers")]
103#[derive(Default, Clone, Debug)]
104pub struct Resolvers {
105    resolvers: Vec<ArcResolver>,
106}
107
108#[cfg(feature = "resolvers")]
109impl Display for Resolvers {
110    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
111        f.write_str("Resolvers(")?;
112        if self.resolvers.is_empty() {
113            f.write_str("empty")?;
114        } else {
115            for (i, resolver) in self.resolvers.iter().enumerate() {
116                if i > 0 {
117                    f.write_str(", ")?;
118                }
119                fmt::Display::fmt(resolver.as_ref(), f)?;
120            }
121        }
122        f.write_str(")")
123    }
124}
125
126#[cfg(feature = "resolvers")]
127#[derive(Debug)]
128pub struct ResolversError {
129    errors: Vec<(String, io::Error)>,
130}
131
132#[cfg(feature = "resolvers")]
133fn format_dns_error_sources(
134    f: &mut fmt::Formatter<'_>,
135    error: &(dyn Error + 'static),
136) -> fmt::Result {
137    let mut index = 1;
138    let mut current = error.source();
139
140    while let Some(source) = current {
141        write!(f, "\n    {index}. {source}")?;
142        index += 1;
143        current = source.source();
144    }
145
146    Ok(())
147}
148
149#[cfg(feature = "resolvers")]
150fn format_dns_error_entry(
151    f: &mut fmt::Formatter<'_>,
152    resolver: &str,
153    error: &io::Error,
154) -> fmt::Result {
155    write!(f, "\n  - {resolver}: {error}")?;
156    format_dns_error_sources(f, error)
157}
158
159#[cfg(feature = "resolvers")]
160impl fmt::Display for ResolversError {
161    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
162        if self.errors.is_empty() {
163            return write!(f, "no DNS resolvers available");
164        }
165
166        write!(f, "all DNS resolvers failed")?;
167        for (resolver, error) in &self.errors {
168            format_dns_error_entry(f, resolver, error)?;
169        }
170        Ok(())
171    }
172}
173
174#[cfg(feature = "resolvers")]
175impl Error for ResolversError {}
176
177#[cfg(feature = "resolvers")]
178#[derive(Default)]
179pub struct ResolversBuilder {
180    resolvers: Resolvers,
181}
182
183#[cfg(feature = "resolvers")]
184impl ResolversBuilder {
185    pub fn resolver(mut self, resolver: ArcResolver) -> Self {
186        self.resolvers.push(resolver);
187        self
188    }
189
190    #[cfg(all(feature = "mdns", feature = "dquic-network"))]
191    pub async fn mdns(
192        mut self,
193        network: Arc<h3x::dquic::Network>,
194        patterns: Arc<Vec<h3x::dquic::binds::BindPattern>>,
195    ) -> Self {
196        let mdns: ArcResolver =
197            Arc::new(MdnsResolvers::bind(network, patterns, DHTTP_MDNS_SERVICE).await);
198        self.resolvers.push(mdns);
199        self
200    }
201
202    #[cfg(feature = "h3")]
203    pub fn h3<C>(
204        self,
205        endpoint: Arc<h3x::endpoint::H3Endpoint<C, C::Connection>>,
206    ) -> io::Result<Self>
207    where
208        C: h3x::quic::Connect + h3x::quic::WithLocalAuthority + Send + Sync + 'static,
209        C::Error: Send + Sync + 'static,
210        C::Connection: Send + 'static,
211    {
212        self.h3_with_base_url(DHTTP_H3_DNS_SERVER, endpoint)
213    }
214
215    #[cfg(feature = "h3")]
216    pub fn h3_with_base_url<C>(
217        mut self,
218        base_url: impl AsRef<str>,
219        endpoint: Arc<h3x::endpoint::H3Endpoint<C, C::Connection>>,
220    ) -> io::Result<Self>
221    where
222        C: h3x::quic::Connect + h3x::quic::WithLocalAuthority + Send + Sync + 'static,
223        C::Error: Send + Sync + 'static,
224        C::Connection: Send + 'static,
225    {
226        let resolver = H3Resolver::from_endpoint(base_url, endpoint)?;
227        self.resolvers.push(Arc::new(resolver));
228        Ok(self)
229    }
230
231    #[cfg(feature = "http")]
232    pub fn http(self) -> io::Result<Self> {
233        self.http_with_base_url(DHTTP_HTTP_DNS_SERVER)
234    }
235
236    #[cfg(feature = "http")]
237    pub fn http_with_base_url(mut self, base_url: impl AsRef<str>) -> io::Result<Self> {
238        let resolver = HttpResolver::new(base_url.as_ref())?;
239        self.resolvers.push(Arc::new(resolver));
240        Ok(self)
241    }
242
243    pub fn system(mut self) -> Self {
244        self.resolvers
245            .push(Arc::new(dquic::qresolve::SystemResolver));
246        self
247    }
248
249    pub fn build(self) -> Resolvers {
250        self.resolvers
251    }
252}
253
254#[cfg(feature = "resolvers")]
255impl Resolvers {
256    pub fn builder() -> ResolversBuilder {
257        ResolversBuilder::default()
258    }
259
260    pub fn new() -> Self {
261        Self::default()
262    }
263
264    pub fn with(mut self, resolver: ArcResolver) -> Self {
265        self.push(resolver);
266        self
267    }
268
269    pub fn push(&mut self, resolver: ArcResolver) {
270        self.resolvers.push(resolver);
271    }
272
273    pub fn iter(&self) -> impl Iterator<Item = &ArcResolver> {
274        self.resolvers.iter()
275    }
276
277    pub async fn lookup(
278        &self,
279        name: &str,
280    ) -> Result<impl Stream<Item = (Source, EndpointAddr)> + use<>, ResolversError> {
281        let mut errors = vec![];
282
283        let mut lookups = stream::FuturesUnordered::from_iter(
284            (self.resolvers.clone().into_iter()).map(|resolver| {
285                let resolver = resolver.clone();
286                let name = name.to_string();
287                async move { (resolver.lookup(&name).await, resolver.clone()) }
288            }),
289        );
290
291        let endpoints = loop {
292            match lookups.next().await {
293                Some((Ok(endpoints), _)) => break endpoints,
294                Some((Err(error), resolver)) => errors.push((resolver.to_string(), error)),
295                None => return Err(ResolversError { errors }),
296            }
297        };
298
299        Ok(endpoints.chain(lookups.flat_map(|(endpoints, _)| stream::iter(endpoints).flatten())))
300    }
301}
302
303#[cfg(feature = "resolvers")]
304impl Resolve for Resolvers {
305    fn lookup<'l>(&'l self, name: &'l str) -> ResolveFuture<'l> {
306        self.lookup(name)
307            .map_ok(StreamExt::boxed)
308            .map_err(io::Error::other)
309            .boxed()
310    }
311}
312
313#[cfg(test)]
314mod tests {
315    #[cfg(all(feature = "mdns", feature = "dquic-network", feature = "resolvers"))]
316    use std::str::FromStr;
317    #[cfg(feature = "resolvers")]
318    use std::{error::Error as StdError, fmt, io};
319
320    #[cfg(all(feature = "mdns", feature = "dquic-network", feature = "resolvers"))]
321    use super::MdnsResolvers;
322    #[cfg(feature = "resolvers")]
323    use super::Resolvers;
324    use super::{DHTTP_H3_DNS_SERVER, DHTTP_HTTP_DNS_SERVER, DHTTP_MDNS_SERVICE, resolvable_name};
325    #[cfg(feature = "resolvers")]
326    use super::{DnsScheme, ResolversError};
327
328    #[cfg(feature = "resolvers")]
329    #[derive(Debug)]
330    struct TestSourceError {
331        message: &'static str,
332        source: Option<Box<TestSourceError>>,
333    }
334
335    #[cfg(feature = "resolvers")]
336    impl TestSourceError {
337        fn leaf(message: &'static str) -> Self {
338            Self {
339                message,
340                source: None,
341            }
342        }
343
344        fn with_source(message: &'static str, source: TestSourceError) -> Self {
345            Self {
346                message,
347                source: Some(Box::new(source)),
348            }
349        }
350    }
351
352    #[cfg(feature = "resolvers")]
353    impl fmt::Display for TestSourceError {
354        fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
355            f.write_str(self.message)
356        }
357    }
358
359    #[cfg(feature = "resolvers")]
360    impl StdError for TestSourceError {
361        fn source(&self) -> Option<&(dyn StdError + 'static)> {
362            self.source
363                .as_deref()
364                .map(|source| source as &(dyn StdError + 'static))
365        }
366    }
367
368    #[cfg(feature = "resolvers")]
369    fn other_error(message: &'static str) -> io::Error {
370        io::Error::other(message)
371    }
372
373    #[cfg(feature = "resolvers")]
374    fn chained_other_error(root: TestSourceError) -> io::Error {
375        io::Error::other(root)
376    }
377
378    #[test]
379    fn resolver_defaults_come_from_compile_time_environment() {
380        if let Some(expected) = option_env!("DHTTP_H3_DNS_SERVER") {
381            assert_eq!(DHTTP_H3_DNS_SERVER, expected);
382        }
383        if let Some(expected) = option_env!("DHTTP_HTTP_DNS_SERVER") {
384            assert_eq!(DHTTP_HTTP_DNS_SERVER, expected);
385        }
386        if let Some(expected) = option_env!("DHTTP_MDNS_SERVICE") {
387            assert_eq!(DHTTP_MDNS_SERVICE, expected);
388        }
389    }
390
391    #[test]
392    fn resolvable_name_accepts_dns_name_with_numeric_port() {
393        assert_eq!(
394            resolvable_name("example.dhttp.net:443"),
395            Some("example.dhttp.net")
396        );
397    }
398
399    #[test]
400    fn resolvable_name_accepts_stun_authority_with_numeric_port() {
401        assert_eq!(
402            resolvable_name("nat.genmeta.net:20004"),
403            Some("nat.genmeta.net")
404        );
405    }
406
407    #[test]
408    fn resolvable_name_rejects_ip_literals() {
409        assert_eq!(resolvable_name("127.0.0.1:443"), None);
410        assert_eq!(resolvable_name("[::1]:443"), None);
411    }
412
413    #[cfg(feature = "resolvers")]
414    #[test]
415    fn dns_scheme_round_trips_supported_schemes_and_rejects_dht() {
416        let cases = [
417            ("mdns", DnsScheme::Mdns),
418            ("http", DnsScheme::Http),
419            ("h3", DnsScheme::H3),
420            ("system", DnsScheme::System),
421        ];
422
423        for (text, scheme) in cases {
424            assert_eq!(DnsScheme::from_str(text).expect("supported scheme"), scheme);
425            assert_eq!(scheme.to_string(), text);
426        }
427
428        assert!(DnsScheme::from_str("dht").is_err());
429    }
430
431    #[cfg(feature = "resolvers")]
432    #[test]
433    fn resolvers_error_renders_no_resolvers_available_when_empty() {
434        let error = ResolversError { errors: vec![] };
435
436        assert_eq!(error.to_string(), "no DNS resolvers available");
437    }
438
439    #[cfg(feature = "resolvers")]
440    #[test]
441    fn resolvers_error_renders_resolver_bullets_in_stored_order() {
442        let error = ResolversError {
443            errors: vec![
444                (
445                    "System DNS Resolver".to_string(),
446                    other_error("invalid socket address"),
447                ),
448                ("mDNS resolvers".to_string(), other_error("timed out")),
449            ],
450        };
451
452        assert_eq!(
453            error.to_string(),
454            concat!(
455                "all DNS resolvers failed\n",
456                "  - System DNS Resolver: invalid socket address\n",
457                "  - mDNS resolvers: timed out"
458            )
459        );
460    }
461
462    #[cfg(feature = "resolvers")]
463    #[test]
464    fn resolvers_error_renders_numbered_source_chain_for_one_resolver() {
465        let error = ResolversError {
466            errors: vec![(
467                "DeferredResolver(H3 DNS Resolver(https://dns.genmeta.net:4433/))".to_string(),
468                chained_other_error(TestSourceError::with_source(
469                    "deferred resolver lookup failed",
470                    TestSourceError::leaf("no DNS record found"),
471                )),
472            )],
473        };
474
475        assert_eq!(
476            error.to_string(),
477            concat!(
478                "all DNS resolvers failed\n",
479                "  - DeferredResolver(H3 DNS Resolver(https://dns.genmeta.net:4433/)): deferred resolver lookup failed\n",
480                "    1. no DNS record found"
481            )
482        );
483    }
484
485    #[cfg(feature = "resolvers")]
486    #[test]
487    fn resolvers_error_renders_repeated_source_messages_without_deduplication() {
488        let error = ResolversError {
489            errors: vec![(
490                "DeferredResolver(H3 DNS Resolver(https://dns.genmeta.net:4433/))".to_string(),
491                chained_other_error(TestSourceError::with_source(
492                    "deferred resolver lookup failed",
493                    TestSourceError::with_source(
494                        "deferred resolver lookup failed",
495                        TestSourceError::leaf("no DNS record found"),
496                    ),
497                )),
498            )],
499        };
500
501        assert_eq!(
502            error.to_string(),
503            concat!(
504                "all DNS resolvers failed\n",
505                "  - DeferredResolver(H3 DNS Resolver(https://dns.genmeta.net:4433/)): deferred resolver lookup failed\n",
506                "    1. deferred resolver lookup failed\n",
507                "    2. no DNS record found"
508            )
509        );
510    }
511
512    #[cfg(all(feature = "mdns", feature = "dquic-network", feature = "resolvers"))]
513    #[tokio::test]
514    async fn resolvers_builder_can_enable_mdns() {
515        use std::sync::Arc;
516
517        use h3x::dquic::{Network, binds::BindPattern};
518
519        let network = Network::builder().build();
520        let pattern = BindPattern::from_str("iface://v4.lo:0").expect("valid pattern");
521
522        let resolvers = Resolvers::builder()
523            .mdns(network, Arc::new(vec![pattern]))
524            .await
525            .build();
526
527        assert!(resolvers.to_string().contains("mDNS resolvers"));
528    }
529
530    #[cfg(all(feature = "h3", feature = "resolvers", feature = "dquic-network"))]
531    #[tokio::test]
532    async fn resolvers_builder_accepts_custom_h3_base_url() {
533        use std::sync::Arc;
534
535        let endpoint = Arc::new(h3x::endpoint::H3Endpoint::new(
536            h3x::dquic::QuicEndpoint::builder().build().await,
537        ));
538
539        let resolvers = Resolvers::builder()
540            .h3_with_base_url("https://custom-dns.example:4433", endpoint)
541            .expect("valid h3 dns url")
542            .build();
543
544        assert!(resolvers.to_string().contains("custom-dns.example"));
545    }
546
547    #[cfg(all(feature = "http", feature = "resolvers"))]
548    #[test]
549    fn resolvers_builder_accepts_custom_http_base_url() {
550        let resolvers = Resolvers::builder()
551            .http_with_base_url("https://custom-dns.example")
552            .expect("valid http dns url")
553            .build();
554
555        assert!(resolvers.to_string().contains("custom-dns.example"));
556    }
557
558    #[cfg(all(feature = "mdns", feature = "dquic-network", feature = "resolvers"))]
559    #[tokio::test]
560    async fn mdns_resolvers_bind_installs_mdns_on_null_io_binding() {
561        use std::sync::Arc;
562
563        use dquic::qinterface::io::IO;
564        use h3x::dquic::{Network, binds::BindPattern};
565
566        let network = Network::builder().build();
567        let pattern = BindPattern::from_str("iface://v4.lo:0").expect("valid pattern");
568        let resolvers = MdnsResolvers::bind(
569            network.clone(),
570            Arc::new(vec![pattern.clone()]),
571            DHTTP_MDNS_SERVICE,
572        )
573        .await;
574
575        let ifaces = resolvers
576            .bound_interfaces(&pattern)
577            .expect("bound interfaces");
578        if ifaces.is_empty() {
579            return;
580        }
581        assert!(ifaces[0].borrow().bound_addr().is_err());
582        assert!(
583            ifaces[0]
584                .with_components(|components, _| components.exist::<crate::mdns::service::Mdns>())
585        );
586    }
587}