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
27pub(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
65pub const DHTTP_H3_DNS_SERVER: &str = crate::bootstrap::DHTTP_H3_DNS_SERVER;
67
68pub const DHTTP_HTTP_DNS_SERVER: &str = crate::bootstrap::DHTTP_HTTP_DNS_SERVER;
70
71pub 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}