rewrite txt_lookup and srv_lookup

rewrite

rewrite

lookup

l
This commit is contained in:
Luna Yao
2026-04-19 00:34:41 +02:00
parent 5a516a195c
commit 7ae7725fbf
3 changed files with 74 additions and 66 deletions
+2 -7
View File
@@ -22,8 +22,8 @@ use stun_codec::{Message, MessageClass, MessageDecoder, MessageEncoder};
use crate::common::error::Error; use crate::common::error::Error;
use crate::utils::dns::resolve_txt_record;
use super::stun_codec_ext::*; use super::stun_codec_ext::*;
use crate::utils::dns::txt_resolve;
const DEFAULT_UDP_STUN_SERVERS: &[&str] = &[ const DEFAULT_UDP_STUN_SERVERS: &[&str] = &[
"txt:stun.easytier.cn", "txt:stun.easytier.cn",
@@ -61,11 +61,6 @@ impl HostResolverIter {
} }
} }
async fn get_txt_record(domain_name: &str) -> Result<Vec<String>, 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_recursion::async_recursion]
async fn next(&mut self) -> Option<SocketAddr> { async fn next(&mut self) -> Option<SocketAddr> {
if self.ips.is_empty() { if self.ips.is_empty() {
@@ -82,7 +77,7 @@ impl HostResolverIter {
if host.starts_with("txt:") { if host.starts_with("txt:") {
let domain_name = host.trim_start_matches("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) => { Ok(hosts) => {
tracing::info!( tracing::info!(
?domain_name, ?domain_name,
+9 -27
View File
@@ -1,22 +1,17 @@
use std::{net::SocketAddr, sync::Arc}; use std::{net::SocketAddr, sync::Arc};
use super::{create_connector_by_url, http_connector::TunnelWithInfo}; use super::{create_connector_by_url, http_connector::TunnelWithInfo};
use crate::utils::dns::{srv_lookup, txt_resolve};
use crate::{ use crate::{
common::{ common::{error::Error, global_ctx::ArcGlobalCtx, log},
error::Error,
global_ctx::ArcGlobalCtx,
log,
},
proto::common::TunnelInfo, proto::common::TunnelInfo,
tunnel::{IpScheme, IpVersion, Tunnel, TunnelConnector, TunnelError, TunnelScheme}, tunnel::{IpScheme, IpVersion, Tunnel, TunnelConnector, TunnelError, TunnelScheme},
}; };
use anyhow::Context; use anyhow::Context;
use dashmap::DashSet; use dashmap::DashSet;
use hickory_proto::rr::RData;
use hickory_resolver::proto::rr::rdata::SRV; use hickory_resolver::proto::rr::rdata::SRV;
use rand::{seq::SliceRandom, Rng as _}; use rand::{Rng as _, seq::SliceRandom};
use strum::VariantArray; use strum::VariantArray;
use crate::utils::dns::{resolve_txt_record, RESOLVER};
fn weighted_choice<T>(options: &[(T, u64)]) -> Option<&T> { fn weighted_choice<T>(options: &[(T, u64)]) -> Option<&T> {
let total_weight = options.iter().map(|(_, weight)| *weight).sum(); let total_weight = options.iter().map(|(_, weight)| *weight).sum();
@@ -59,14 +54,13 @@ impl DnsTunnelConnector {
&self, &self,
domain_name: &str, domain_name: &str,
) -> Result<Box<dyn TunnelConnector>, Error> { ) -> Result<Box<dyn TunnelConnector>, Error> {
let txt_data = resolve_txt_record(domain_name) let txt_data = txt_resolve(domain_name)
.await .await
.with_context(|| format!("resolve txt record failed, domain_name: {}", domain_name))?; .with_context(|| format!("resolve txt record failed, domain_name: {}", domain_name))?;
let candidate_urls = txt_data let candidate_urls = txt_data
.split(" ") .iter()
.map(|s| s.to_string()) .filter_map(|s| url::Url::parse(s).ok())
.filter_map(|s| url::Url::parse(s.as_str()).ok())
.collect::<Vec<_>>(); .collect::<Vec<_>>();
// shuffle candidate_urls and get the first one // shuffle candidate_urls and get the first one
@@ -74,7 +68,7 @@ impl DnsTunnelConnector {
.choose(&mut rand::thread_rng()) .choose(&mut rand::thread_rng())
.with_context(|| { .with_context(|| {
format!( 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 txt_data
) )
})?; })?;
@@ -84,7 +78,7 @@ impl DnsTunnelConnector {
Ok(connector) 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 // port must be non-zero
if record.port == 0 { if record.port == 0 {
return Err(anyhow::anyhow!("port must be non-zero").into()); return Err(anyhow::anyhow!("port must be non-zero").into());
@@ -120,21 +114,9 @@ impl DnsTunnelConnector {
let srv_lookup_tasks = srv_domains let srv_lookup_tasks = srv_domains
.iter() .iter()
.map(|(protocol, srv_domain)| { .map(|(protocol, srv_domain)| {
let resolver = RESOLVER.clone();
let responses = responses.clone(); let responses = responses.clone();
async move { async move {
let response = resolver.srv_lookup(srv_domain).await.with_context(|| { for record in srv_lookup(srv_domain).await? {
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,
})
{
let parsed_record = Self::handle_one_srv_record(record, **protocol); let parsed_record = Self::handle_one_srv_record(record, **protocol);
tracing::info!(?parsed_record, ?srv_domain, "parsed_record"); tracing::info!(?parsed_record, ?srv_domain, "parsed_record");
if let Err(e) = &parsed_record { if let Err(e) = &parsed_record {
+63 -32
View File
@@ -1,73 +1,104 @@
use std::net::SocketAddr; use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::AtomicBool; use std::sync::atomic::AtomicBool;
use std::sync::Arc;
use anyhow::Context; use anyhow::Context;
use hickory_net::runtime::TokioRuntimeProvider; 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::{ use hickory_resolver::config::{
ConnectionConfig, LookupIpStrategy, NameServerConfig, ResolverConfig, ResolverOpts, ConnectionConfig, LookupIpStrategy, NameServerConfig, ResolverConfig, ResolverOpts,
}; };
use hickory_resolver::system_conf::read_system_conf; use hickory_resolver::system_conf::read_system_conf;
use hickory_resolver::{Resolver, TokioResolver}; use hickory_resolver::{Resolver, TokioResolver};
use once_cell::sync::Lazy; use once_cell::sync::Lazy;
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use tokio::net::lookup_host; use tokio::net::lookup_host;
use super::error::Error; use crate::common::error::Error;
pub fn get_default_resolver_config() -> ResolverConfig { pub fn get_default_resolver_config() -> ResolverConfig {
let mut default_resolve_config = ResolverConfig::default(); let mut config = ResolverConfig::default();
default_resolve_config.add_name_server(NameServerConfig::new( for server in ["223.5.5.5", "180.184.1.1"] {
"223.5.5.5".parse().unwrap(), config.add_name_server(NameServerConfig::new(
true, server.parse().unwrap(),
vec![ConnectionConfig::udp()], true,
)); vec![ConnectionConfig::udp()],
default_resolve_config.add_name_server(NameServerConfig::new( ));
"180.184.1.1".parse().unwrap(), }
true, config
vec![ConnectionConfig::udp()],
));
default_resolve_config
} }
pub static ALLOW_USE_SYSTEM_DNS_RESOLVER: Lazy<AtomicBool> = Lazy::new(|| AtomicBool::new(true)); pub static ALLOW_USE_SYSTEM_DNS_RESOLVER: AtomicBool = AtomicBool::new(true);
pub static RESOLVER: Lazy<Arc<Resolver<TokioRuntimeProvider>>> = Lazy::new(|| { pub static RESOLVER: Lazy<Arc<Resolver<TokioRuntimeProvider>>> = Lazy::new(|| {
let system_cfg = read_system_conf();
let mut cfg = get_default_resolver_config(); let mut cfg = get_default_resolver_config();
let mut opt = ResolverOpts::default(); let mut opt = ResolverOpts::default();
if let Ok(s) = system_cfg { if let Ok((sys_cfg, sys_opt)) = read_system_conf() {
for ns in s.0.name_servers() { for ns in sys_cfg.name_servers() {
cfg.add_name_server(ns.clone()); cfg.add_name_server(ns.clone());
} }
opt = s.1; opt = sys_opt;
} }
opt.ip_strategy = LookupIpStrategy::Ipv4AndIpv6; opt.ip_strategy = LookupIpStrategy::Ipv4AndIpv6;
let builder = let builder =
TokioResolver::builder_with_config(cfg, TokioRuntimeProvider::default()).with_options(opt); 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<String, Error> { pub async fn txt_lookup(name: impl IntoName) -> Result<Vec<String>, Error> {
let r = RESOLVER.clone(); let response = RESOLVER
let response = r .txt_lookup(name)
.txt_lookup(domain_name)
.await .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() .answers()
.iter() .iter()
.filter_map(|record| match record.data { .filter_map(|record| match record.data {
RData::TXT(ref txt) => Some(txt), RData::TXT(ref txt) => Some(txt.to_string()),
_ => None, _ => None,
}) })
.next() .collect();
.with_context(|| format!("no txt record found, domain_name: {}", domain_name))?;
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<Vec<String>, 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<Vec<SRV>, 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( pub async fn socket_addrs(