From e21046a5fa1fcbb304003a5a8cdb88de858b8ccf Mon Sep 17 00:00:00 2001 From: fanyang Date: Wed, 3 Jun 2026 22:24:21 +0800 Subject: [PATCH] feat(dns): support configurable resolver chain --- Cargo.lock | 10 + easytier/Cargo.toml | 2 +- easytier/src/common/config.rs | 82 +++++- easytier/src/common/global_ctx.rs | 6 + easytier/src/common/stun.rs | 102 ++++++-- easytier/src/connector/dns_connector.rs | 2 +- easytier/src/utils/dns.rs | 333 ++++++++++++++++++++---- 7 files changed, 455 insertions(+), 82 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 88755425..5f79798c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3760,23 +3760,29 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8a6fe56c0038198998a6f217ca4e7ef3a5e51f46163bd6dd60b5c71ca6c6502" dependencies = [ "async-trait", + "bytes", "cfg-if", "data-encoding", "enum-as-inner", "futures-channel", "futures-io", "futures-util", + "h2", + "http", "idna 1.0.3", "ipnet", "once_cell", "rand 0.9.1", "ring", + "rustls", "serde", "thiserror 2.0.11", "tinyvec", "tokio", + "tokio-rustls", "tracing", "url", + "webpki-roots 0.26.3", ] [[package]] @@ -3794,11 +3800,14 @@ dependencies = [ "parking_lot", "rand 0.9.1", "resolv-conf", + "rustls", "serde", "smallvec", "thiserror 2.0.11", "tokio", + "tokio-rustls", "tracing", + "webpki-roots 0.26.3", ] [[package]] @@ -7728,6 +7737,7 @@ version = "0.23.27" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "730944ca083c1c233a75c09f199e973ca499344a2b7ba9e755c457e86fb4a321" dependencies = [ + "log", "once_cell", "ring", "rustls-pki-types", diff --git a/easytier/Cargo.toml b/easytier/Cargo.toml index b94a45ab..75807044 100644 --- a/easytier/Cargo.toml +++ b/easytier/Cargo.toml @@ -245,7 +245,7 @@ http_req = { git = "https://github.com/EasyTier/http_req.git", default-features # for dns connector hickory-proto = "0.26.0" hickory-net = { version = "0.26.0", features = ["serde"] } -hickory-resolver = "0.26.0" +hickory-resolver = { version = "0.26.0", features = ["https-ring", "webpki-roots"] } # for magic dns hickory-server = { version = "0.26.0", features = ["resolver"], optional = true } diff --git a/easytier/src/common/config.rs b/easytier/src/common/config.rs index 38b15eeb..f7eca94f 100644 --- a/easytier/src/common/config.rs +++ b/easytier/src/common/config.rs @@ -1,5 +1,5 @@ use super::env_parser; -use crate::utils::dns::sanitize; +use crate::utils::dns; use crate::{ common::stun::StunInfoCollector, proto::{ @@ -325,6 +325,14 @@ pub trait ConfigLoader: Send + Sync + DnsConfigLoaderExt { fn get_stun_servers_v6(&self) -> Option>; fn set_stun_servers_v6(&self, servers: Option>); + fn get_dns_resolvers(&self) -> Vec { + dns::get_default_dns_resolvers() + } + fn get_dns_resolvers_config(&self) -> Option> { + None + } + fn set_dns_resolvers(&self, _resolvers: Option>) {} + fn get_secure_mode(&self) -> Option; fn set_secure_mode(&self, secure_mode: Option); @@ -652,6 +660,7 @@ struct Config { udp_whitelist: Option>, stun_servers: Option>, stun_servers_v6: Option>, + dns_resolvers: Option>, credential_file: Option, source: Option, @@ -685,6 +694,10 @@ impl TomlConfigLoader { Self::normalize_config_source(&mut config); config.flags_struct = Some(Self::gen_flags(config.flags.clone().unwrap_or_default())); + if let Some(dns_resolvers) = &config.dns_resolvers { + dns::validate_dns_resolvers(dns_resolvers) + .with_context(|| "invalid dns_resolvers config")?; + } let config = TomlConfigLoader { config: Arc::new(Mutex::new(config)), @@ -766,7 +779,7 @@ impl ConfigLoader for TomlConfigLoader { .filter(|h| !h.is_empty()); self.set_hostname(hostname.clone()); - hostname.unwrap_or_else(|| sanitize(utils::hostname())) + hostname.unwrap_or_else(|| utils::dns::sanitize(utils::hostname())) } fn set_hostname(&self, name: Option) { @@ -1094,6 +1107,23 @@ impl ConfigLoader for TomlConfigLoader { self.config.lock().unwrap().stun_servers_v6 = servers; } + fn get_dns_resolvers(&self) -> Vec { + self.config + .lock() + .unwrap() + .dns_resolvers + .clone() + .unwrap_or_else(dns::get_default_dns_resolvers) + } + + fn get_dns_resolvers_config(&self) -> Option> { + self.config.lock().unwrap().dns_resolvers.clone() + } + + fn set_dns_resolvers(&self, resolvers: Option>) { + self.config.lock().unwrap().dns_resolvers = resolvers; + } + fn get_secure_mode(&self) -> Option { self.config.lock().unwrap().secure_mode.clone() } @@ -1156,6 +1186,9 @@ impl ConfigLoader for TomlConfigLoader { if config.stun_servers_v6 == Some(StunInfoCollector::get_default_servers_v6()) { config.stun_servers_v6 = None; } + if config.dns_resolvers == Some(dns::get_default_dns_resolvers()) { + config.dns_resolvers = None; + } toml::to_string_pretty(&config).unwrap() } } @@ -1387,6 +1420,51 @@ stun_servers = [ assert_eq!(stun_servers[2], "txt:stun.easytier.cn"); } + #[test] + fn test_dns_resolvers_default_and_roundtrip() { + let config = TomlConfigLoader::default(); + assert_eq!(config.get_dns_resolvers_config(), None); + assert_eq!(config.get_dns_resolvers(), vec!["system".to_string()]); + assert!(!config.dump().contains("dns_resolvers")); + + let config = TomlConfigLoader::new_from_str( + r#" +dns_resolvers = ["system", "https://dns.alidns.com/dns-query"] +"#, + ) + .unwrap(); + assert_eq!( + config.get_dns_resolvers_config().unwrap(), + vec![ + "system".to_string(), + "https://dns.alidns.com/dns-query".to_string() + ] + ); + assert_eq!( + config.get_dns_resolvers(), + vec![ + "system".to_string(), + "https://dns.alidns.com/dns-query".to_string() + ] + ); + + let dumped = config.dump(); + assert!(dumped.contains("dns_resolvers")); + let loaded = TomlConfigLoader::new_from_str(&dumped).unwrap(); + assert_eq!(loaded.get_dns_resolvers(), config.get_dns_resolvers()); + } + + #[test] + fn test_dns_resolvers_reject_unknown_doh_without_bootstrap() { + let err = TomlConfigLoader::new_from_str( + r#" +dns_resolvers = ["https://example.com/dns-query"] +"#, + ) + .unwrap_err(); + assert!(err.to_string().contains("invalid dns_resolvers")); + } + #[test] fn test_network_config_source_toml_roundtrip() { let config = TomlConfigLoader::default(); diff --git a/easytier/src/common/global_ctx.rs b/easytier/src/common/global_ctx.rs index 3791b9c4..3e5971fc 100644 --- a/easytier/src/common/global_ctx.rs +++ b/easytier/src/common/global_ctx.rs @@ -294,6 +294,12 @@ impl GlobalCtx { let (event_bus, _) = tokio::sync::broadcast::channel(16); + if let Some(dns_resolvers) = config_fs.get_dns_resolvers_config() + && let Err(e) = crate::utils::dns::set_dns_resolvers(dns_resolvers) + { + crate::common::log::warn!("failed to set dns resolvers: {:?}", e); + } + let stun_info_collector = StunInfoCollector::new_with_default_servers(); if let Some(stun_servers) = config_fs.get_stun_servers() { diff --git a/easytier/src/common/stun.rs b/easytier/src/common/stun.rs index 461594bf..b21550a2 100644 --- a/easytier/src/common/stun.rs +++ b/easytier/src/common/stun.rs @@ -11,7 +11,7 @@ use crossbeam::atomic::AtomicCell; use rand::seq::IteratorRandom; use socket2::{SockAddr, SockRef}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; -use tokio::net::{UdpSocket, lookup_host}; +use tokio::net::UdpSocket; use tokio::sync::{Mutex, broadcast}; use tokio::task::JoinSet; use tracing::{Instrument, Level}; @@ -20,10 +20,8 @@ use bytecodec::{DecodeExt, EncodeExt}; use stun_codec::rfc5389::methods::BINDING; use stun_codec::{Message, MessageClass, MessageDecoder, MessageEncoder}; -use crate::common::error::Error; - use super::stun_codec_ext::*; -use crate::utils::dns::txt_resolve; +use crate::utils::dns::{resolve_host, txt_resolve}; const DEFAULT_UDP_STUN_SERVERS: &[&str] = &[ "txt:stun.easytier.cn", @@ -61,6 +59,18 @@ impl HostResolverIter { } } + fn parse_ipv6_socket_addr_without_brackets(host: &str) -> Option { + if host.parse::().is_ok() { + return None; + } + + let (ip, port) = host.rsplit_once(':')?; + Some(SocketAddr::new( + IpAddr::V6(ip.parse().ok()?), + port.parse().ok()?, + )) + } + #[async_recursion::async_recursion] async fn next(&mut self) -> Option { if self.ips.is_empty() { @@ -69,11 +79,6 @@ impl HostResolverIter { } let host = self.hostnames.remove(0); - let host = if host.contains(':') { - host - } else { - format!("{}:3478", host) - }; if host.starts_with("txt:") { let domain_name = host.trim_start_matches("txt:"); @@ -99,22 +104,53 @@ impl HostResolverIter { } let use_ipv6 = self.use_ipv6; - - match lookup_host(&host).await { - Ok(ips) => { - self.ips = ips - .filter(|x| if use_ipv6 { x.is_ipv6() } else { x.is_ipv4() }) - .choose_multiple(&mut rand::thread_rng(), self.max_ip_per_domain as usize); - - if self.ips.is_empty() { - return self.next().await; - } + if let Ok(addr) = host.parse::() { + if (use_ipv6 && addr.is_ipv6()) || (!use_ipv6 && addr.is_ipv4()) { + self.ips = vec![addr]; } - Err(e) => { - tracing::warn!(?host, ?e, "lookup host for stun failed"); + if self.ips.is_empty() { return self.next().await; } - }; + } else if let Some(addr) = Self::parse_ipv6_socket_addr_without_brackets(&host) { + if use_ipv6 { + self.ips = vec![addr]; + } + if self.ips.is_empty() { + return self.next().await; + } + } else { + let (host, port) = if let Ok(ip) = host.parse::() { + (ip.to_string(), 3478) + } else if let Ok(url) = url::Url::parse(&format!("stun://{}", host)) { + let Some(parsed_host) = url.host_str() else { + tracing::warn!(?host, "parse stun host failed"); + return self.next().await; + }; + (parsed_host.to_string(), url.port().unwrap_or(3478)) + } else { + (host, 3478) + }; + + match resolve_host(&host, port).await { + Ok(ips) => { + self.ips = ips + .into_iter() + .filter(|x| if use_ipv6 { x.is_ipv6() } else { x.is_ipv4() }) + .choose_multiple( + &mut rand::thread_rng(), + self.max_ip_per_domain as usize, + ); + + if self.ips.is_empty() { + return self.next().await; + } + } + Err(e) => { + tracing::warn!(?host, ?e, "resolve host for stun failed"); + return self.next().await; + } + }; + } } Some(self.ips.remove(0)) @@ -1344,6 +1380,26 @@ mod tests { use super::*; + #[test] + fn parse_ipv6_socket_addr_without_brackets_rejects_plain_ipv6_literals() { + assert_eq!( + HostResolverIter::parse_ipv6_socket_addr_without_brackets("2001:db8::1"), + None + ); + assert_eq!( + HostResolverIter::parse_ipv6_socket_addr_without_brackets("2001:db8:0:0:0:0:0:1"), + None + ); + } + + #[test] + fn parse_ipv6_socket_addr_without_brackets_accepts_unambiguous_port() { + assert_eq!( + HostResolverIter::parse_ipv6_socket_addr_without_brackets("::1:55355"), + Some("[::1]:55355".parse().unwrap()) + ); + } + #[tokio::test] async fn test_udp_nat_type_detector() { let collector = StunInfoCollector::new( @@ -1558,6 +1614,6 @@ mod tests { }); let stun_servers = vec!["::1:55355".to_string()]; let ret = StunInfoCollector::get_public_ipv6(&stun_servers).await; - println!("{:#?}", ret); + assert_eq!(ret, Some(Ipv6Addr::LOCALHOST)); } } diff --git a/easytier/src/connector/dns_connector.rs b/easytier/src/connector/dns_connector.rs index ec1b3b4a..9e523ec5 100644 --- a/easytier/src/connector/dns_connector.rs +++ b/easytier/src/connector/dns_connector.rs @@ -9,7 +9,7 @@ use crate::{ }; use anyhow::Context; use dashmap::DashSet; -use hickory_resolver::proto::rr::rdata::SRV; +use hickory_proto::rr::rdata::SRV; use rand::{Rng as _, seq::SliceRandom}; use strum::VariantArray; diff --git a/easytier/src/utils/dns.rs b/easytier/src/utils/dns.rs index 14d3167a..0eb2b17d 100644 --- a/easytier/src/utils/dns.rs +++ b/easytier/src/utils/dns.rs @@ -10,11 +10,139 @@ use hickory_resolver::system_conf::read_system_conf; use hickory_resolver::{Resolver, TokioResolver}; use idna::AsciiDenyList; use once_cell::sync::Lazy; -use std::net::SocketAddr; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; use std::sync::Arc; use std::sync::atomic::AtomicBool; +use std::sync::RwLock; use tokio::net::lookup_host; +const SYSTEM_DNS_RESOLVER: &str = "system"; + +#[derive(Clone)] +enum DnsResolver { + System, + Hickory { + uri: String, + resolver: Arc>, + }, +} + +pub fn get_default_dns_resolvers() -> Vec { + vec![SYSTEM_DNS_RESOLVER.to_string()] +} + +fn bootstrap_ips_for_doh_host(host: &str) -> Option> { + if let Ok(ip) = host.parse::() { + return Some(vec![ip]); + } + + match host { + "dns.alidns.com" => Some(vec![ + IpAddr::V4(Ipv4Addr::new(223, 5, 5, 5)), + IpAddr::V4(Ipv4Addr::new(223, 6, 6, 6)), + IpAddr::V6(Ipv6Addr::new(0x2400, 0x3200, 0, 0, 0, 0, 0, 1)), + IpAddr::V6(Ipv6Addr::new(0x2400, 0x3200, 0xbaba, 0, 0, 0, 0, 1)), + ]), + _ => None, + } +} + +fn build_hickory_resolver( + config: ResolverConfig, +) -> Result>, Error> { + let mut opts = ResolverOpts::default(); + opts.ip_strategy = LookupIpStrategy::Ipv4AndIpv6; + Ok(Arc::new( + TokioResolver::builder_with_config(config, TokioRuntimeProvider::default()) + .with_options(opts) + .build()?, + )) +} + +fn build_doh_resolver(raw: &str) -> Result { + let url = url::Url::parse(raw).map_err(|e| Error::InvalidUrl(e.to_string()))?; + if url.scheme() != "https" { + return Err(anyhow::anyhow!("unsupported dns resolver scheme: {}", url.scheme()).into()); + } + + let host = url + .host_str() + .with_context(|| format!("DoH resolver host is empty: {}", raw))?; + let ips = bootstrap_ips_for_doh_host(host).with_context(|| { + format!( + "DoH resolver {} requires a known bootstrap IP; currently only dns.alidns.com or IP literals are supported", + raw + ) + })?; + let port = url.port().unwrap_or(443); + let server_name: Arc = Arc::from(host.to_string()); + let http_endpoint: Option> = match url.path() { + "" | "/" | "/dns-query" => None, + path => Some(Arc::from(path.to_string())), + }; + + let name_servers = ips + .into_iter() + .map(|ip| { + NameServerConfig::new( + SocketAddr::new(ip, port), + true, + vec![ConnectionConfig::https( + server_name.clone(), + http_endpoint.clone(), + )], + ) + }) + .collect::>(); + + Ok(DnsResolver::Hickory { + uri: raw.to_string(), + resolver: build_hickory_resolver(ResolverConfig::from_parts( + None, + Vec::new(), + name_servers, + ))?, + }) +} + +fn build_dns_resolver(raw: &str) -> Result { + if raw.eq_ignore_ascii_case(SYSTEM_DNS_RESOLVER) { + return Ok(DnsResolver::System); + } + + build_doh_resolver(raw) +} + +fn build_dns_resolvers(raw_resolvers: &[String]) -> Result, Error> { + let raw_resolvers = if raw_resolvers.is_empty() { + get_default_dns_resolvers() + } else { + raw_resolvers.to_vec() + }; + + raw_resolvers + .iter() + .map(|raw| build_dns_resolver(raw)) + .collect() +} + +pub fn validate_dns_resolvers(raw_resolvers: &[String]) -> Result<(), Error> { + build_dns_resolvers(raw_resolvers).map(|_| ()) +} + +pub fn set_dns_resolvers(raw_resolvers: Vec) -> Result<(), Error> { + let resolvers = build_dns_resolvers(&raw_resolvers)?; + *DNS_RESOLVERS.write().unwrap() = resolvers; + Ok(()) +} + +static DNS_RESOLVERS: Lazy>> = + Lazy::new(|| RwLock::new(build_dns_resolvers(&get_default_dns_resolvers()).unwrap())); + +fn configured_dns_resolvers() -> Vec { + DNS_RESOLVERS.read().unwrap().clone() +} + pub fn sanitize(name: impl AsRef) -> String { let name = name.as_ref(); let dot = name.ends_with('.'); @@ -91,23 +219,46 @@ static RESOLVER: Lazy>> = Lazy::new(|| { }); pub async fn txt_lookup(name: impl IntoName) -> Result, Error> { - let response = RESOLVER - .txt_lookup(name) - .await - .context("failed to lookup txt record")?; + let name = name.into_name().context("invalid txt record name")?; + let mut last_err = None; - let data = response - .answers() - .iter() - .filter_map(|record| match record.data { - RData::TXT(ref txt) => Some(txt.to_string()), - _ => None, - }) - .collect(); + for resolver in configured_dns_resolvers() { + let response = match resolver { + DnsResolver::System => RESOLVER.txt_lookup(name.clone()).await, + DnsResolver::Hickory { uri, resolver } => { + let response = resolver.txt_lookup(name.clone()).await; + if response.is_err() { + tracing::debug!(?uri, ?name, "txt lookup failed with resolver"); + } + response + } + }; - tracing::info!(?data, "got txt record(s)"); + let Ok(response) = response else { + last_err = Some(anyhow::anyhow!("failed to lookup txt record").into()); + continue; + }; - Ok(data) + let data = response + .answers() + .iter() + .filter_map(|record| match record.data { + RData::TXT(ref txt) => Some(txt.to_string()), + _ => None, + }) + .collect::>(); + + if data.is_empty() { + last_err = Some(Error::NotFound); + continue; + } + + tracing::info!(?data, "got txt record(s)"); + + return Ok(data); + } + + Err(last_err.unwrap_or(Error::NotFound)) } pub async fn txt_resolve(name: impl IntoName) -> Result, Error> { @@ -120,23 +271,101 @@ pub async fn txt_resolve(name: impl IntoName) -> Result, Error> { } pub async fn srv_lookup(name: impl IntoName) -> Result, Error> { - let response = RESOLVER - .srv_lookup(name) - .await - .context("failed to lookup srv record")?; + let name = name.into_name().context("invalid srv record name")?; + let mut last_err = None; - let data = response - .answers() - .iter() - .filter_map(|record| match record.data { - RData::SRV(ref srv) => Some(srv.clone()), - _ => None, - }) - .collect(); + for resolver in configured_dns_resolvers() { + let response = match resolver { + DnsResolver::System => RESOLVER.srv_lookup(name.clone()).await, + DnsResolver::Hickory { uri, resolver } => { + let response = resolver.srv_lookup(name.clone()).await; + if response.is_err() { + tracing::debug!(?uri, ?name, "srv lookup failed with resolver"); + } + response + } + }; - tracing::info!(?data, "got srv record(s)"); + let Ok(response) = response else { + last_err = Some(anyhow::anyhow!("failed to lookup srv record").into()); + continue; + }; - Ok(data) + let data = response + .answers() + .iter() + .filter_map(|record| match record.data { + RData::SRV(ref srv) => Some(srv.clone()), + _ => None, + }) + .collect::>(); + + if data.is_empty() { + last_err = Some(Error::NotFound); + continue; + } + + tracing::info!(?data, "got srv record(s)"); + + return Ok(data); + } + + Err(last_err.unwrap_or(Error::NotFound)) +} + +pub async fn resolve_host(host: &str, port: u16) -> Result, Error> { + if let Ok(ip) = host.parse::() { + return Ok(vec![SocketAddr::new(ip, port)]); + } + + let mut last_err = None; + for resolver in configured_dns_resolvers() { + match resolver { + DnsResolver::System => { + if !ALLOW_USE_SYSTEM_DNS_RESOLVER.load(std::sync::atomic::Ordering::Relaxed) { + continue; + } + + match lookup_host(format!("{}:{}", host, port)).await { + Ok(addrs) => { + let addrs = addrs.collect::>(); + if !addrs.is_empty() { + tracing::debug!(?addrs, "system dns lookup done"); + return Ok(addrs); + } + } + Err(error) => { + tracing::debug!(?error, "system dns lookup failed"); + last_err = Some(Error::from(error)); + } + } + } + DnsResolver::Hickory { uri, resolver } => match resolver.lookup_ip(host).await { + Ok(lookup) => { + let addrs = lookup + .iter() + .map(|ip| SocketAddr::new(ip, port)) + .collect::>(); + if !addrs.is_empty() { + return Ok(addrs); + } + } + Err(error) => { + tracing::debug!(?uri, ?host, ?error, "hickory dns lookup failed"); + last_err = Some( + anyhow::anyhow!( + "hickory dns lookup_ip failed, host: {}, port: {}", + host, + port + ) + .into(), + ); + } + }, + } + } + + Err(last_err.unwrap_or(Error::NotFound)) } pub async fn socket_addrs( @@ -160,32 +389,7 @@ pub async fn socket_addrs( return Ok(vec![SocketAddr::new(ip, port)]); } - let host = host.to_string(); - - if ALLOW_USE_SYSTEM_DNS_RESOLVER.load(std::sync::atomic::Ordering::Relaxed) { - match lookup_host(format!("{}:{}", host, port)).await { - Ok(addrs) => { - let addrs = addrs.collect(); - tracing::debug!(?addrs, "system dns lookup done"); - return Ok(addrs); - } - Err(error) => { - tracing::error!(?error, "system dns lookup failed"); - } - } - } - - // use hickory_resolver - let ips = RESOLVER.lookup_ip(&host).await.with_context(|| { - format!( - "hickory dns lookup_ip failed, host: {}, port: {}", - host, port - ) - })?; - Ok(ips - .iter() - .map(|ip| SocketAddr::new(ip, port)) - .collect::>()) + resolve_host(&host.to_string(), port).await } #[cfg(test)] @@ -218,6 +422,25 @@ mod tests { } } + #[test] + fn default_dns_resolver_is_system() { + assert_eq!(get_default_dns_resolvers(), vec!["system".to_string()]); + assert!(matches!( + build_dns_resolvers(&get_default_dns_resolvers()).unwrap()[0], + DnsResolver::System + )); + } + + #[test] + fn alidns_doh_resolver_is_supported() { + validate_dns_resolvers(&["https://dns.alidns.com/dns-query".to_string()]).unwrap(); + } + + #[test] + fn unknown_doh_resolver_requires_bootstrap_ip() { + assert!(validate_dns_resolvers(&["https://example.com/dns-query".to_string()]).is_err()); + } + #[tokio::test] async fn test_socket_addrs() { let url = url::Url::parse("tcp://github-ci-test.easytier.cn:80").unwrap();