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