diff --git a/easytier/src/common/stun.rs b/easytier/src/common/stun.rs index f2f16b23..4eb0c27e 100644 --- a/easytier/src/common/stun.rs +++ b/easytier/src/common/stun.rs @@ -22,8 +22,8 @@ use stun_codec::{Message, MessageClass, MessageDecoder, MessageEncoder}; use crate::common::error::Error; -use crate::utils::dns::resolve_txt_record; use super::stun_codec_ext::*; +use crate::utils::dns::txt_resolve; const DEFAULT_UDP_STUN_SERVERS: &[&str] = &[ "txt:stun.easytier.cn", @@ -61,11 +61,6 @@ impl HostResolverIter { } } - async fn get_txt_record(domain_name: &str) -> Result, Error> { - let txt_data = resolve_txt_record(domain_name).await?; - Ok(txt_data.split(" ").map(|x| x.to_string()).collect()) - } - #[async_recursion::async_recursion] async fn next(&mut self) -> Option { if self.ips.is_empty() { @@ -82,7 +77,7 @@ impl HostResolverIter { if host.starts_with("txt:") { let domain_name = host.trim_start_matches("txt:"); - match Self::get_txt_record(domain_name).await { + match txt_resolve(domain_name).await { Ok(hosts) => { tracing::info!( ?domain_name, diff --git a/easytier/src/connector/dns_connector.rs b/easytier/src/connector/dns_connector.rs index f4fb357a..ec1b3b4a 100644 --- a/easytier/src/connector/dns_connector.rs +++ b/easytier/src/connector/dns_connector.rs @@ -1,22 +1,17 @@ use std::{net::SocketAddr, sync::Arc}; use super::{create_connector_by_url, http_connector::TunnelWithInfo}; +use crate::utils::dns::{srv_lookup, txt_resolve}; use crate::{ - common::{ - error::Error, - global_ctx::ArcGlobalCtx, - log, - }, + common::{error::Error, global_ctx::ArcGlobalCtx, log}, proto::common::TunnelInfo, tunnel::{IpScheme, IpVersion, Tunnel, TunnelConnector, TunnelError, TunnelScheme}, }; use anyhow::Context; use dashmap::DashSet; -use hickory_proto::rr::RData; use hickory_resolver::proto::rr::rdata::SRV; -use rand::{seq::SliceRandom, Rng as _}; +use rand::{Rng as _, seq::SliceRandom}; use strum::VariantArray; -use crate::utils::dns::{resolve_txt_record, RESOLVER}; fn weighted_choice(options: &[(T, u64)]) -> Option<&T> { let total_weight = options.iter().map(|(_, weight)| *weight).sum(); @@ -59,14 +54,13 @@ impl DnsTunnelConnector { &self, domain_name: &str, ) -> Result, Error> { - let txt_data = resolve_txt_record(domain_name) + let txt_data = txt_resolve(domain_name) .await .with_context(|| format!("resolve txt record failed, domain_name: {}", domain_name))?; let candidate_urls = txt_data - .split(" ") - .map(|s| s.to_string()) - .filter_map(|s| url::Url::parse(s.as_str()).ok()) + .iter() + .filter_map(|s| url::Url::parse(s).ok()) .collect::>(); // shuffle candidate_urls and get the first one @@ -74,7 +68,7 @@ impl DnsTunnelConnector { .choose(&mut rand::thread_rng()) .with_context(|| { format!( - "no valid url found, txt_data: {}, expecting an url list splitted by space", + "no valid url found, txt_data: {:?}, expecting an url list split by space", txt_data ) })?; @@ -84,7 +78,7 @@ impl DnsTunnelConnector { Ok(connector) } - fn handle_one_srv_record(record: &SRV, protocol: IpScheme) -> Result<(url::Url, u64), Error> { + fn handle_one_srv_record(record: SRV, protocol: IpScheme) -> Result<(url::Url, u64), Error> { // port must be non-zero if record.port == 0 { return Err(anyhow::anyhow!("port must be non-zero").into()); @@ -120,21 +114,9 @@ impl DnsTunnelConnector { let srv_lookup_tasks = srv_domains .iter() .map(|(protocol, srv_domain)| { - let resolver = RESOLVER.clone(); let responses = responses.clone(); async move { - let response = resolver.srv_lookup(srv_domain).await.with_context(|| { - format!("srv_lookup failed, srv_domain: {}", srv_domain) - })?; - tracing::info!(?response, ?srv_domain, "srv_lookup response"); - for record in response - .answers() - .iter() - .filter_map(|record| match record.data { - RData::SRV(ref srv) => Some(srv), - _ => None, - }) - { + for record in srv_lookup(srv_domain).await? { let parsed_record = Self::handle_one_srv_record(record, **protocol); tracing::info!(?parsed_record, ?srv_domain, "parsed_record"); if let Err(e) = &parsed_record { diff --git a/easytier/src/utils/dns.rs b/easytier/src/utils/dns.rs index 8533edf3..c3855404 100644 --- a/easytier/src/utils/dns.rs +++ b/easytier/src/utils/dns.rs @@ -1,73 +1,104 @@ use std::net::SocketAddr; -use std::sync::Arc; use std::sync::atomic::AtomicBool; +use std::sync::Arc; use anyhow::Context; use hickory_net::runtime::TokioRuntimeProvider; -use hickory_proto::rr::RData; +use hickory_proto::rr::{IntoName, RData}; +use hickory_proto::rr::rdata::SRV; use hickory_resolver::config::{ ConnectionConfig, LookupIpStrategy, NameServerConfig, ResolverConfig, ResolverOpts, }; use hickory_resolver::system_conf::read_system_conf; use hickory_resolver::{Resolver, TokioResolver}; use once_cell::sync::Lazy; +use std::net::SocketAddr; +use std::sync::Arc; +use std::sync::atomic::AtomicBool; use tokio::net::lookup_host; -use super::error::Error; +use crate::common::error::Error; pub fn get_default_resolver_config() -> ResolverConfig { - let mut default_resolve_config = ResolverConfig::default(); - default_resolve_config.add_name_server(NameServerConfig::new( - "223.5.5.5".parse().unwrap(), - true, - vec![ConnectionConfig::udp()], - )); - default_resolve_config.add_name_server(NameServerConfig::new( - "180.184.1.1".parse().unwrap(), - true, - vec![ConnectionConfig::udp()], - )); - default_resolve_config + let mut config = ResolverConfig::default(); + for server in ["223.5.5.5", "180.184.1.1"] { + config.add_name_server(NameServerConfig::new( + server.parse().unwrap(), + true, + vec![ConnectionConfig::udp()], + )); + } + config } -pub static ALLOW_USE_SYSTEM_DNS_RESOLVER: Lazy = Lazy::new(|| AtomicBool::new(true)); +pub static ALLOW_USE_SYSTEM_DNS_RESOLVER: AtomicBool = AtomicBool::new(true); pub static RESOLVER: Lazy>> = Lazy::new(|| { - let system_cfg = read_system_conf(); let mut cfg = get_default_resolver_config(); let mut opt = ResolverOpts::default(); - if let Ok(s) = system_cfg { - for ns in s.0.name_servers() { + if let Ok((sys_cfg, sys_opt)) = read_system_conf() { + for ns in sys_cfg.name_servers() { cfg.add_name_server(ns.clone()); } - opt = s.1; + opt = sys_opt; } opt.ip_strategy = LookupIpStrategy::Ipv4AndIpv6; let builder = TokioResolver::builder_with_config(cfg, TokioRuntimeProvider::default()).with_options(opt); - Arc::new(builder.build().unwrap()) + Arc::new( + builder + .build() + .expect("failed to initialize global DNS resolver"), + ) }); -pub async fn resolve_txt_record(domain_name: &str) -> Result { - let r = RESOLVER.clone(); - let response = r - .txt_lookup(domain_name) +pub async fn txt_lookup(name: impl IntoName) -> Result, Error> { + let response = RESOLVER + .txt_lookup(name) .await - .with_context(|| format!("txt_lookup failed, domain_name: {}", domain_name))?; + .context("failed to lookup txt record")?; - let txt_data = response + let data = response .answers() .iter() .filter_map(|record| match record.data { - RData::TXT(ref txt) => Some(txt), + RData::TXT(ref txt) => Some(txt.to_string()), _ => None, }) - .next() - .with_context(|| format!("no txt record found, domain_name: {}", domain_name))?; + .collect(); - tracing::info!(?txt_data, ?domain_name, "get txt record"); + tracing::info!(?data, "got txt record(s)"); - Ok(txt_data.to_string()) + Ok(data) +} + +pub async fn txt_resolve(name: impl IntoName) -> Result, Error> { + Ok(txt_lookup(name) + .await? + .iter() + .flat_map(|s| s.split_whitespace()) + .map(String::from) + .collect()) +} + +pub async fn srv_lookup(name: impl IntoName) -> Result, Error> { + let response = RESOLVER + .srv_lookup(name) + .await + .context("failed to lookup srv record")?; + + let data = response + .answers() + .iter() + .filter_map(|record| match record.data { + RData::SRV(ref srv) => Some(srv.clone()), + _ => None, + }) + .collect(); + + tracing::info!(?data, "got srv record(s)"); + + Ok(data) } pub async fn socket_addrs(