Skip to main content

ddns/
resolvers.rs

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