feat(dns): support configurable resolver chain

This commit is contained in:
fanyang
2026-06-03 22:24:21 +08:00
parent b207b8a7bb
commit e21046a5fa
7 changed files with 455 additions and 82 deletions
Generated
+10
View File
@@ -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",
+1 -1
View File
@@ -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 }
+80 -2
View File
@@ -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<Vec<String>>;
fn set_stun_servers_v6(&self, servers: Option<Vec<String>>);
fn get_dns_resolvers(&self) -> Vec<String> {
dns::get_default_dns_resolvers()
}
fn get_dns_resolvers_config(&self) -> Option<Vec<String>> {
None
}
fn set_dns_resolvers(&self, _resolvers: Option<Vec<String>>) {}
fn get_secure_mode(&self) -> Option<SecureModeConfig>;
fn set_secure_mode(&self, secure_mode: Option<SecureModeConfig>);
@@ -652,6 +660,7 @@ struct Config {
udp_whitelist: Option<Vec<String>>,
stun_servers: Option<Vec<String>>,
stun_servers_v6: Option<Vec<String>>,
dns_resolvers: Option<Vec<String>>,
credential_file: Option<PathBuf>,
source: Option<ConfigSourceConfig>,
@@ -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<String>) {
@@ -1094,6 +1107,23 @@ impl ConfigLoader for TomlConfigLoader {
self.config.lock().unwrap().stun_servers_v6 = servers;
}
fn get_dns_resolvers(&self) -> Vec<String> {
self.config
.lock()
.unwrap()
.dns_resolvers
.clone()
.unwrap_or_else(dns::get_default_dns_resolvers)
}
fn get_dns_resolvers_config(&self) -> Option<Vec<String>> {
self.config.lock().unwrap().dns_resolvers.clone()
}
fn set_dns_resolvers(&self, resolvers: Option<Vec<String>>) {
self.config.lock().unwrap().dns_resolvers = resolvers;
}
fn get_secure_mode(&self) -> Option<SecureModeConfig> {
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();
+6
View File
@@ -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() {
+79 -23
View File
@@ -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<SocketAddr> {
if host.parse::<IpAddr>().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<SocketAddr> {
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::<SocketAddr>() {
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::<IpAddr>() {
(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));
}
}
+1 -1
View File
@@ -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;
+278 -55
View File
@@ -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<Resolver<TokioRuntimeProvider>>,
},
}
pub fn get_default_dns_resolvers() -> Vec<String> {
vec![SYSTEM_DNS_RESOLVER.to_string()]
}
fn bootstrap_ips_for_doh_host(host: &str) -> Option<Vec<IpAddr>> {
if let Ok(ip) = host.parse::<IpAddr>() {
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<Arc<Resolver<TokioRuntimeProvider>>, 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<DnsResolver, Error> {
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<str> = Arc::from(host.to_string());
let http_endpoint: Option<Arc<str>> = 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::<Vec<_>>();
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<DnsResolver, Error> {
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<Vec<DnsResolver>, 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<String>) -> Result<(), Error> {
let resolvers = build_dns_resolvers(&raw_resolvers)?;
*DNS_RESOLVERS.write().unwrap() = resolvers;
Ok(())
}
static DNS_RESOLVERS: Lazy<RwLock<Vec<DnsResolver>>> =
Lazy::new(|| RwLock::new(build_dns_resolvers(&get_default_dns_resolvers()).unwrap()));
fn configured_dns_resolvers() -> Vec<DnsResolver> {
DNS_RESOLVERS.read().unwrap().clone()
}
pub fn sanitize(name: impl AsRef<str>) -> String {
let name = name.as_ref();
let dot = name.ends_with('.');
@@ -91,23 +219,46 @@ static RESOLVER: Lazy<Arc<Resolver<TokioRuntimeProvider>>> = Lazy::new(|| {
});
pub async fn txt_lookup(name: impl IntoName) -> Result<Vec<String>, 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::<Vec<_>>();
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<Vec<String>, Error> {
@@ -120,23 +271,101 @@ pub async fn txt_resolve(name: impl IntoName) -> Result<Vec<String>, Error> {
}
pub async fn srv_lookup(name: impl IntoName) -> Result<Vec<SRV>, 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::<Vec<_>>();
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<Vec<SocketAddr>, Error> {
if let Ok(ip) = host.parse::<IpAddr>() {
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::<Vec<_>>();
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::<Vec<_>>();
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::<Vec<_>>())
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();